1use async_trait::async_trait;
2use futures::stream::{Stream, StreamExt};
3use parking_lot::RwLock;
4use serde_json::Value;
5use std::collections::{HashMap, HashSet};
6use std::future::Future;
7use std::pin::Pin;
8use std::sync::atomic::{AtomicBool, AtomicU64, Ordering};
9use std::sync::{Arc, Weak};
10use std::time::{Duration, Instant};
11use tracing::{debug, error, info, instrument, warn};
12
13const DISAMBIGUATION_STATE_GENERATION_KEY: &str = "_runtime.disambiguation_state_generation";
14const MAX_TOOL_FALLBACK_HOPS: usize = 16;
15
16pub(crate) type RootTurnGate = Arc<tokio::sync::Mutex<()>>;
18
19pub(crate) type RootTurnGateIdentityStack = Arc<[RootTurnGate]>;
21
22tokio::task_local! {
23 static RUNTIME_GATE_IDENTITY_STACK: RootTurnGateIdentityStack;
24}
25
26pub(crate) fn current_runtime_gate_identity_stack() -> RootTurnGateIdentityStack {
28 RUNTIME_GATE_IDENTITY_STACK
29 .try_with(Arc::clone)
30 .unwrap_or_default()
31}
32
33pub(crate) async fn scope_runtime_gate_identity_stack<F, T>(
35 identity_stack: &RootTurnGateIdentityStack,
36 future: F,
37) -> T
38where
39 F: Future<Output = T>,
40{
41 RUNTIME_GATE_IDENTITY_STACK
42 .scope(Arc::clone(identity_stack), future)
43 .await
44}
45
46pub(crate) type ToolResourceLocks = Arc<RwLock<HashMap<String, Weak<tokio::sync::Mutex<()>>>>>;
48
49struct ToolResourceGuards {
53 guards: Vec<tokio::sync::OwnedMutexGuard<()>>,
54 locks: ToolResourceLocks,
55}
56
57struct RootTurnAdmission {
61 guard: tokio::sync::OwnedMutexGuard<()>,
62 identity_stack: RootTurnGateIdentityStack,
63}
64
65#[derive(Clone)]
66struct StoredSessionRestore {
67 snapshot: AgentSnapshot,
68 metadata: Option<ai_agents_core::SessionMetadata>,
69}
70
71struct RuntimeSessionRestorePoint {
72 snapshot: AgentSnapshot,
73 metadata: ai_agents_core::SessionMetadata,
74 actor_id: Option<String>,
75 session_id: Option<String>,
76}
77
78impl Drop for ToolResourceGuards {
79 fn drop(&mut self) {
80 self.guards.clear();
81 self.locks.write().retain(|_, lock| lock.strong_count() > 0);
82 }
83}
84
85#[derive(Clone)]
89struct RuntimeSafetySnapshot {
90 version: u64,
91 emergency_deny: bool,
92 tool_security: ToolSecurityEngine,
93 tool_scope_override: Option<Vec<String>>,
94}
95
96#[derive(Clone, Copy)]
100struct ToolDecisionVersions {
101 policy: u64,
102 registry: u64,
103 runtime_control: u64,
104 state: Option<u64>,
105}
106
107#[derive(Clone, Debug, Default)]
111struct ToolFallbackState {
112 visited_canonical_ids: Vec<String>,
113}
114
115impl ToolFallbackState {
116 fn rejection_reason(&self, canonical_id: &str) -> Option<String> {
120 if self
121 .visited_canonical_ids
122 .iter()
123 .any(|visited| visited == canonical_id)
124 {
125 return Some(format!(
126 "Tool fallback cycle detected at '{canonical_id}' after [{}]",
127 self.visited_canonical_ids.join(" -> ")
128 ));
129 }
130 if self.visited_canonical_ids.len() > MAX_TOOL_FALLBACK_HOPS {
131 return Some(format!(
132 "Tool fallback chain exceeds the maximum of {MAX_TOOL_FALLBACK_HOPS} hops"
133 ));
134 }
135 None
136 }
137
138 fn with_current(mut self, canonical_id: String) -> Self {
142 self.visited_canonical_ids.push(canonical_id);
143 self
144 }
145
146 fn final_rejection_reason(
150 &self,
151 admitted_canonical_id: &str,
152 final_canonical_id: &str,
153 ) -> Option<String> {
154 if admitted_canonical_id == final_canonical_id {
155 return None;
156 }
157 if self
158 .visited_canonical_ids
159 .iter()
160 .any(|visited| visited == final_canonical_id)
161 {
162 return Some(format!(
163 "Tool fallback cycle detected after final resolution changed '{admitted_canonical_id}' to '{final_canonical_id}'"
164 ));
165 }
166 Some(format!(
167 "Tool canonical target changed after initial admission from '{admitted_canonical_id}' to '{final_canonical_id}'"
168 ))
169 }
170}
171
172#[derive(Clone, Copy, Debug)]
176struct ValidatedToolTimeout {
177 timer: Duration,
178 deadline_delta: chrono::Duration,
179}
180
181struct AvailableToolIdsSnapshot {
185 tool_ids: Vec<String>,
186 state_generation: Option<u64>,
187}
188
189#[derive(Clone)]
193struct ToolApprovalBinding {
194 canonical_id: String,
195 arguments: Value,
196 confirmation_required: bool,
197 policy_version: u64,
198 runtime_control_version: u64,
199 state_generation: Option<u64>,
200 reviewed_tool: Arc<dyn ai_agents_core::Tool>,
201}
202
203fn merge_approved_record(record: &mut Option<ToolApprovalRecord>) {
207 if record
208 .as_ref()
209 .is_some_and(|record| matches!(record.status, ToolApprovalStatus::Modified))
210 {
211 return;
212 }
213 *record = Some(ToolApprovalRecord {
214 status: ToolApprovalStatus::Approved,
215 reason: None,
216 modified_arguments: None,
217 });
218}
219
220impl ToolApprovalBinding {
221 fn is_stale(
223 &self,
224 canonical_id: &str,
225 arguments: &Value,
226 confirmation_required: bool,
227 versions: ToolDecisionVersions,
228 resolved_tool: &Arc<dyn ai_agents_core::Tool>,
229 ) -> bool {
230 self.canonical_id != canonical_id
231 || self.arguments != *arguments
232 || self.confirmation_required != confirmation_required
233 || self.policy_version != versions.policy
234 || self.runtime_control_version != versions.runtime_control
235 || self.state_generation != versions.state
236 || !Arc::ptr_eq(&self.reviewed_tool, resolved_tool)
237 }
238}
239
240use crate::turn_context::{current_turn_actor_context, scope_actor_context};
241
242use ai_agents_context::{ContextManager, ContextProvider, TemplateRenderer};
243use ai_agents_core::traits::storage::StorageCapability;
244use ai_agents_core::{
245 AgentError, AgentSnapshot, AgentStorage, ChatMessage, FinishReason, LLMChunk, LLMError,
246 LLMFeature, LLMProvider, LLMResponse, LLMToolDefinition, LLMToolRequest, PermissionOutcome,
247 Result, ToolActorContext, ToolApprovalRecord, ToolApprovalStatus, ToolCallClassification,
248 ToolCallSource, ToolCancellationToken, ToolChoice, ToolExecutionContext, ToolExecutionLimits,
249 ToolExecutionRecord, ToolExecutionRequest, ToolInvoker, ToolPolicyDecisionRecord, ToolResult,
250 ToolSafetyMetadata, decode_native_tool_call_markers, encode_native_tool_call_markers,
251 encode_native_tool_result_marker, inspect_native_history, native_readable_projection,
252};
253use ai_agents_disambiguation::{
254 AmbiguityDetectionResult, ClarificationObserver, ClarificationParseFuture,
255 ClarificationQuestion, ClarificationQuestionFuture, ConfirmationParseFuture,
256 DisambiguationConfig, DisambiguationContext, DisambiguationManager, DisambiguationResult,
257};
258use ai_agents_hitl::{
259 ApprovalHandler, ApprovalResolvedOutcome, ApprovalResult, ApprovalTrigger, HITLCheckResult,
260 HITLEngine, RejectAllHandler, TimeoutAction,
261};
262use ai_agents_hooks::{AgentHooks, NoopHooks};
263use ai_agents_llm::LLMRegistry;
264use ai_agents_memory::{
265 CompressResult, EvictionReason, Memory, MemoryBudgetEvent, MemoryCompressEvent,
266 MemoryEvictEvent, MemoryTokenBudget, OverflowStrategy,
267};
268use ai_agents_observability::{
269 EventStatus, EventType, ObservabilityManager, ObservationPurpose, SpanContext,
270 current_observation_context, new_session_id as new_observation_session_id,
271 resolve_language_from_context, with_observation_context, with_observation_purpose,
272};
273use ai_agents_process::{
274 ProcessData, ProcessProcessor, ProcessPurposeHint, ProcessStageFuture, ProcessStageObserver,
275};
276use ai_agents_reasoning::{
277 CriterionResult, EvaluationResult, Plan, PlanAction, PlanStatus, PlanStep, ReasoningConfig,
278 ReasoningMetadata, ReasoningMode, ReasoningOutput, ReflectionAttempt, ReflectionConfig,
279 ReflectionMetadata, StepFailureAction,
280};
281use ai_agents_recovery::{
282 ByRoleFilter, ContextOverflowAction, FilterConfig, KeepRecentFilter, LLMFailureAction,
283 MessageFilter, RecoveryManager, SkipPatternFilter, ToolFailureAction,
284};
285use ai_agents_relationships::RelationshipManager;
286use ai_agents_skills::{SkillDefinition, SkillExecutor, SkillRouter};
287use ai_agents_state::{
288 PromptMode, StateAction, StateMachine, StateMachineSnapshot, StateTransitionEvent, Transition,
289 TransitionContext, TransitionEvaluator, TransitionTiming, evaluate_guard,
290};
291use ai_agents_storage::{StorageConfig as StorageStorageConfig, create_storage};
292use ai_agents_tools::{
293 CommandRunner, ConditionEvaluator, DiagnosticsProvider, EvaluationContext, LLMGetter,
294 MAX_TOOL_TIMEOUT_MS, QuestionHandler, SecurityCheckResult, TodoItem, ToolCallRecord,
295 ToolRegistry, ToolSecurityConfig, ToolSecurityEngine,
296};
297
298use super::{
299 Agent, AgentInfo, AgentResponse, AgentStreamEvent, ParallelToolsConfig, StreamChunk,
300 StreamingConfig, ToolCall,
301};
302use crate::optimization::{
303 AwaitBeforeNextTurn, BackgroundMaintenanceQueue, BackgroundOverflowPolicy, MainResponseDraft,
304 MaintenanceMode, MaintenanceSequenceKey, RuntimeBranch, RuntimeBranchResult,
305 RuntimeBranchStatus, RuntimeCommitBehavior, RuntimeConfig, RuntimeOptimizationKind,
306 RuntimeTaskPriority, RuntimeTaskPurpose, ScheduledBranchSet, SkillCandidate,
307 StreamingDraftResult, TransitionCandidate, TurnBranchScheduler, TurnOptimizationContext,
308};
309use crate::spec::StorageConfig;
310
311enum ToolCallOutcome {
313 Continue,
315 TransitionFired,
317 Rejected(AgentResponse),
319}
320
321#[derive(Clone)]
322struct MainToolProtocol {
323 choice: Option<ToolChoice>,
324 tool_ids: Vec<String>,
325 definitions: Vec<LLMToolDefinition>,
326}
327
328struct MainProviderResponse {
329 response: LLMResponse,
330 used_native_tools: bool,
331}
332
333enum MainStreamSource {
338 Stream(Box<dyn Stream<Item = std::result::Result<LLMChunk, LLMError>> + Unpin + Send>),
339 StaticResponse(String),
340}
341
342#[derive(Clone)]
344struct ActiveNativeExchange {
345 exchange_id: String,
346 call_ids: Vec<String>,
347}
348
349struct CommittedTextResponse<'a> {
353 processed_input: &'a str,
354 input_context: &'a HashMap<String, Value>,
355 answer: String,
356 reasoning_mode: ReasoningMode,
357 auto_detected: bool,
358 iterations: u32,
359 thinking_content: Option<String>,
360 all_tool_calls: Vec<ToolCall>,
361}
362
363struct AgentResponseParts {
367 content: String,
368 all_tool_calls: Vec<ToolCall>,
369 reasoning_mode: ReasoningMode,
370 auto_detected: bool,
371 iterations: u32,
372 thinking: Option<String>,
373 reflection_metadata: Option<ReflectionMetadata>,
374}
375
376type RuntimeStreamTerminalSlot = Arc<RwLock<Option<AgentResponse>>>;
380
381fn new_runtime_stream_terminal_slot() -> RuntimeStreamTerminalSlot {
385 Arc::new(RwLock::new(None))
386}
387
388fn record_runtime_stream_final(slot: &RuntimeStreamTerminalSlot, response: AgentResponse) {
392 *slot.write() = Some(response);
393}
394
395#[derive(Clone, Copy)]
396struct DisambiguationOwnership {
397 epoch: u64,
398 state_generation: Option<u64>,
399}
400
401enum SkillRouteResult {
403 NoMatch,
405 Response { skill_id: String, content: String },
407 NeedsClarification {
409 response: AgentResponse,
410 ownership: Option<DisambiguationOwnership>,
411 },
412}
413
414enum ParallelTransitionSelection {
416 Candidate(TransitionCandidate),
418 NoMatch,
420 ReservationExhausted,
422}
423
424enum DisambiguationDispatch {
430 Proceed(String),
432 Terminal(AgentResponse),
434 RecheckSkill {
436 skill_id: String,
437 enriched_input: String,
438 disambiguation_epoch: u64,
439 state_generation: Option<u64>,
440 },
441}
442
443enum PostLoopResult {
444 NoTransition(String),
446 Transitioned { content: String, regenerated: bool },
449 NeedsRedispatch,
452}
453
454struct AppliedPostLoop {
456 content: String,
457 transitioned: bool,
458 regenerated: bool,
460}
461
462struct StateTransitionReservation<'a> {
463 reserved: &'a AtomicBool,
464}
465
466impl Drop for StateTransitionReservation<'_> {
467 fn drop(&mut self) {
468 self.reserved.store(false, Ordering::SeqCst);
469 }
470}
471
472struct RootTurnCleanup<'a> {
473 agent: &'a RuntimeAgent,
474}
475
476impl<'a> RootTurnCleanup<'a> {
477 fn new(agent: &'a RuntimeAgent) -> Self {
478 Self { agent }
479 }
480}
481
482impl Drop for RootTurnCleanup<'_> {
483 fn drop(&mut self) {
484 self.agent.end_root_turn();
485 }
486}
487
488#[derive(Debug)]
490struct RuntimeControlState {
491 snapshot_guard: RwLock<()>,
493 version: AtomicU64,
495 emergency_deny: Arc<AtomicBool>,
497 tool_security_override: RwLock<Option<ToolSecurityEngine>>,
499 tool_scope_override: RwLock<Option<Vec<String>>>,
501}
502
503impl Default for RuntimeControlState {
504 fn default() -> Self {
505 Self {
506 snapshot_guard: RwLock::new(()),
507 version: AtomicU64::new(1),
508 emergency_deny: Arc::new(AtomicBool::new(false)),
509 tool_security_override: RwLock::new(None),
510 tool_scope_override: RwLock::new(None),
511 }
512 }
513}
514
515#[derive(Clone)]
517pub struct RuntimeControlHandle {
518 state: Arc<RuntimeControlState>,
519}
520
521impl RuntimeControlHandle {
522 pub fn version(&self) -> u64 {
524 self.state.version.load(Ordering::SeqCst)
525 }
526
527 fn bump(&self) -> u64 {
528 self.state.version.fetch_add(1, Ordering::SeqCst) + 1
529 }
530
531 pub fn set_tool_security(&self, config: ToolSecurityConfig) -> u64 {
533 self.try_set_tool_security(config)
534 .expect("invalid tool security configuration")
535 }
536
537 pub fn try_set_tool_security(&self, config: ToolSecurityConfig) -> Result<u64> {
539 config.validate()?;
540 let _guard = self.state.snapshot_guard.write();
541 let generation = self.bump();
542 *self.state.tool_security_override.write() = Some(
543 ToolSecurityEngine::new_with_policy_version(config, generation),
544 );
545 Ok(generation)
546 }
547
548 pub fn clear_tool_security_override(&self) -> u64 {
550 let _guard = self.state.snapshot_guard.write();
551 *self.state.tool_security_override.write() = None;
552 self.bump()
553 }
554
555 pub fn set_tool_scope(&self, tool_ids: Vec<String>) -> u64 {
557 let _guard = self.state.snapshot_guard.write();
558 *self.state.tool_scope_override.write() = Some(tool_ids);
559 self.bump()
560 }
561
562 pub fn clear_tool_scope_override(&self) -> u64 {
564 let _guard = self.state.snapshot_guard.write();
565 *self.state.tool_scope_override.write() = None;
566 self.bump()
567 }
568
569 pub fn set_emergency_deny(&self, enabled: bool) -> u64 {
571 let _guard = self.state.snapshot_guard.write();
572 self.state.emergency_deny.store(enabled, Ordering::SeqCst);
573 self.bump()
574 }
575
576 pub fn cancel_all(&self) -> u64 {
578 self.set_emergency_deny(true)
579 }
580}
581
582pub struct RuntimeAgent {
583 info: AgentInfo,
584 llm_registry: Arc<LLMRegistry>,
585 memory: Arc<dyn Memory>,
586 tools: Arc<ToolRegistry>,
587 skills: Vec<SkillDefinition>,
588 skill_router: Option<SkillRouter>,
589 skill_executor: Option<SkillExecutor>,
590 base_system_prompt: String,
591 max_iterations: u32,
592 iteration_count: RwLock<u32>,
593 max_context_tokens: u32,
594 memory_token_budget: Option<MemoryTokenBudget>,
595 recovery_manager: RecoveryManager,
596 tool_security: ToolSecurityEngine,
597 process_processor: Option<ProcessProcessor>,
598 message_filters: RwLock<HashMap<String, Arc<dyn MessageFilter>>>,
599 state_machine: Option<Arc<StateMachine>>,
600 transition_evaluator: Option<Arc<dyn TransitionEvaluator>>,
601 context_manager: Arc<ContextManager>,
602 template_renderer: TemplateRenderer,
603 tool_call_history: RwLock<Vec<ToolCallRecord>>,
604 parallel_tools: ParallelToolsConfig,
605 streaming: StreamingConfig,
606 hooks: Arc<dyn AgentHooks>,
607 hitl_engine: Option<HITLEngine>,
608 approval_handler: Arc<dyn ApprovalHandler>,
609 storage_config: StorageConfig,
610 storage: RwLock<Option<Arc<dyn AgentStorage>>>,
611 storage_init: tokio::sync::Mutex<()>,
612 reasoning_config: ReasoningConfig,
613 reflection_config: ReflectionConfig,
614 disambiguation_manager: Option<DisambiguationManager>,
615 disambiguation_epoch: AtomicU64,
617 disambiguation_admission: tokio::sync::RwLock<()>,
619 state_transition_reserved: AtomicBool,
621 persona_manager: Option<Arc<ai_agents_persona::PersonaManager>>,
623 pending_skill_id: RwLock<Option<String>>,
627 current_plan: RwLock<Option<Plan>>,
628 declared_tool_ids: Option<Vec<String>>,
630 context_initialized: AtomicBool,
632 spawner: Option<Arc<crate::spawner::AgentSpawner>>,
634 spawner_registry: Option<Arc<crate::spawner::AgentRegistry>>,
636 redispatch_depth: RwLock<u32>,
639 active_turn_context: RwLock<Option<TurnOptimizationContext>>,
641 root_user_message_committed: AtomicBool,
643 active_native_exchanges: RwLock<Vec<ActiveNativeExchange>>,
645 actor_id: RwLock<Option<String>>,
647 fact_store: RwLock<Option<Arc<ai_agents_facts::FactStore>>>,
649 fact_extractor: RwLock<Option<Arc<dyn ai_agents_facts::FactExtractor>>>,
652 actor_facts_cache: Arc<RwLock<HashMap<String, Vec<ai_agents_core::KeyFact>>>>,
654 messages_since_extraction: Arc<RwLock<usize>>,
656 actor_memory_config: Option<ai_agents_facts::ActorMemoryConfig>,
658 facts_config: Option<ai_agents_facts::FactsConfig>,
660 session_metadata: RwLock<ai_agents_core::SessionMetadata>,
662 current_session_id: RwLock<Option<String>>,
664 relationship_manager: Option<Arc<RelationshipManager>>,
666 observability_manager: Option<Arc<ObservabilityManager>>,
668 runtime_config: RuntimeConfig,
670 background_maintenance: Arc<BackgroundMaintenanceQueue>,
672 resource_locks: ToolResourceLocks,
674 runtime_control: Arc<RuntimeControlState>,
676 root_turn_gate: RootTurnGate,
678}
679
680impl std::fmt::Debug for RuntimeAgent {
681 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
682 f.debug_struct("RuntimeAgent")
683 .field("info", &self.info)
684 .field("base_system_prompt", &self.base_system_prompt)
685 .field("max_iterations", &self.max_iterations)
686 .field("skills_count", &self.skills.len())
687 .field("max_context_tokens", &self.max_context_tokens)
688 .field("has_state_machine", &self.state_machine.is_some())
689 .field("parallel_tools", &self.parallel_tools)
690 .field("streaming", &self.streaming)
691 .field("has_hooks", &true)
692 .field("has_hitl", &self.hitl_engine.is_some())
693 .field("storage_type", &self.storage_config.storage_type())
694 .field("reasoning_mode", &self.reasoning_config.mode)
695 .field("reflection_enabled", &self.reflection_config.enabled)
696 .field("declared_tool_ids", &self.declared_tool_ids)
697 .field("has_persona", &self.persona_manager.is_some())
698 .field("has_observability", &self.observability_manager.is_some())
699 .finish_non_exhaustive()
700 }
701}
702
703struct ObservabilityClarificationObserver;
704
705impl ClarificationObserver for ObservabilityClarificationObserver {
706 fn observe_question<'a>(
708 &'a self,
709 future: ClarificationQuestionFuture<'a>,
710 ) -> ClarificationQuestionFuture<'a> {
711 Box::pin(async move {
712 with_observation_purpose(ObservationPurpose::DisambiguationClarification, future).await
713 })
714 }
715
716 fn observe_parse<'a>(
718 &'a self,
719 future: ClarificationParseFuture<'a>,
720 ) -> ClarificationParseFuture<'a> {
721 Box::pin(async move {
722 with_observation_purpose(ObservationPurpose::DisambiguationClarification, future).await
723 })
724 }
725
726 fn observe_confirmation_parse<'a>(
728 &'a self,
729 future: ConfirmationParseFuture<'a>,
730 ) -> ConfirmationParseFuture<'a> {
731 Box::pin(async move {
732 with_observation_purpose(ObservationPurpose::DisambiguationClarification, future).await
733 })
734 }
735}
736
737struct ObservabilityProcessStageObserver;
738
739impl ProcessStageObserver for ObservabilityProcessStageObserver {
740 fn observe<'a>(
742 &'a self,
743 hint: ProcessPurposeHint,
744 future: ProcessStageFuture<'a>,
745 ) -> ProcessStageFuture<'a> {
746 Box::pin(async move {
747 with_observation_purpose(observation_purpose_for_process(hint), future).await
748 })
749 }
750}
751
752struct RegistryLLMGetter {
753 registry: Arc<LLMRegistry>,
754}
755
756impl LLMGetter for RegistryLLMGetter {
757 fn get_llm(&self, alias: &str) -> Option<Arc<dyn LLMProvider>> {
758 self.registry.get(alias).ok()
759 }
760}
761
762impl RuntimeAgent {
763 #[allow(clippy::too_many_arguments)]
765 pub fn new(
766 info: AgentInfo,
767 llm_registry: Arc<LLMRegistry>,
768 memory: Arc<dyn Memory>,
769 tools: Arc<ToolRegistry>,
770 skills: Vec<SkillDefinition>,
771 system_prompt: String,
772 max_iterations: u32,
773 ) -> Self {
774 let (skill_router, skill_executor) = if !skills.is_empty() {
775 let router_llm = llm_registry.router().ok();
776 let router = router_llm.map(|llm| SkillRouter::new(llm, skills.clone()));
777 let executor = SkillExecutor::new(llm_registry.clone(), tools.clone());
778 (router, Some(executor))
779 } else {
780 (None, None)
781 };
782
783 let context_manager =
784 ContextManager::new(HashMap::new(), info.name.clone(), info.version.clone());
785
786 Self {
787 info,
788 llm_registry,
789 memory,
790 tools,
791 skills,
792 skill_router,
793 skill_executor,
794 base_system_prompt: system_prompt,
795 max_iterations,
796 iteration_count: RwLock::new(0),
797 max_context_tokens: 128000,
798 memory_token_budget: None,
799 recovery_manager: RecoveryManager::default(),
800 tool_security: ToolSecurityEngine::default(),
801 process_processor: None,
802 message_filters: RwLock::new(HashMap::new()),
803 state_machine: None,
804 transition_evaluator: None,
805 context_manager: Arc::new(context_manager),
806 template_renderer: TemplateRenderer::new(),
807 tool_call_history: RwLock::new(Vec::new()),
808 parallel_tools: ParallelToolsConfig::default(),
809 streaming: StreamingConfig::default(),
810 hooks: Arc::new(NoopHooks),
811 hitl_engine: None,
812 approval_handler: Arc::new(RejectAllHandler::new()),
813 storage_config: StorageConfig::default(),
814 storage: RwLock::new(None),
815 storage_init: tokio::sync::Mutex::new(()),
816 reasoning_config: ReasoningConfig::default(),
817 reflection_config: ReflectionConfig::default(),
818 disambiguation_manager: None,
819 disambiguation_epoch: AtomicU64::new(0),
820 disambiguation_admission: tokio::sync::RwLock::new(()),
821 state_transition_reserved: AtomicBool::new(false),
822 persona_manager: None,
823 pending_skill_id: RwLock::new(None),
824 current_plan: RwLock::new(None),
825 declared_tool_ids: None,
826 context_initialized: AtomicBool::new(false),
827 spawner: None,
828 spawner_registry: None,
829 redispatch_depth: RwLock::new(0),
830 active_turn_context: RwLock::new(None),
831 root_user_message_committed: AtomicBool::new(false),
832 active_native_exchanges: RwLock::new(Vec::new()),
833 actor_id: RwLock::new(None),
834 fact_store: RwLock::new(None),
835 fact_extractor: RwLock::new(None),
836 actor_facts_cache: Arc::new(RwLock::new(HashMap::new())),
837 messages_since_extraction: Arc::new(RwLock::new(0)),
838 actor_memory_config: None,
839 facts_config: None,
840 session_metadata: RwLock::new(ai_agents_core::SessionMetadata::default()),
841 current_session_id: RwLock::new(None),
842 relationship_manager: None,
843 observability_manager: None,
844 runtime_config: RuntimeConfig::default(),
845 background_maintenance: Arc::new(BackgroundMaintenanceQueue::default()),
846 resource_locks: new_tool_resource_locks(),
847 runtime_control: Arc::new(RuntimeControlState::default()),
848 root_turn_gate: Arc::new(tokio::sync::Mutex::new(())),
849 }
850 }
851
852 pub fn with_declared_tool_ids(mut self, ids: Option<Vec<String>>) -> Self {
853 self.declared_tool_ids = ids;
854 self
855 }
856
857 pub fn with_storage_config(mut self, config: StorageConfig) -> Self {
858 self.storage_config = config;
859 self
860 }
861
862 pub fn with_storage(self, storage: Arc<dyn AgentStorage>) -> Self {
863 *self.storage.write() = Some(storage);
864 self
865 }
866
867 pub(crate) fn with_shared_resource_locks(mut self, locks: ToolResourceLocks) -> Self {
868 self.resource_locks = locks;
869 self
870 }
871
872 pub fn with_reasoning(mut self, config: ReasoningConfig) -> Self {
873 self.reasoning_config = config;
874 self
875 }
876
877 pub fn with_reflection(mut self, config: ReflectionConfig) -> Self {
878 self.reflection_config = config;
879 self
880 }
881
882 pub fn with_relationships(mut self, manager: Arc<RelationshipManager>) -> Self {
884 self.relationship_manager = Some(manager);
885 self
886 }
887
888 pub fn with_observability(mut self, manager: Arc<ObservabilityManager>) -> Self {
890 self.observability_manager = Some(manager);
891 self
892 }
893
894 pub fn with_runtime_config(mut self, config: RuntimeConfig) -> Self {
896 let max_tasks = config.optimization.post_turn.max_background_tasks;
897 self.background_maintenance = Arc::new(BackgroundMaintenanceQueue::new(max_tasks));
898 self.runtime_config = config;
899 self
900 }
901
902 pub fn runtime_config(&self) -> &RuntimeConfig {
904 &self.runtime_config
905 }
906
907 pub async fn flush_background_tasks(&self) -> Result<()> {
909 self.background_maintenance.flush_all().await
910 }
911
912 pub async fn flush_background_tasks_for_actor(&self, actor_id: &str) -> Result<()> {
914 self.background_maintenance.flush_scope(actor_id).await
915 }
916
917 pub async fn flush_background_tasks_for_purpose(
919 &self,
920 purpose: RuntimeTaskPurpose,
921 ) -> Result<()> {
922 self.background_maintenance.flush_purpose(purpose).await
923 }
924
925 pub async fn flush_background_tasks_for_actor_purpose(
927 &self,
928 actor_id: &str,
929 purpose: RuntimeTaskPurpose,
930 ) -> Result<()> {
931 self.background_maintenance
932 .flush_scope_purpose(actor_id, purpose)
933 .await
934 }
935
936 pub async fn shutdown_background_tasks(&self) -> Result<()> {
938 self.flush_background_tasks().await
939 }
940
941 pub fn observability(&self) -> Option<Arc<ObservabilityManager>> {
943 self.observability_manager.clone()
944 }
945
946 async fn export_observability_if_configured(&self) {
948 let Some(manager) = self.observability_manager.as_ref() else {
949 return;
950 };
951 let export = &manager.config().export;
952 if !export.write_report && !export.write_raw_events {
953 return;
954 }
955 if let Err(error) = manager.export().await {
956 warn!(error = %error, "Observability export failed");
957 }
958 }
959
960 pub fn relationship_manager(&self) -> Option<Arc<RelationshipManager>> {
962 self.relationship_manager.clone()
963 }
964
965 fn current_turn_actor_context(&self) -> Option<crate::TurnActorContext> {
966 current_turn_actor_context()
967 }
968
969 fn effective_actor_id(&self) -> Option<String> {
970 self.current_turn_actor_context()
971 .and_then(|ctx| ctx.effective_actor_id().map(|id| id.to_string()))
972 .or_else(|| self.actor_id.read().clone())
973 }
974
975 fn effective_origin_actor_id(&self) -> Option<String> {
976 self.current_turn_actor_context()
977 .and_then(|ctx| ctx.origin_actor_id.clone())
978 .or_else(|| self.actor_id.read().clone())
979 }
980
981 fn record_session_actor_if_needed(&self) {
982 if let Some(actor_id) = self.effective_origin_actor_id() {
983 let mut meta = self.session_metadata.write();
984 meta.actor_id = Some(actor_id.clone());
985 if !meta.actors.iter().any(|a| a == &actor_id) {
986 meta.actors.push(actor_id);
987 }
988 }
989 }
990
991 fn outbound_actor_context(&self) -> crate::TurnActorContext {
992 let mut context = self.current_turn_actor_context().unwrap_or_default();
993 if context.origin_actor_id.is_none() {
994 context.origin_actor_id = self.effective_origin_actor_id();
995 }
996 context.sender_agent_id = Some(self.info.id.clone());
997 context
998 }
999
1000 fn observation_session_id(&self) -> Option<String> {
1002 let mut current = self.current_session_id.write();
1003 if current.is_none() {
1004 *current = Some(new_observation_session_id());
1005 }
1006 current.clone()
1007 }
1008
1009 fn build_observation_context(&self, actor_id: Option<String>) -> Option<SpanContext> {
1011 let manager = self.observability_manager.as_ref()?;
1012 let context = self.build_context_with_overlays();
1013 let language = resolve_language_from_context(manager.config(), &context);
1014 let context = current_observation_context()
1015 .map(|parent| parent.child_for_agent(self.info.id.clone()).with_new_turn())
1016 .unwrap_or_else(|| SpanContext::new_root(self.info.id.clone()));
1017 Some(
1018 context
1019 .with_actor(actor_id.or_else(|| self.effective_actor_id()))
1020 .with_session(self.observation_session_id())
1021 .with_state(self.current_state())
1022 .with_language(Some(language)),
1023 )
1024 }
1025
1026 fn current_runtime_observation_context(
1028 &self,
1029 purpose: ObservationPurpose,
1030 ) -> Option<SpanContext> {
1031 let manager = self.observability_manager.as_ref()?;
1032 let context = self.build_context_with_overlays();
1033 let language = resolve_language_from_context(manager.config(), &context);
1034 let mut observation = current_observation_context()
1035 .unwrap_or_else(|| SpanContext::new_root(self.info.id.clone()));
1036 observation.agent_id = self.info.id.clone();
1037 observation.actor_id = self.effective_actor_id();
1038 observation.session_id = self.observation_session_id();
1039 observation.state = self.current_state();
1040 observation.language = Some(language);
1041 observation.purpose = purpose;
1042 Some(observation)
1043 }
1044
1045 async fn observe_purpose<F, T>(&self, purpose: ObservationPurpose, future: F) -> T
1047 where
1048 F: Future<Output = T>,
1049 {
1050 if let Some(context) = self.current_runtime_observation_context(purpose) {
1051 with_observation_context(context, future).await
1052 } else {
1053 future.await
1054 }
1055 }
1056
1057 fn chat_with_actor_context_boxed<'a>(
1061 &'a self,
1062 input: &'a str,
1063 actor_context: crate::TurnActorContext,
1064 ) -> Pin<Box<dyn Future<Output = Result<AgentResponse>> + Send + 'a>> {
1065 Box::pin(async move {
1066 let RootTurnAdmission {
1067 guard,
1068 identity_stack,
1069 } = self.acquire_root_turn().await?;
1070 let result = scope_runtime_gate_identity_stack(&identity_stack, async move {
1071 let actor_id = actor_context.effective_actor_id().map(str::to_string);
1072 let run = async move {
1073 scope_actor_context(
1074 actor_context,
1075 Box::pin(async move { self.run_loop(input).await }),
1076 )
1077 .await
1078 };
1079 let result = if let Some(context) = self.build_observation_context(actor_id) {
1080 with_observation_context(context, run).await
1081 } else {
1082 run.await
1083 };
1084 self.export_observability_if_configured().await;
1085 result
1086 })
1087 .await;
1088 drop(guard);
1089 result
1090 })
1091 }
1092
1093 async fn acquire_root_turn(&self) -> Result<RootTurnAdmission> {
1095 let gate_identity = Arc::clone(&self.root_turn_gate);
1096 let current_identity_stack = current_runtime_gate_identity_stack();
1097 if current_identity_stack
1098 .iter()
1099 .any(|owned_gate| Arc::ptr_eq(owned_gate, &gate_identity))
1100 {
1101 return Err(AgentError::Other(format!(
1102 "RuntimeAgent '{}' rejected reentrant root turn ownership",
1103 self.info.id
1104 )));
1105 }
1106 let guard = Arc::clone(&gate_identity).lock_owned().await;
1107 let mut identity_stack = Vec::with_capacity(current_identity_stack.len() + 1);
1111 identity_stack.extend(current_identity_stack.iter().cloned());
1112 identity_stack.push(gate_identity);
1113 Ok(RootTurnAdmission {
1114 guard,
1115 identity_stack: identity_stack.into(),
1116 })
1117 }
1118
1119 pub async fn chat_with_actor_context(
1123 &self,
1124 input: &str,
1125 actor_context: crate::TurnActorContext,
1126 ) -> Result<AgentResponse> {
1127 self.chat_with_actor_context_boxed(input, actor_context)
1128 .await
1129 }
1130
1131 pub async fn chat_as_actor(&self, actor_id: &str, input: &str) -> Result<AgentResponse> {
1133 let actor_context = crate::TurnActorContext::new().with_origin_actor(actor_id);
1134 self.chat_with_actor_context(input, actor_context).await
1135 }
1136
1137 pub async fn load_actor_relationship(&self) -> Result<()> {
1139 self.maybe_load_actor_relationship().await;
1140 Ok(())
1141 }
1142
1143 pub async fn update_relationship_dimension(
1145 &self,
1146 dimension: &str,
1147 delta: f64,
1148 reason: Option<&str>,
1149 ) -> Result<ai_agents_relationships::DimensionChange> {
1150 self.update_relationship_dimension_for_perspective(
1151 ai_agents_relationships::RelationshipPerspective::AgentToActor,
1152 dimension,
1153 delta,
1154 reason,
1155 )
1156 .await
1157 }
1158
1159 pub async fn update_relationship_dimension_for_perspective(
1163 &self,
1164 perspective: ai_agents_relationships::RelationshipPerspective,
1165 dimension: &str,
1166 delta: f64,
1167 reason: Option<&str>,
1168 ) -> Result<ai_agents_relationships::DimensionChange> {
1169 let manager = self
1170 .relationship_manager
1171 .as_ref()
1172 .ok_or_else(|| AgentError::Config("Relationship memory is not configured".into()))?;
1173 let actor_id = self.effective_actor_id().ok_or_else(|| {
1174 AgentError::Config("No actor ID set. Use set_actor_id() first".into())
1175 })?;
1176 let change = manager.update_dimension_for_perspective(
1177 &actor_id,
1178 perspective,
1179 dimension,
1180 delta,
1181 1.0,
1182 reason.unwrap_or("manual relationship update"),
1183 )?;
1184 self.persist_actor_relationship(&actor_id).await?;
1185 info!(
1186 actor_id = %actor_id,
1187 perspective = %change.perspective,
1188 dimension = %change.dimension,
1189 delta = change.delta,
1190 current = change.current,
1191 "relationship updated manually"
1192 );
1193 self.hooks
1194 .on_relationship_change(&actor_id, std::slice::from_ref(&change))
1195 .await;
1196 Ok(change)
1197 }
1198
1199 pub fn reasoning_config(&self) -> &ReasoningConfig {
1200 &self.reasoning_config
1201 }
1202
1203 pub fn reflection_config(&self) -> &ReflectionConfig {
1204 &self.reflection_config
1205 }
1206
1207 pub fn with_facts_config(
1210 mut self,
1211 actor_memory_config: Option<ai_agents_facts::ActorMemoryConfig>,
1212 facts_config: Option<ai_agents_facts::FactsConfig>,
1213 ) -> Self {
1214 self.actor_memory_config = actor_memory_config;
1215 self.facts_config = facts_config;
1216 self
1217 }
1218
1219 pub fn with_facts(
1222 mut self,
1223 store: Arc<ai_agents_facts::FactStore>,
1224 extractor: Option<Arc<dyn ai_agents_facts::FactExtractor>>,
1225 actor_memory_config: Option<ai_agents_facts::ActorMemoryConfig>,
1226 facts_config: Option<ai_agents_facts::FactsConfig>,
1227 ) -> Self {
1228 *self.fact_store.write() = Some(store);
1229 *self.fact_extractor.write() = extractor;
1230 self.actor_memory_config = actor_memory_config;
1231 self.facts_config = facts_config;
1232 self
1233 }
1234
1235 pub fn fact_store(&self) -> Option<Arc<ai_agents_facts::FactStore>> {
1237 self.fact_store.read().clone()
1238 }
1239
1240 pub fn actor_id(&self) -> Option<String> {
1242 self.actor_id.read().clone()
1243 }
1244
1245 pub fn set_actor_id(&self, actor_id: &str) -> ai_agents_core::Result<()> {
1247 *self.actor_id.write() = Some(actor_id.to_string());
1248 {
1249 let mut meta = self.session_metadata.write();
1250 meta.actor_id = Some(actor_id.to_string());
1251 if !meta.actors.iter().any(|a| a == actor_id) {
1252 meta.actors.push(actor_id.to_string());
1253 }
1254 }
1255 Ok(())
1256 }
1257
1258 pub fn clear_actor_id(&self) {
1260 *self.actor_id.write() = None;
1261 self.session_metadata.write().actor_id = None;
1262 }
1263
1264 pub fn set_user_id(&self, user_id: &str) -> ai_agents_core::Result<()> {
1266 self.set_actor_id(user_id)
1267 }
1268
1269 pub async fn load_actor_memory(&self) -> ai_agents_core::Result<()> {
1271 let actor_id = match self.effective_actor_id() {
1272 Some(id) => id,
1273 None => return Ok(()),
1274 };
1275
1276 let store_opt = self.fact_store.read().clone();
1277 if let Some(store) = store_opt {
1278 let facts = store.get_facts(&actor_id).await?;
1279 let count = facts.len();
1280 self.actor_facts_cache
1281 .write()
1282 .insert(actor_id.clone(), facts);
1283 self.hooks.on_actor_memory_loaded(&actor_id, count).await;
1284 tracing::debug!("loaded {} facts for actor {}", count, actor_id);
1285 }
1286
1287 Ok(())
1288 }
1289
1290 async fn maybe_load_actor_memory(&self) {
1292 let Some(actor_id) = self.effective_actor_id() else {
1293 return;
1294 };
1295 if self.actor_facts_cache.read().contains_key(&actor_id) {
1296 return;
1297 }
1298 let _ = self.load_actor_memory().await;
1299 }
1300
1301 async fn pre_turn_session_lifecycle(&self) {
1303 if *self.redispatch_depth.read() > 0 {
1304 return;
1305 }
1306 self.resolve_actor_id_from_context();
1307 self.await_background_before_next_turn().await;
1308 self.record_session_actor_if_needed();
1309 self.maybe_load_actor_memory().await;
1310 self.maybe_load_actor_relationship().await;
1311 *self.messages_since_extraction.write() += 1;
1312 }
1313
1314 async fn post_turn_session_lifecycle(&self) -> Result<()> {
1316 if *self.redispatch_depth.read() > 0 {
1317 return Ok(());
1318 }
1319 *self.messages_since_extraction.write() += 1;
1320 self.run_post_turn_maintenance().await
1321 }
1322
1323 fn begin_root_turn(&self) {
1325 if *self.redispatch_depth.read() == 0 {
1326 let mut guard = self.active_turn_context.write();
1327 if guard.is_none() {
1328 self.root_user_message_committed
1329 .store(false, Ordering::SeqCst);
1330 self.active_native_exchanges.write().clear();
1331 let max_calls = self
1332 .runtime_config
1333 .optimization
1334 .max_speculative_llm_calls_per_turn;
1335 *guard = Some(TurnOptimizationContext::new(
1336 String::new(),
1337 HashMap::new(),
1338 max_calls,
1339 ));
1340 }
1341 }
1342 }
1343
1344 fn update_active_turn_context(
1345 &self,
1346 processed_input: &str,
1347 input_context: HashMap<String, Value>,
1348 ) {
1349 if *self.redispatch_depth.read() > 0 {
1350 return;
1351 }
1352 let max_calls = self
1353 .runtime_config
1354 .optimization
1355 .max_speculative_llm_calls_per_turn;
1356 let mut guard = self.active_turn_context.write();
1357 match guard.as_mut() {
1358 Some(context) => {
1359 context.processed_input = processed_input.to_string();
1360 context.input_context = input_context;
1361 context.max_speculative_llm_calls = max_calls;
1362 }
1363 None => {
1364 *guard = Some(TurnOptimizationContext::new(
1365 processed_input,
1366 input_context,
1367 max_calls,
1368 ));
1369 }
1370 }
1371 }
1372
1373 async fn commit_root_user_message(&self, processed_input: &str) -> Result<()> {
1375 if *self.redispatch_depth.read() > 0 {
1376 return Ok(());
1377 }
1378 if !self
1379 .root_user_message_committed
1380 .swap(true, Ordering::SeqCst)
1381 {
1382 self.memory
1383 .add_message(ChatMessage::user(processed_input))
1384 .await?;
1385 if let Some(context) = self.active_turn_context.write().as_mut() {
1386 context.mark_user_message_committed();
1387 }
1388 }
1389 Ok(())
1390 }
1391
1392 fn end_root_turn(&self) {
1394 if *self.redispatch_depth.read() == 0 {
1395 self.root_user_message_committed
1396 .store(false, Ordering::SeqCst);
1397 *self.active_turn_context.write() = None;
1398 self.active_native_exchanges.write().clear();
1399 }
1400 }
1401
1402 fn reserve_active_speculative_llm_call(&self, kind: RuntimeOptimizationKind) -> bool {
1403 self.begin_root_turn();
1404 let mut guard = self.active_turn_context.write();
1405 let Some(context) = guard.as_mut() else {
1406 return false;
1407 };
1408 context.reserve_speculative_llm_call_for(kind)
1409 }
1410
1411 fn branch_context_preview(&self) -> String {
1412 let context = self.build_context_with_overlays();
1413 let mut value = serde_json::to_string_pretty(&context).unwrap_or_else(|_| "{}".to_string());
1414 const MAX_CONTEXT_PREVIEW_CHARS: usize = 2048;
1415 if value.chars().count() > MAX_CONTEXT_PREVIEW_CHARS {
1416 value = value
1417 .chars()
1418 .take(MAX_CONTEXT_PREVIEW_CHARS)
1419 .collect::<String>();
1420 value.push_str("...");
1421 }
1422 value
1423 }
1424
1425 async fn await_background_before_next_turn(&self) {
1427 let optimization = &self.runtime_config.optimization;
1428 if !optimization.enabled {
1429 return;
1430 }
1431 let actor_id = self.effective_actor_id();
1432 let post = &optimization.post_turn;
1433 self.await_background_task(
1434 post.facts.await_before_next_turn,
1435 RuntimeTaskPurpose::PostTurnFacts,
1436 actor_id.as_deref(),
1437 "facts",
1438 )
1439 .await;
1440 self.await_background_task(
1441 post.relationships.await_before_next_turn,
1442 RuntimeTaskPurpose::PostTurnRelationship,
1443 actor_id.as_deref(),
1444 "relationships",
1445 )
1446 .await;
1447 }
1448
1449 async fn await_background_task(
1450 &self,
1451 policy: AwaitBeforeNextTurn,
1452 purpose: RuntimeTaskPurpose,
1453 actor_id: Option<&str>,
1454 label: &str,
1455 ) {
1456 match policy {
1457 AwaitBeforeNextTurn::Never => {}
1458 AwaitBeforeNextTurn::Always => {
1459 if let Err(error) = self.flush_background_tasks_for_purpose(purpose).await {
1460 warn!(label = label, error = %error, "background maintenance flush failed");
1461 }
1462 }
1463 AwaitBeforeNextTurn::SameActor => {
1464 if let Some(actor_id) = actor_id
1465 && let Err(error) = self
1466 .flush_background_tasks_for_actor_purpose(actor_id, purpose)
1467 .await
1468 {
1469 warn!(label = label, actor_id = %actor_id, error = %error, "actor background maintenance flush failed");
1470 }
1471 }
1472 }
1473 }
1474
1475 async fn run_post_turn_maintenance(&self) -> Result<()> {
1477 let optimization = &self.runtime_config.optimization;
1478 if !optimization.enabled {
1479 self.auto_extract_facts().await;
1480 self.auto_update_relationship().await;
1481 return Ok(());
1482 }
1483
1484 let facts_mode = effective_maintenance_mode(
1485 optimization.post_turn.facts.mode,
1486 optimization.parallel_post_turn_memory,
1487 );
1488 let relationships_mode = effective_maintenance_mode(
1489 optimization.post_turn.relationships.mode,
1490 optimization.parallel_post_turn_memory,
1491 );
1492
1493 match (facts_mode, relationships_mode) {
1494 (MaintenanceMode::InlineSerial, MaintenanceMode::InlineSerial) => {
1495 self.auto_extract_facts().await;
1496 self.auto_update_relationship().await;
1497 }
1498 (MaintenanceMode::InlineParallel, MaintenanceMode::InlineParallel) => {
1499 let facts = self.auto_extract_facts();
1500 let relationships = self.auto_update_relationship();
1501 tokio::join!(facts, relationships);
1502 }
1503 (MaintenanceMode::Background, MaintenanceMode::Background) => {
1504 self.schedule_facts_background().await?;
1505 self.schedule_relationship_background().await?;
1506 }
1507 (MaintenanceMode::Background, MaintenanceMode::InlineParallel)
1508 | (MaintenanceMode::Background, MaintenanceMode::InlineSerial) => {
1509 self.schedule_facts_background().await?;
1510 self.auto_update_relationship().await;
1511 }
1512 (MaintenanceMode::InlineParallel, MaintenanceMode::Background)
1513 | (MaintenanceMode::InlineSerial, MaintenanceMode::Background) => {
1514 self.auto_extract_facts().await;
1515 self.schedule_relationship_background().await?;
1516 }
1517 _ => {
1518 self.auto_extract_facts().await;
1519 self.auto_update_relationship().await;
1520 }
1521 }
1522 Ok(())
1523 }
1524
1525 async fn schedule_facts_background(&self) -> Result<()> {
1526 let policy = self.runtime_config.optimization.post_turn.facts.clone();
1527 let should_extract = self
1528 .facts_config
1529 .as_ref()
1530 .map(|c| c.enabled && c.auto_extract)
1531 .unwrap_or(false);
1532 if !should_extract {
1533 return Ok(());
1534 }
1535 let msgs_since = *self.messages_since_extraction.read();
1536 if msgs_since < 2 {
1537 return Ok(());
1538 }
1539 let Some(actor_id) = self.effective_actor_id() else {
1540 self.record_skipped_maintenance(
1541 "facts",
1542 ObservationPurpose::FactsExtraction,
1543 "missing_actor",
1544 Some(&policy),
1545 );
1546 return Ok(());
1547 };
1548 let Some(extractor) = self.fact_extractor.read().clone() else {
1549 return Ok(());
1550 };
1551 let messages = match self.memory.get_messages(None).await {
1552 Ok(messages) => messages,
1553 Err(error) => {
1554 warn!(error = %error, "failed to snapshot messages for fact extraction");
1555 return Ok(());
1556 }
1557 };
1558 let messages = Self::readable_native_messages(messages)?;
1559 let recent: Vec<_> = messages
1560 .iter()
1561 .rev()
1562 .take(msgs_since)
1563 .rev()
1564 .cloned()
1565 .collect();
1566 if recent.is_empty() {
1567 return Ok(());
1568 }
1569 let existing = self
1570 .actor_facts_cache
1571 .read()
1572 .get(&actor_id)
1573 .cloned()
1574 .unwrap_or_default();
1575 let categories = self
1576 .facts_config
1577 .as_ref()
1578 .map(|c| c.custom_categories.clone())
1579 .unwrap_or_default();
1580 let store = self.fact_store.read().clone();
1581 let cache = Arc::clone(&self.actor_facts_cache);
1582 let counter = Arc::clone(&self.messages_since_extraction);
1583 let hooks = Arc::clone(&self.hooks);
1584 let agent_id = self.info.id.clone();
1585 let observation = current_observation_context();
1586 let key = MaintenanceSequenceKey::actor(
1587 agent_id,
1588 actor_id.clone(),
1589 RuntimeTaskPurpose::PostTurnFacts,
1590 );
1591 let actor_for_task = actor_id.clone();
1592 let task = async move {
1593 let run = async move {
1594 let facts = extractor
1595 .extract(&recent, &existing, Some(&actor_for_task), &categories)
1596 .await?;
1597 if !facts.is_empty() {
1598 if let Some(store) = store {
1599 let authoritative = store.add_facts(&actor_for_task, facts.clone()).await?;
1600 cache.write().insert(actor_for_task.clone(), authoritative);
1601 } else {
1602 cache
1603 .write()
1604 .entry(actor_for_task.clone())
1605 .or_default()
1606 .extend(facts.clone());
1607 }
1608 {
1609 let mut count = counter.write();
1610 if *count <= msgs_since {
1611 *count = 0;
1612 } else {
1613 *count -= msgs_since;
1614 }
1615 }
1616 hooks.on_facts_extracted(&actor_for_task, &facts).await;
1617 }
1618 Ok(())
1619 };
1620 if let Some(context) = observation {
1621 with_observation_context(
1622 context.with_purpose(ObservationPurpose::FactsExtraction),
1623 run,
1624 )
1625 .await
1626 } else {
1627 run.await
1628 }
1629 };
1630 self.spawn_or_handle_background(Some(key), task, "facts", &policy)
1631 .await
1632 }
1633
1634 async fn schedule_relationship_background(&self) -> Result<()> {
1635 let policy = self
1636 .runtime_config
1637 .optimization
1638 .post_turn
1639 .relationships
1640 .clone();
1641 let Some(manager) = self.relationship_manager.as_ref().cloned() else {
1642 return Ok(());
1643 };
1644 let Some(actor_id) = self.effective_actor_id() else {
1645 self.record_skipped_maintenance(
1646 "relationships",
1647 ObservationPurpose::RelationshipUpdate,
1648 "missing_actor",
1649 Some(&policy),
1650 );
1651 return Ok(());
1652 };
1653 let recent_messages = manager.config().auto_update.recent_messages;
1654 let messages = match self.memory.get_messages(Some(recent_messages)).await {
1655 Ok(messages) => messages,
1656 Err(error) => {
1657 warn!(actor = %actor_id, error = %error, "failed to snapshot messages for relationship update");
1658 return Ok(());
1659 }
1660 };
1661 let messages = Self::readable_native_messages(messages)?;
1662 let storage = self.storage.read().clone();
1663 let hooks = Arc::clone(&self.hooks);
1664 let agent_id = self.info.id.clone();
1665 let observation = current_observation_context();
1666 let key = MaintenanceSequenceKey::actor(
1667 agent_id.clone(),
1668 actor_id.clone(),
1669 RuntimeTaskPurpose::PostTurnRelationship,
1670 );
1671 let actor_for_task = actor_id.clone();
1672 let task = async move {
1673 let run = async move {
1674 if manager.config().auto_update.enabled {
1675 let update = manager.auto_update(&actor_for_task, &messages).await?;
1676 if !update.changes.is_empty() {
1677 hooks
1678 .on_relationship_change(&actor_for_task, &update.changes)
1679 .await;
1680 }
1681 if let Some(ref event) = update.event {
1682 hooks.on_notable_event(&actor_for_task, event).await;
1683 }
1684 }
1685 if manager.config().persistence.enabled
1686 && let (Some(storage), Some(value)) =
1687 (storage, manager.relationship_as_value(&actor_for_task)?)
1688 {
1689 storage
1690 .save_relationship(&agent_id, &actor_for_task, &value)
1691 .await?;
1692 }
1693 Ok(())
1694 };
1695 if let Some(context) = observation {
1696 with_observation_context(
1697 context.with_purpose(ObservationPurpose::RelationshipUpdate),
1698 run,
1699 )
1700 .await
1701 } else {
1702 run.await
1703 }
1704 };
1705 self.spawn_or_handle_background(Some(key), task, "relationships", &policy)
1706 .await
1707 }
1708
1709 async fn spawn_or_handle_background<F>(
1711 &self,
1712 key: Option<MaintenanceSequenceKey>,
1713 task: F,
1714 label: &'static str,
1715 policy: &crate::optimization::config::MaintenanceTaskPolicy,
1716 ) -> Result<()>
1717 where
1718 F: Future<Output = Result<()>> + Send + 'static,
1719 {
1720 if self.background_maintenance.is_full() {
1721 match self
1722 .runtime_config
1723 .optimization
1724 .post_turn
1725 .on_background_overflow
1726 {
1727 BackgroundOverflowPolicy::RunInline => {
1728 record_background_maintenance_event(
1729 self.observability_manager.as_ref(),
1730 label,
1731 EventStatus::Success,
1732 0,
1733 "inline_overflow",
1734 None,
1735 Some(policy),
1736 );
1737 let start = Instant::now();
1738 match task.await {
1739 Ok(()) => record_background_maintenance_event(
1740 self.observability_manager.as_ref(),
1741 label,
1742 EventStatus::Success,
1743 start.elapsed().as_millis() as u64,
1744 "inline_completed",
1745 None,
1746 Some(policy),
1747 ),
1748 Err(error) => {
1749 warn!(label = label, error = %error, "inline maintenance fallback failed");
1750 record_background_maintenance_event(
1751 self.observability_manager.as_ref(),
1752 label,
1753 EventStatus::Error,
1754 start.elapsed().as_millis() as u64,
1755 "inline_failed",
1756 Some(error.to_string()),
1757 Some(policy),
1758 );
1759 return Err(error);
1760 }
1761 }
1762 }
1763 BackgroundOverflowPolicy::Drop => {
1764 self.record_skipped_maintenance(
1765 label,
1766 ObservationPurpose::Other(label.to_string()),
1767 "queue_full",
1768 Some(policy),
1769 );
1770 }
1771 BackgroundOverflowPolicy::Error => {
1772 record_background_maintenance_event(
1773 self.observability_manager.as_ref(),
1774 label,
1775 EventStatus::Error,
1776 0,
1777 "queue_full",
1778 None,
1779 Some(policy),
1780 );
1781 warn!(label = label, "background maintenance queue full");
1782 return Err(AgentError::Other(format!(
1783 "background maintenance queue is full for {}",
1784 label
1785 )));
1786 }
1787 }
1788 return Ok(());
1789 }
1790
1791 record_background_maintenance_event(
1792 self.observability_manager.as_ref(),
1793 label,
1794 EventStatus::Success,
1795 0,
1796 "scheduled",
1797 None,
1798 Some(policy),
1799 );
1800 let manager = self.observability_manager.clone();
1801 let policy_for_task = policy.clone();
1802 let observed_task = async move {
1803 let start = Instant::now();
1804 let result = task.await;
1805 match &result {
1806 Ok(()) => record_background_maintenance_event(
1807 manager.as_ref(),
1808 label,
1809 EventStatus::Success,
1810 start.elapsed().as_millis() as u64,
1811 "completed",
1812 None,
1813 Some(&policy_for_task),
1814 ),
1815 Err(error) => record_background_maintenance_event(
1816 manager.as_ref(),
1817 label,
1818 EventStatus::Error,
1819 start.elapsed().as_millis() as u64,
1820 "failed",
1821 Some(error.to_string()),
1822 Some(&policy_for_task),
1823 ),
1824 }
1825 result
1826 };
1827
1828 if let Err(error) = self.background_maintenance.spawn(key, observed_task) {
1829 record_background_maintenance_event(
1830 self.observability_manager.as_ref(),
1831 label,
1832 EventStatus::Error,
1833 0,
1834 "spawn_failed",
1835 Some(error.to_string()),
1836 Some(policy),
1837 );
1838 warn!(label = label, error = %error, "background maintenance spawn failed");
1839 return Err(error);
1840 }
1841 Ok(())
1842 }
1843
1844 fn record_skipped_maintenance(
1846 &self,
1847 label: &str,
1848 purpose: ObservationPurpose,
1849 reason: &str,
1850 policy: Option<&crate::optimization::config::MaintenanceTaskPolicy>,
1851 ) {
1852 if let Some(manager) = self.observability_manager.as_ref() {
1853 let mut tags = background_maintenance_tags(label, "skipped", Some(reason), policy);
1854 tags.insert("runtime.skip_reason".to_string(), reason.to_string());
1855 manager.record_lifecycle_event(
1856 EventType::MemoryOperation {
1857 operation: format!("{}_maintenance", label),
1858 },
1859 purpose,
1860 EventStatus::Skipped,
1861 0,
1862 tags,
1863 None,
1864 );
1865 }
1866 }
1867
1868 pub fn actor_facts(&self) -> Vec<ai_agents_core::KeyFact> {
1870 let Some(actor_id) = self.effective_actor_id() else {
1871 return Vec::new();
1872 };
1873 self.actor_facts_cache
1874 .read()
1875 .get(&actor_id)
1876 .cloned()
1877 .unwrap_or_default()
1878 }
1879
1880 pub fn relationship_memory_text(&self) -> Option<String> {
1882 self.format_relationship_for_context().map(|(_, text)| text)
1883 }
1884
1885 pub async fn extract_facts(
1887 &self,
1888 last_n: usize,
1889 ) -> ai_agents_core::Result<Vec<ai_agents_core::KeyFact>> {
1890 self.extract_facts_with_source(last_n, "manual").await
1891 }
1892
1893 async fn extract_facts_with_source(
1894 &self,
1895 last_n: usize,
1896 source: &'static str,
1897 ) -> ai_agents_core::Result<Vec<ai_agents_core::KeyFact>> {
1898 let extractor = match self.fact_extractor.read().clone() {
1899 Some(e) => e,
1900 None => return Ok(vec![]),
1901 };
1902
1903 let messages = Self::readable_native_messages(self.memory.get_messages(None).await?)?;
1904 let recent: Vec<_> = messages.iter().rev().take(last_n).rev().cloned().collect();
1905
1906 if recent.is_empty() {
1907 return Ok(vec![]);
1908 }
1909
1910 let actor_id = self.effective_actor_id();
1911 let existing = actor_id
1912 .as_ref()
1913 .and_then(|aid| self.actor_facts_cache.read().get(aid).cloned())
1914 .unwrap_or_default();
1915
1916 let categories = self
1917 .facts_config
1918 .as_ref()
1919 .map(|c| c.custom_categories.clone())
1920 .unwrap_or_default();
1921
1922 let facts = self
1923 .observe_purpose(
1924 ObservationPurpose::FactsExtraction,
1925 extractor.extract(&recent, &existing, actor_id.as_deref(), &categories),
1926 )
1927 .await?;
1928
1929 if !facts.is_empty() {
1931 let fact_store_opt = self.fact_store.read().clone();
1932 let mut stored_total = 0usize;
1933 let mut cache_updated = false;
1934 if let (Some(store), Some(aid)) = (fact_store_opt, &actor_id) {
1935 let authoritative = store.add_facts(aid, facts.clone()).await?;
1937 stored_total = authoritative.len();
1938 self.actor_facts_cache
1939 .write()
1940 .insert(aid.clone(), authoritative);
1941 cache_updated = true;
1942 } else if let Some(aid) = &actor_id {
1943 let mut cache = self.actor_facts_cache.write();
1944 let entry = cache.entry(aid.clone()).or_default();
1945 entry.extend(facts.clone());
1946 stored_total = entry.len();
1947 cache_updated = true;
1948 }
1949
1950 info!(
1951 actor_id = %actor_id.as_deref().unwrap_or("<none>"),
1952 source = source,
1953 requested_messages = last_n,
1954 message_count = recent.len(),
1955 extracted_count = facts.len(),
1956 cache_updated = cache_updated,
1957 stored_total = stored_total,
1958 "facts extracted"
1959 );
1960
1961 if let Some(ref aid) = actor_id {
1962 self.hooks.on_facts_extracted(aid, &facts).await;
1963 }
1964 }
1965
1966 Ok(facts)
1967 }
1968
1969 fn resolve_actor_id_from_context(&self) {
1972 if self
1973 .current_turn_actor_context()
1974 .and_then(|ctx| ctx.effective_actor_id().map(str::to_string))
1975 .is_some()
1976 {
1977 return;
1978 }
1979
1980 if let Some(ref am_config) = self.actor_memory_config
1981 && am_config.identification.method == ai_agents_facts::IdentificationMethod::FromContext
1982 && let Some(ref path) = am_config.identification.context_path
1983 {
1984 let val = self
1986 .context_manager
1987 .get_path(path)
1988 .or_else(|| self.context_manager.get(path));
1989 if let Some(val) = val
1990 && let Some(id_str) = val.as_str()
1991 {
1992 let current = self.actor_id.read().clone();
1993 if current.as_deref() != Some(id_str) {
1994 *self.actor_id.write() = Some(id_str.to_string());
1995 let mut meta = self.session_metadata.write();
1996 meta.actor_id = Some(id_str.to_string());
1997 if !meta.actors.iter().any(|a| a == id_str) {
1998 meta.actors.push(id_str.to_string());
1999 }
2000 }
2001 }
2002 }
2003 }
2004
2005 fn format_actor_facts_for_context(&self) -> String {
2007 let should_inject = self
2009 .facts_config
2010 .as_ref()
2011 .map(|c| c.inject_in_context)
2012 .unwrap_or(true);
2013 if !should_inject {
2014 return String::new();
2015 }
2016
2017 let Some(actor_id) = self.effective_actor_id() else {
2018 return String::new();
2019 };
2020
2021 let facts = self
2022 .actor_facts_cache
2023 .read()
2024 .get(&actor_id)
2025 .cloned()
2026 .unwrap_or_default();
2027 if facts.is_empty() {
2028 return String::new();
2029 }
2030
2031 let am_config = self.actor_memory_config.as_ref();
2032 let facts_budget = self
2035 .memory_token_budget
2036 .as_ref()
2037 .map(|b| b.allocation.facts as usize)
2038 .filter(|n| *n > 0);
2039 let default_max = am_config.map(|c| c.injection.max_tokens).unwrap_or(800);
2040 let max_tokens = facts_budget.unwrap_or(default_max);
2041
2042 let filtered: Vec<ai_agents_core::KeyFact> = if let Some(cfg) = am_config {
2044 if cfg.injection.mode == ai_agents_facts::InjectionMode::OnDemand {
2045 return String::new();
2046 }
2047 if cfg.injection.mode == ai_agents_facts::InjectionMode::Category
2048 && !cfg.injection.categories.is_empty()
2049 {
2050 facts
2051 .iter()
2052 .filter(|f| {
2053 cfg.injection
2054 .categories
2055 .iter()
2056 .any(|c| f.category.to_string() == *c)
2057 })
2058 .cloned()
2059 .collect()
2060 } else {
2061 facts.clone()
2062 }
2063 } else {
2064 facts.clone()
2065 };
2066
2067 if filtered.is_empty() {
2068 return String::new();
2069 }
2070
2071 if let Some(store) = self.fact_store.read().clone() {
2072 store.format_for_context(&filtered, max_tokens)
2073 } else {
2074 String::new()
2075 }
2076 }
2077
2078 fn build_context_with_staged(&self, staged: &HashMap<String, Value>) -> HashMap<String, Value> {
2079 let context = self.build_context_with_overlays();
2080 let mut root = Value::Object(context.into_iter().collect());
2081 for (path, value) in staged {
2082 if let Ok(updated) = ai_agents_core::set_dot_path(root.clone(), path, value.clone()) {
2083 root = updated;
2084 }
2085 }
2086 match root {
2087 Value::Object(obj) => obj.into_iter().collect(),
2088 _ => HashMap::new(),
2089 }
2090 }
2091
2092 fn build_context_with_overlays(&self) -> HashMap<String, Value> {
2093 let mut context = self.context_manager.get_all();
2094 let mut root = Value::Object(context.clone().into_iter().collect());
2095
2096 if let Some(turn_ctx) = self.current_turn_actor_context() {
2097 if let Some(ref origin_actor_id) = turn_ctx.origin_actor_id
2098 && let Ok(updated) = ai_agents_core::set_dot_path(
2099 root.clone(),
2100 "interaction.origin_actor_id",
2101 serde_json::json!(origin_actor_id),
2102 )
2103 {
2104 root = updated;
2105 }
2106 if let Some(ref sender_agent_id) = turn_ctx.sender_agent_id
2107 && let Ok(updated) = ai_agents_core::set_dot_path(
2108 root.clone(),
2109 "interaction.sender_agent_id",
2110 serde_json::json!(sender_agent_id),
2111 )
2112 {
2113 root = updated;
2114 }
2115 }
2116
2117 if let Some(ref actor_id) = self.effective_actor_id()
2118 && let Ok(updated) = ai_agents_core::set_dot_path(
2119 root.clone(),
2120 "interaction.actor_id",
2121 serde_json::json!(actor_id),
2122 )
2123 {
2124 root = updated;
2125 }
2126
2127 if let Some(manager) = self.relationship_manager.as_ref()
2128 && let Some(actor_id) = self.effective_actor_id()
2129 && let Some(value) = manager.to_context_value(&actor_id)
2130 && let Ok(updated) = ai_agents_core::set_dot_path(
2131 root.clone(),
2132 &manager.config().injection.context_path,
2133 value,
2134 )
2135 {
2136 root = updated;
2137 }
2138
2139 if let Value::Object(obj) = root {
2140 context = obj.into_iter().collect();
2141 }
2142
2143 context
2144 }
2145
2146 fn resolve_actor_name_from_context(&self) -> Option<String> {
2147 for path in ["actor.name", "user.name", "player.name", "customer.name"] {
2148 if let Some(value) = self.context_manager.get_path(path)
2149 && let Some(name) = value.as_str()
2150 {
2151 return Some(name.to_string());
2152 }
2153 }
2154 None
2155 }
2156
2157 async fn maybe_load_actor_relationship(&self) {
2158 let Some(manager) = self.relationship_manager.as_ref() else {
2159 return;
2160 };
2161 let Some(actor_id) = self.effective_actor_id() else {
2162 return;
2163 };
2164
2165 let mut should_fire_loaded = false;
2166 if manager.get(&actor_id).is_none() {
2167 let mut loaded = false;
2168 if manager.config().persistence.enabled {
2169 let storage = self.storage.read().clone();
2170 if let Some(storage) = storage {
2171 match storage.load_relationship(&self.info.id, &actor_id).await {
2172 Ok(Some(value)) => match manager.insert_from_value(value) {
2173 Ok(_) => loaded = true,
2174 Err(e) => {
2175 warn!(actor = %actor_id, error = %e, "failed to restore relationship")
2176 }
2177 },
2178 Ok(None) => {}
2179 Err(e) => {
2180 warn!(actor = %actor_id, error = %e, "failed to load relationship")
2181 }
2182 }
2183 }
2184 }
2185
2186 if !loaded {
2187 manager.get_or_create(&actor_id, self.resolve_actor_name_from_context().as_deref());
2188 }
2189 should_fire_loaded = true;
2190 }
2191
2192 let actor_name = self.resolve_actor_name_from_context();
2193 let relationship = manager.touch_interaction(&actor_id, actor_name.as_deref());
2194 if should_fire_loaded {
2195 self.hooks
2196 .on_relationship_loaded(&actor_id, &relationship)
2197 .await;
2198 }
2199 }
2200
2201 fn format_relationship_for_context(&self) -> Option<(String, String)> {
2202 let manager = self.relationship_manager.as_ref()?;
2203 if !manager.config().injection.enabled {
2204 return None;
2205 }
2206 let actor_id = self.effective_actor_id()?;
2207 let relationship = manager.get(&actor_id)?;
2208 let local_cap = manager.config().injection.max_tokens;
2209 let global_cap = self
2210 .memory_token_budget
2211 .as_ref()
2212 .map(|b| b.allocation.relationships as usize)
2213 .filter(|n| *n > 0);
2214 let max_tokens = global_cap.map(|g| g.min(local_cap)).unwrap_or(local_cap);
2215 let text = ai_agents_relationships::format_relationship(
2216 &relationship,
2217 &manager.config().injection.format,
2218 max_tokens,
2219 );
2220 if text.is_empty() {
2221 None
2222 } else {
2223 Some((manager.config().injection.prompt_variable.clone(), text))
2224 }
2225 }
2226
2227 async fn persist_actor_relationship(&self, actor_id: &str) -> Result<()> {
2228 let Some(manager) = self.relationship_manager.as_ref() else {
2229 return Ok(());
2230 };
2231 if !manager.config().persistence.enabled {
2232 return Ok(());
2233 }
2234 let storage = self.storage.read().clone();
2235 let Some(storage) = storage else {
2236 return Ok(());
2237 };
2238 if let Some(value) = manager.relationship_as_value(actor_id)? {
2239 storage
2240 .save_relationship(&self.info.id, actor_id, &value)
2241 .await?;
2242 }
2243 Ok(())
2244 }
2245
2246 async fn auto_update_relationship(&self) {
2247 let Some(manager) = self.relationship_manager.as_ref() else {
2248 return;
2249 };
2250 let Some(actor_id) = self.effective_actor_id() else {
2251 return;
2252 };
2253 if !manager.config().auto_update.enabled {
2254 let _ = self.persist_actor_relationship(&actor_id).await;
2255 return;
2256 }
2257
2258 let recent_messages = manager.config().auto_update.recent_messages;
2259 let messages = match self.memory.get_messages(Some(recent_messages)).await {
2260 Ok(messages) => messages,
2261 Err(e) => {
2262 warn!(actor = %actor_id, error = %e, "failed to read messages for relationship update");
2263 return;
2264 }
2265 };
2266 let messages = match Self::readable_native_messages(messages) {
2267 Ok(messages) => messages,
2268 Err(error) => {
2269 warn!(actor = %actor_id, error = %error, "failed to project native history for relationship update");
2270 return;
2271 }
2272 };
2273
2274 match self
2275 .observe_purpose(
2276 ObservationPurpose::RelationshipUpdate,
2277 manager.auto_update(&actor_id, &messages),
2278 )
2279 .await
2280 {
2281 Ok(update) => {
2282 if !update.changes.is_empty() {
2283 self.hooks
2284 .on_relationship_change(&actor_id, &update.changes)
2285 .await;
2286 }
2287 if let Some(ref event) = update.event {
2288 self.hooks.on_notable_event(&actor_id, event).await;
2289 }
2290 let persisted = match self.persist_actor_relationship(&actor_id).await {
2291 Ok(()) => true,
2292 Err(e) => {
2293 warn!(actor = %actor_id, error = %e, "failed to persist relationship");
2294 false
2295 }
2296 };
2297 if !update.changes.is_empty() || update.event.is_some() {
2298 let changed_dimensions: Vec<String> = update
2299 .changes
2300 .iter()
2301 .map(|change| format!("{}:{}", change.perspective, change.dimension))
2302 .collect();
2303 info!(
2304 actor_id = %actor_id,
2305 change_count = update.changes.len(),
2306 changed_dimensions = ?changed_dimensions,
2307 event_present = update.event.is_some(),
2308 persisted = persisted,
2309 "relationship updated"
2310 );
2311 } else {
2312 debug!(actor_id = %actor_id, persisted = persisted, "relationship evaluation ran but found no changes");
2313 }
2314 }
2315 Err(e) => warn!(actor = %actor_id, error = %e, "relationship update failed"),
2316 }
2317 }
2318
2319 async fn auto_extract_facts(&self) {
2321 let should_extract = self
2322 .facts_config
2323 .as_ref()
2324 .map(|c| c.enabled && c.auto_extract)
2325 .unwrap_or(false);
2326
2327 if !should_extract {
2328 debug!("fact extraction skipped because auto extraction is disabled");
2329 return;
2330 }
2331
2332 let msgs_since = *self.messages_since_extraction.read();
2333 if msgs_since < 2 {
2334 debug!(
2335 messages_since_extraction = msgs_since,
2336 "fact extraction skipped until threshold is reached"
2337 );
2338 return;
2339 }
2340
2341 match self.extract_facts_with_source(msgs_since, "auto").await {
2342 Ok(facts) => {
2343 if !facts.is_empty() {
2344 *self.messages_since_extraction.write() = 0;
2345 } else {
2346 debug!("fact extraction ran but found no new facts");
2347 }
2348 }
2349 Err(e) => {
2350 warn!("fact extraction failed: {}", e);
2351 }
2352 }
2353 }
2354
2355 pub fn with_persona(mut self, manager: Arc<ai_agents_persona::PersonaManager>) -> Self {
2356 self.persona_manager = Some(manager);
2357 self
2358 }
2359
2360 pub fn persona_manager(&self) -> Option<&Arc<ai_agents_persona::PersonaManager>> {
2361 self.persona_manager.as_ref()
2362 }
2363
2364 pub fn with_disambiguation(mut self, config: DisambiguationConfig) -> Self {
2365 if config.is_enabled() {
2366 let manager = DisambiguationManager::new(config, Arc::clone(&self.llm_registry))
2367 .with_clarification_observer(Arc::new(ObservabilityClarificationObserver));
2368 self.disambiguation_manager = Some(manager);
2369 }
2370 self
2371 }
2372
2373 pub fn disambiguation_manager(&self) -> Option<&DisambiguationManager> {
2374 self.disambiguation_manager.as_ref()
2375 }
2376
2377 pub fn has_disambiguation(&self) -> bool {
2378 self.disambiguation_manager
2379 .as_ref()
2380 .is_some_and(|m| m.is_enabled())
2381 }
2382
2383 pub async fn init_storage(&self) -> Result<()> {
2384 let _guard = self.storage_init.lock().await;
2388 let mut storage = self.storage.read().clone();
2389 if storage.is_none() && !self.storage_config.is_none() {
2390 let storage_config = self.convert_storage_config();
2391 storage = create_storage(&storage_config).await?;
2392 *self.storage.write() = storage.clone();
2393 }
2394
2395 self.validate_storage_requirements(storage.as_deref())?;
2396 self.complete_facts_init().await;
2397 Ok(())
2398 }
2399
2400 fn validate_storage_requirements(&self, storage: Option<&dyn AgentStorage>) -> Result<()> {
2401 let facts_required = self
2402 .facts_config
2403 .as_ref()
2404 .is_some_and(|config| config.enabled)
2405 || self
2406 .actor_memory_config
2407 .as_ref()
2408 .is_some_and(|config| config.enabled);
2409 let relationships_required = self
2410 .relationship_manager
2411 .as_ref()
2412 .is_some_and(|manager| manager.config().persistence.enabled);
2413
2414 let Some(storage) = storage else {
2415 let mut requirements = Vec::new();
2416 if facts_required {
2417 requirements.push("actor facts or actor memory");
2418 }
2419 if relationships_required {
2420 requirements.push("persistent relationships");
2421 }
2422 if requirements.is_empty() {
2423 return Ok(());
2424 }
2425 return Err(AgentError::Config(format!(
2426 "Storage is required for enabled {} but none is configured or injected",
2427 requirements.join(" and ")
2428 )));
2429 };
2430
2431 if facts_required && !storage.supports(StorageCapability::ActorFacts) {
2435 return Err(AgentError::UnsupportedStorageCapability(
2436 StorageCapability::ActorFacts,
2437 ));
2438 }
2439 if relationships_required && !storage.supports(StorageCapability::ActorRelationships) {
2440 return Err(AgentError::UnsupportedStorageCapability(
2441 StorageCapability::ActorRelationships,
2442 ));
2443 }
2444 Ok(())
2445 }
2446
2447 async fn complete_facts_init(&self) {
2450 if self.fact_store.read().is_some() {
2451 return;
2452 }
2453 let storage = match self.storage.read().clone() {
2454 Some(s) => s,
2455 None => return,
2456 };
2457
2458 let facts_enabled = self
2459 .facts_config
2460 .as_ref()
2461 .map(|f| f.enabled)
2462 .unwrap_or(false);
2463 let actor_memory_enabled = self
2464 .actor_memory_config
2465 .as_ref()
2466 .map(|a| a.enabled)
2467 .unwrap_or(false);
2468
2469 if !facts_enabled && !actor_memory_enabled {
2470 return;
2471 }
2472
2473 let fc = self.facts_config.clone().unwrap_or_default();
2474 let store = Arc::new(ai_agents_facts::FactStore::new(
2475 storage,
2476 self.info.id.clone(),
2477 fc.clone(),
2478 ));
2479
2480 let extractor: Option<Arc<dyn ai_agents_facts::FactExtractor>> = if facts_enabled {
2481 let extractor_llm = fc
2482 .extractor_llm
2483 .as_ref()
2484 .and_then(|alias| self.llm_registry.get(alias).ok())
2485 .or_else(|| self.llm_registry.router().ok())
2486 .or_else(|| self.llm_registry.default().ok());
2487 extractor_llm.map(|llm| {
2488 Arc::new(ai_agents_facts::LLMFactExtractor::new(llm, fc.clone()))
2489 as Arc<dyn ai_agents_facts::FactExtractor>
2490 })
2491 } else {
2492 None
2493 };
2494
2495 *self.fact_store.write() = Some(store);
2496 *self.fact_extractor.write() = extractor;
2497 debug!(
2498 agent = %self.info.id,
2499 facts_enabled,
2500 actor_memory_enabled,
2501 "facts storage initialized"
2502 );
2503 }
2504
2505 fn convert_storage_config(&self) -> StorageStorageConfig {
2506 crate::spec::storage::to_storage_config(&self.storage_config)
2507 }
2508
2509 pub fn storage(&self) -> Option<Arc<dyn AgentStorage>> {
2510 self.storage.read().clone()
2511 }
2512
2513 pub fn storage_config(&self) -> &StorageConfig {
2514 &self.storage_config
2515 }
2516
2517 pub fn spawner(&self) -> Option<&Arc<crate::spawner::AgentSpawner>> {
2519 self.spawner.as_ref()
2520 }
2521
2522 pub fn spawner_registry(&self) -> Option<&Arc<crate::spawner::AgentRegistry>> {
2524 self.spawner_registry.as_ref()
2525 }
2526
2527 pub fn has_spawner(&self) -> bool {
2528 self.spawner_registry.is_some()
2529 }
2530
2531 pub fn with_spawner_handles(
2532 mut self,
2533 spawner: Arc<crate::spawner::AgentSpawner>,
2534 registry: Arc<crate::spawner::AgentRegistry>,
2535 ) -> Self {
2536 self.spawner = Some(spawner);
2537 self.spawner_registry = Some(registry);
2538 self
2539 }
2540
2541 pub fn with_hooks(mut self, hooks: Arc<dyn AgentHooks>) -> Self {
2542 self.hooks = hooks;
2543 self
2544 }
2545
2546 pub fn with_parallel_tools(mut self, config: ParallelToolsConfig) -> Self {
2547 self.parallel_tools = config;
2548 self
2549 }
2550
2551 pub fn with_streaming(mut self, config: StreamingConfig) -> Self {
2552 self.streaming = config;
2553 self
2554 }
2555
2556 pub fn with_hitl(mut self, engine: HITLEngine, handler: Arc<dyn ApprovalHandler>) -> Self {
2557 self.hitl_engine = Some(engine);
2558 self.approval_handler = handler;
2559 self
2560 }
2561
2562 pub fn with_max_context_tokens(mut self, tokens: u32) -> Self {
2563 self.max_context_tokens = tokens;
2564 self
2565 }
2566
2567 pub fn with_memory_token_budget(mut self, budget: MemoryTokenBudget) -> Self {
2568 self.memory_token_budget = Some(budget);
2569 self
2570 }
2571
2572 pub fn with_recovery_manager(mut self, manager: RecoveryManager) -> Self {
2573 self.recovery_manager = manager;
2574 self
2575 }
2576
2577 pub fn with_tool_security(mut self, engine: ToolSecurityEngine) -> Self {
2578 self.tool_security = engine;
2579 self
2580 }
2581
2582 pub fn runtime_control(&self) -> RuntimeControlHandle {
2584 RuntimeControlHandle {
2585 state: Arc::clone(&self.runtime_control),
2586 }
2587 }
2588
2589 pub fn set_question_handler(&self, handler: Option<Arc<dyn QuestionHandler>>) {
2591 self.tools.set_question_handler(handler);
2592 }
2593
2594 pub fn set_diagnostics_provider(&self, provider: Arc<dyn DiagnosticsProvider>) {
2596 self.tools.set_diagnostics_provider(provider);
2597 }
2598
2599 pub fn set_command_runner(&self, runner: Arc<dyn CommandRunner>) {
2601 self.tools.set_command_runner(runner);
2602 }
2603
2604 pub fn set_web_search_provider(&self, provider: Arc<dyn ai_agents_tools::WebSearchProvider>) {
2606 self.tools.set_web_search_provider(provider);
2607 }
2608
2609 pub fn todos(&self) -> Vec<TodoItem> {
2611 self.tools.todos()
2612 }
2613
2614 fn active_tool_security(&self) -> ToolSecurityEngine {
2616 self.runtime_control
2617 .tool_security_override
2618 .read()
2619 .clone()
2620 .unwrap_or_else(|| self.tool_security.clone())
2621 }
2622
2623 fn runtime_safety_snapshot(&self) -> RuntimeSafetySnapshot {
2625 let _guard = self.runtime_control.snapshot_guard.read();
2626 RuntimeSafetySnapshot {
2627 version: self.runtime_control.version.load(Ordering::SeqCst),
2628 emergency_deny: self.runtime_control.emergency_deny.load(Ordering::SeqCst),
2629 tool_security: self
2630 .runtime_control
2631 .tool_security_override
2632 .read()
2633 .clone()
2634 .unwrap_or_else(|| self.tool_security.clone()),
2635 tool_scope_override: self.runtime_control.tool_scope_override.read().clone(),
2636 }
2637 }
2638
2639 fn admit_tool_execution(
2641 &self,
2642 expected_runtime_version: u64,
2643 expected_policy_version: u64,
2644 expected_state_generation: Option<u64>,
2645 canonical_id: &str,
2646 ) -> SecurityCheckResult {
2647 let _guard = self.runtime_control.snapshot_guard.read();
2648 if self.runtime_control.emergency_deny.load(Ordering::SeqCst) {
2649 return SecurityCheckResult::Block {
2650 reason: "runtime emergency deny is enabled".to_string(),
2651 };
2652 }
2653 let runtime_version = self.runtime_control.version.load(Ordering::SeqCst);
2654 let security_engine = self
2655 .runtime_control
2656 .tool_security_override
2657 .read()
2658 .clone()
2659 .unwrap_or_else(|| self.tool_security.clone());
2660 if runtime_version != expected_runtime_version
2661 || security_engine.policy_version() != expected_policy_version
2662 {
2663 return SecurityCheckResult::Block {
2664 reason: "runtime safety controls changed before admission".to_string(),
2665 };
2666 }
2667 let current_state_generation = self
2668 .state_machine
2669 .as_ref()
2670 .map(|state_machine| state_machine.generation());
2671 if current_state_generation != expected_state_generation {
2672 return SecurityCheckResult::Block {
2673 reason: "state scope changed before admission".to_string(),
2674 };
2675 }
2676 security_engine.admit_tool_execution(canonical_id)
2677 }
2678
2679 pub fn with_process_processor(mut self, processor: ProcessProcessor) -> Self {
2680 let processor = processor.with_stage_observer(Arc::new(ObservabilityProcessStageObserver));
2681 self.process_processor = Some(processor);
2682 self
2683 }
2684
2685 pub fn with_state_machine(
2686 mut self,
2687 state_machine: Arc<StateMachine>,
2688 evaluator: Arc<dyn TransitionEvaluator>,
2689 ) -> Self {
2690 self.state_machine = Some(state_machine);
2691 self.transition_evaluator = Some(evaluator);
2692 self
2693 }
2694
2695 pub fn with_context_manager(mut self, manager: Arc<ContextManager>) -> Self {
2696 self.context_manager = manager;
2697 self
2698 }
2699
2700 pub fn register_message_filter(&self, name: impl Into<String>, filter: Arc<dyn MessageFilter>) {
2701 self.message_filters.write().insert(name.into(), filter);
2702 }
2703
2704 pub fn set_context(&self, key: &str, value: Value) -> Result<()> {
2705 self.context_manager.update(key, value)
2706 }
2707
2708 pub fn update_context(&self, path: &str, value: Value) -> Result<()> {
2709 self.context_manager.update(path, value)
2710 }
2711
2712 pub fn get_context(&self) -> HashMap<String, Value> {
2713 self.build_context_with_overlays()
2714 }
2715
2716 pub fn remove_context(&self, key: &str) -> Option<Value> {
2717 self.context_manager.remove(key)
2718 }
2719
2720 pub async fn refresh_context(&self, key: &str) -> Result<()> {
2721 self.context_manager.refresh(key).await
2722 }
2723
2724 pub fn register_context_provider(&self, name: &str, provider: Arc<dyn ContextProvider>) {
2725 self.context_manager.register_provider(name, provider);
2726 }
2727
2728 pub fn current_state(&self) -> Option<String> {
2729 self.state_machine.as_ref().map(|sm| sm.current())
2730 }
2731
2732 async fn invalidate_pending_confirmation(&self, reason: &'static str) {
2734 self.disambiguation_epoch.fetch_add(1, Ordering::SeqCst);
2735 let Some(disambiguator) = self.disambiguation_manager.as_ref() else {
2736 return;
2737 };
2738 if disambiguator.has_pending_confirmation().await {
2739 disambiguator.clear_pending().await;
2740 *self.pending_skill_id.write() = None;
2741 info!(
2742 confirmation_event = "invalidated",
2743 invalidation_reason = reason,
2744 "Runtime invalidated pending confirmation"
2745 );
2746 }
2747 }
2748
2749 async fn admit_disambiguation_redispatch(
2751 &self,
2752 expected_epoch: u64,
2753 expected_state_generation: Option<u64>,
2754 ) -> Result<tokio::sync::RwLockReadGuard<'_, ()>> {
2755 let admission = self.disambiguation_admission.read().await;
2756 let state_generation = self
2757 .state_machine
2758 .as_ref()
2759 .map(|state_machine| state_machine.generation());
2760 if self.disambiguation_epoch.load(Ordering::SeqCst) != expected_epoch
2761 || state_generation != expected_state_generation
2762 {
2763 return Err(AgentError::Other(
2764 "Disambiguation ownership changed before redispatch admission".to_string(),
2765 ));
2766 }
2767 Ok(admission)
2768 }
2769
2770 fn reserve_state_transition(&self) -> Option<StateTransitionReservation<'_>> {
2772 self.state_transition_reserved
2773 .compare_exchange(false, true, Ordering::SeqCst, Ordering::SeqCst)
2774 .ok()
2775 .map(|_| StateTransitionReservation {
2776 reserved: &self.state_transition_reserved,
2777 })
2778 }
2779
2780 async fn admit_optional_disambiguation_ownership(
2782 &self,
2783 ownership: Option<DisambiguationOwnership>,
2784 ) -> Result<Option<tokio::sync::RwLockReadGuard<'_, ()>>> {
2785 match ownership {
2786 Some(ownership) => self
2787 .admit_disambiguation_redispatch(ownership.epoch, ownership.state_generation)
2788 .await
2789 .map(Some),
2790 None => Ok(None),
2791 }
2792 }
2793
2794 pub async fn transition_to(&self, state: &str) -> Result<()> {
2796 let Some(ref sm) = self.state_machine else {
2797 return Ok(());
2798 };
2799 let claim_admission = self.disambiguation_admission.write().await;
2800 let reservation = self.reserve_state_transition().ok_or_else(|| {
2801 AgentError::Other("Another state transition is already in progress".to_string())
2802 })?;
2803 let from_state = sm.current();
2804 let expected_state_generation = sm.generation();
2805 let expected_disambiguation_epoch = self.disambiguation_epoch.load(Ordering::SeqCst);
2806 let history_before = sm.history();
2807 drop(claim_admission);
2808
2809 self.execute_state_exit_actions(&from_state).await;
2810
2811 let admission = self.disambiguation_admission.write().await;
2812 if sm.current() != from_state
2813 || sm.generation() != expected_state_generation
2814 || self.disambiguation_epoch.load(Ordering::SeqCst) != expected_disambiguation_epoch
2815 {
2816 return Err(AgentError::Other(
2817 "State ownership changed during manual transition preparation".to_string(),
2818 ));
2819 }
2820 sm.transition_to(state, "manual transition")?;
2821 self.invalidate_pending_confirmation("state_transition")
2822 .await;
2823 let entered = sm.current();
2824 let is_reentry = Self::state_was_previously_entered(&entered, &from_state, &history_before);
2825 drop(admission);
2826
2827 self.execute_state_enter_actions(&entered, is_reentry).await;
2828 drop(reservation);
2829 info!(to = %entered, "Manual state transition");
2830 Ok(())
2831 }
2832
2833 pub fn state_history(&self) -> Vec<StateTransitionEvent> {
2834 self.state_machine
2835 .as_ref()
2836 .map(|sm| sm.history())
2837 .unwrap_or_default()
2838 }
2839
2840 pub fn session_metadata(&self) -> ai_agents_core::SessionMetadata {
2842 self.session_metadata.read().clone()
2843 }
2844
2845 pub async fn delete_actor_data(&self, actor_id: &str) -> Result<()> {
2848 let allowed = self
2849 .actor_memory_config
2850 .as_ref()
2851 .map(|c| c.privacy.allow_deletion)
2852 .unwrap_or(true);
2853 if !allowed {
2854 return Err(AgentError::Config(
2855 "privacy.allow_deletion is false; actor data deletion is not permitted".into(),
2856 ));
2857 }
2858 let storage = self.storage.read().clone();
2859 if let Some(storage) = storage {
2860 if !storage.supports(StorageCapability::ActorDataDeletion) {
2864 return Err(AgentError::UnsupportedStorageCapability(
2865 StorageCapability::ActorDataDeletion,
2866 ));
2867 }
2868 storage.delete_actor_data(&self.info.id, actor_id).await?;
2869 } else {
2870 let store = { self.fact_store.read().clone() };
2874 if let Some(store) = store {
2875 store.delete_actor_data(actor_id).await?;
2876 }
2877 }
2878 if let Some(manager) = self.relationship_manager.as_ref() {
2879 manager.remove(actor_id);
2880 }
2881 self.actor_facts_cache.write().remove(actor_id);
2882 Ok(())
2883 }
2884
2885 pub fn set_session_metadata(&self, meta: ai_agents_core::SessionMetadata) {
2887 *self.session_metadata.write() = meta;
2888 }
2889
2890 pub async fn cleanup_expired_sessions(&self) -> Result<usize> {
2892 let storage = self.storage.read().clone();
2893 match storage {
2894 Some(s) => {
2895 let count = s.cleanup_expired().await?;
2896 if count > 0 {
2897 self.hooks.on_sessions_expired(count).await;
2898 }
2899 Ok(count)
2900 }
2901 None => Err(AgentError::Config(
2902 "No storage configured. Use with_storage_config() or with_storage() first".into(),
2903 )),
2904 }
2905 }
2906
2907 pub async fn list_sessions_filtered(
2909 &self,
2910 filter: &ai_agents_core::SessionFilter,
2911 ) -> Result<Vec<ai_agents_core::SessionSummary>> {
2912 let storage = self.storage.read().clone();
2913 match storage {
2914 Some(s) => s.list_sessions_filtered(filter).await,
2915 None => Err(AgentError::Config(
2916 "No storage configured. Use with_storage_config() or with_storage() first".into(),
2917 )),
2918 }
2919 }
2920
2921 pub async fn save_state(&self) -> Result<AgentSnapshot> {
2922 let memory_snapshot = self.memory.snapshot().await?;
2923 let state_machine_snapshot = self.state_machine.as_ref().map(|sm| sm.snapshot());
2924 let context_snapshot = self.context_manager.snapshot();
2925
2926 let mut snapshot = AgentSnapshot::new(self.info.id.clone())
2927 .with_memory(memory_snapshot)
2928 .with_context(context_snapshot)
2929 .with_state_machine(
2930 state_machine_snapshot.unwrap_or_else(|| StateMachineSnapshot {
2931 current_state: String::new(),
2932 previous_state: None,
2933 turn_count: 0,
2934 no_transition_count: 0,
2935 history: vec![],
2936 }),
2937 );
2938
2939 if let Some(ref persona) = self.persona_manager {
2940 snapshot.persona = Some(persona.snapshot_as_value()?);
2941 }
2942
2943 if let Some(ref relationships) = self.relationship_manager {
2944 snapshot.relationships = Some(relationships.snapshot_as_value()?);
2945 }
2946
2947 Ok(snapshot)
2948 }
2949
2950 pub async fn save_state_full(&self) -> Result<AgentSnapshot> {
2952 let mut snapshot = self.save_state().await?;
2953 if let Some(ref registry) = self.spawner_registry {
2954 let entries = registry.list_with_specs();
2955 if !entries.is_empty() {
2956 snapshot = snapshot.with_spawned_agents(entries);
2957 }
2958 }
2959 Ok(snapshot)
2960 }
2961
2962 pub async fn restore_state(&self, snapshot: AgentSnapshot) -> Result<()> {
2964 let _admission = self.disambiguation_admission.write().await;
2965 if self.state_transition_reserved.load(Ordering::SeqCst) {
2966 return Err(AgentError::Other(
2967 "Cannot restore state while a state transition is in progress".to_string(),
2968 ));
2969 }
2970 self.invalidate_pending_confirmation("state_restore").await;
2971 *self.pending_skill_id.write() = None;
2972 if let Some(disambiguator) = self.disambiguation_manager.as_ref() {
2973 disambiguator.clear_pending().await;
2974 }
2975 self.memory.restore(snapshot.memory).await?;
2976 self.active_native_exchanges.write().clear();
2977
2978 if let (Some(sm), Some(sm_snapshot)) = (&self.state_machine, snapshot.state_machine)
2979 && !sm_snapshot.current_state.is_empty()
2980 {
2981 sm.restore(sm_snapshot)?;
2982 }
2983
2984 self.context_manager.restore(snapshot.context);
2985
2986 if let (Some(persona_value), Some(persona_manager)) =
2987 (snapshot.persona, &self.persona_manager)
2988 {
2989 persona_manager.restore_from_value(persona_value)?;
2990 }
2991
2992 if let (Some(relationship_value), Some(relationship_manager)) =
2993 (snapshot.relationships, &self.relationship_manager)
2994 {
2995 relationship_manager.restore_from_value(relationship_value)?;
2996 }
2997
2998 info!(agent_id = %snapshot.agent_id, "State restored");
2999 Ok(())
3000 }
3001
3002 pub async fn save_to(&self, storage: &dyn AgentStorage, session_id: &str) -> Result<()> {
3003 let snapshot = self.save_state().await?;
3004 storage.save(session_id, &snapshot).await
3005 }
3006
3007 async fn load_session_restore(
3008 storage: &dyn AgentStorage,
3009 session_id: &str,
3010 ) -> Result<Option<StoredSessionRestore>> {
3011 let Some(snapshot) = storage.load(session_id).await? else {
3012 return Ok(None);
3013 };
3014 let metadata = if storage.supports(StorageCapability::SessionMetadata) {
3018 storage.load_metadata(session_id).await?
3019 } else {
3020 None
3021 };
3022 Ok(Some(StoredSessionRestore { snapshot, metadata }))
3023 }
3024
3025 async fn capture_session_restore_point(&self) -> Result<RuntimeSessionRestorePoint> {
3026 Ok(RuntimeSessionRestorePoint {
3027 snapshot: self.save_state().await?,
3028 metadata: self.session_metadata(),
3029 actor_id: self.actor_id(),
3030 session_id: self.current_session_id.read().clone(),
3031 })
3032 }
3033
3034 async fn apply_session_restore_unchecked(
3035 &self,
3036 session_id: &str,
3037 stored: StoredSessionRestore,
3038 ) -> Result<()> {
3039 self.restore_state(stored.snapshot).await?;
3040 let metadata = stored.metadata.unwrap_or_default();
3041 if let Some(actor_id) = metadata.actor_id.as_deref() {
3042 self.set_actor_id(actor_id)?;
3043 } else {
3044 self.clear_actor_id();
3045 }
3046 self.set_session_metadata(metadata);
3047 *self.current_session_id.write() = Some(session_id.to_string());
3048 Ok(())
3049 }
3050
3051 async fn restore_session_restore_point(
3052 &self,
3053 restore_point: &RuntimeSessionRestorePoint,
3054 ) -> Result<()> {
3055 self.restore_state(restore_point.snapshot.clone()).await?;
3056 if let Some(actor_id) = restore_point.actor_id.as_deref() {
3057 self.set_actor_id(actor_id)?;
3058 } else {
3059 self.clear_actor_id();
3060 }
3061 self.set_session_metadata(restore_point.metadata.clone());
3062 *self.current_session_id.write() = restore_point.session_id.clone();
3063 Ok(())
3064 }
3065
3066 async fn apply_session_restore(
3067 &self,
3068 session_id: &str,
3069 stored: StoredSessionRestore,
3070 ) -> Result<()> {
3071 let before = self.capture_session_restore_point().await?;
3072 if let Err(error) = self
3073 .apply_session_restore_unchecked(session_id, stored)
3074 .await
3075 {
3076 return match self.restore_session_restore_point(&before).await {
3077 Ok(()) => Err(error),
3078 Err(rollback_error) => Err(AgentError::Other(format!(
3079 "Session restore failed: {error}; rollback failed: {rollback_error}"
3080 ))),
3081 };
3082 }
3083 Ok(())
3084 }
3085
3086 async fn rollback_session_restore_set(
3087 parent: Option<(&RuntimeAgent, &RuntimeSessionRestorePoint)>,
3088 children: &[(String, Arc<RuntimeAgent>, RuntimeSessionRestorePoint)],
3089 ) -> Vec<String> {
3090 let mut errors = Vec::new();
3091 if let Some((agent, restore_point)) = parent
3092 && let Err(error) = agent.restore_session_restore_point(restore_point).await
3093 {
3094 errors.push(format!("parent: {error}"));
3095 }
3096 for (id, agent, restore_point) in children {
3097 if let Err(error) = agent.restore_session_restore_point(restore_point).await {
3098 errors.push(format!("child '{id}': {error}"));
3099 }
3100 }
3101 errors
3102 }
3103
3104 fn restore_failure(error: impl std::fmt::Display, rollback_errors: Vec<String>) -> AgentError {
3105 if rollback_errors.is_empty() {
3106 AgentError::Other(format!(
3107 "Session restore failed: {error}; runtime state was rolled back"
3108 ))
3109 } else {
3110 AgentError::Other(format!(
3111 "Session restore failed: {error}; rollback also failed for {}",
3112 rollback_errors.join(", ")
3113 ))
3114 }
3115 }
3116
3117 pub async fn load_from(&self, storage: &dyn AgentStorage, session_id: &str) -> Result<bool> {
3118 let Some(stored) = Self::load_session_restore(storage, session_id).await? else {
3119 return Ok(false);
3120 };
3121 self.apply_session_restore(session_id, stored).await?;
3122 Ok(true)
3123 }
3124
3125 pub async fn save_session(&self, session_id: &str) -> Result<()> {
3126 let storage = self.storage.read().clone();
3127 match storage {
3128 Some(s) => {
3129 let is_new = {
3131 let cur = self.current_session_id.read().clone();
3132 cur.as_deref() != Some(session_id)
3133 };
3134 if is_new {
3135 *self.current_session_id.write() = Some(session_id.to_string());
3136 self.hooks.on_session_created(session_id).await;
3137 }
3138
3139 {
3141 let now = chrono::Utc::now();
3142 let msg_count = self
3143 .memory
3144 .get_messages(None)
3145 .await
3146 .map(|v| v.len())
3147 .unwrap_or(0);
3148 let mut meta = self.session_metadata.write();
3149 meta.last_active = now;
3150 meta.message_count = msg_count;
3151 if meta.actor_id.is_none() {
3152 meta.actor_id = self.actor_id.read().clone();
3153 }
3154 }
3155
3156 let snapshot = self.save_state().await?;
3157 if s.supports(StorageCapability::SessionMetadata) {
3161 let metadata = self.session_metadata.read().clone();
3162 s.save_snapshot_with_metadata(session_id, &snapshot, &metadata)
3163 .await
3164 } else {
3165 s.save(session_id, &snapshot).await
3166 }
3167 }
3168 None => Err(AgentError::Config(
3169 "No storage configured. Use with_storage_config() or with_storage() first".into(),
3170 )),
3171 }
3172 }
3173
3174 pub async fn load_session(&self, session_id: &str) -> Result<bool> {
3175 let storage = self.storage.read().clone();
3176 match storage {
3177 Some(storage) => self.load_from(storage.as_ref(), session_id).await,
3178 None => Err(AgentError::Config(
3179 "No storage configured. Use with_storage_config() or with_storage() first".into(),
3180 )),
3181 }
3182 }
3183
3184 pub async fn restore_session_full(&self, session_id: &str) -> Result<usize> {
3186 self.init_storage().await?;
3187 let storage = self.storage.read().clone().ok_or_else(|| {
3188 AgentError::Config(
3189 "No storage configured. Use with_storage_config() or with_storage() first".into(),
3190 )
3191 })?;
3192 let target_parent = Self::load_session_restore(storage.as_ref(), session_id)
3193 .await?
3194 .ok_or_else(|| AgentError::Persistence(format!("Session not found: {session_id}")))?;
3195 let manifest = target_parent
3196 .snapshot
3197 .spawned_agents
3198 .clone()
3199 .unwrap_or_default();
3200
3201 let registry = self.spawner_registry.as_ref().cloned();
3202 let spawner = if manifest.is_empty() {
3203 self.spawner.as_ref().cloned()
3204 } else {
3205 Some(self.spawner.as_ref().cloned().ok_or_else(|| {
3206 AgentError::Config(
3207 "Saved session contains child agents but this runtime has no spawner".into(),
3208 )
3209 })?)
3210 };
3211 let registry = if manifest.is_empty() {
3212 registry
3213 } else {
3214 Some(registry.ok_or_else(|| {
3215 AgentError::Config(
3216 "Saved session contains child agents but this runtime has no registry".into(),
3217 )
3218 })?)
3219 };
3220
3221 let mut target_ids = HashSet::with_capacity(manifest.len());
3222 let mut prepared = Vec::with_capacity(manifest.len());
3223 for entry in manifest {
3224 if !target_ids.insert(entry.id.clone()) {
3225 return Err(AgentError::InvalidSpec(format!(
3226 "Saved child manifest contains duplicate ID: {}",
3227 entry.id
3228 )));
3229 }
3230 let spec = crate::spec::AgentSpec::from_yaml_strict(&entry.spec_yaml)?;
3231 spawner
3232 .as_ref()
3233 .expect("non-empty manifests require a spawner")
3234 .validate_explicit_child(&entry.id, &spec)?;
3235 prepared.push((entry.id, spec));
3236 }
3237
3238 let current_ids = registry
3239 .as_ref()
3240 .map(|registry| {
3241 registry
3242 .list()
3243 .into_iter()
3244 .map(|info| info.id)
3245 .collect::<HashSet<_>>()
3246 })
3247 .unwrap_or_default();
3248 let removal_count = current_ids.difference(&target_ids).count();
3249 let additions = prepared
3250 .iter()
3251 .filter(|(id, _)| !current_ids.contains(id))
3252 .cloned()
3253 .collect::<Vec<_>>();
3254
3255 let mut existing = Vec::new();
3256 if let Some(registry) = registry.as_ref() {
3257 for (id, _) in prepared.iter().filter(|(id, _)| current_ids.contains(id)) {
3258 let agent = registry.get(id).ok_or_else(|| {
3259 AgentError::Config(format!("Retained child disappeared during restore: {id}"))
3260 })?;
3261 let child_storage = agent.storage().ok_or_else(|| {
3262 AgentError::Config(format!("Child '{id}' has no storage for session restore"))
3263 })?;
3264 let stored = Self::load_session_restore(child_storage.as_ref(), session_id)
3265 .await?
3266 .ok_or_else(|| {
3267 AgentError::Persistence(format!(
3268 "Child '{id}' has no saved session '{session_id}'"
3269 ))
3270 })?;
3271 existing.push((id.clone(), agent, stored));
3272 }
3273 }
3274
3275 let mut staged = Vec::with_capacity(additions.len());
3276 if !additions.is_empty() {
3277 let spawner = spawner
3278 .as_ref()
3279 .expect("restored additions require a spawner");
3280 let reservations = spawner.reserve_restore_capacity(additions.len(), removal_count)?;
3281 for ((id, spec), reservation) in additions.into_iter().zip(reservations) {
3282 let spawned = spawner
3283 .spawn_with_reserved_capacity(id.clone(), spec, reservation)
3284 .await?;
3285 let child_storage = spawned.agent.storage().ok_or_else(|| {
3286 AgentError::Config(format!("Child '{id}' has no storage for session restore"))
3287 })?;
3288 let stored = Self::load_session_restore(child_storage.as_ref(), session_id)
3289 .await?
3290 .ok_or_else(|| {
3291 AgentError::Persistence(format!(
3292 "Child '{id}' has no saved session '{session_id}'"
3293 ))
3294 })?;
3295 staged.push((spawned, stored));
3296 }
3297 } else if let Some(spawner) = spawner.as_ref() {
3298 spawner.reserve_restore_capacity(0, removal_count)?;
3299 }
3300
3301 let parent_before = self.capture_session_restore_point().await?;
3302 let mut existing_before = Vec::with_capacity(existing.len());
3303 for (id, agent, _) in &existing {
3304 existing_before.push((
3305 id.clone(),
3306 Arc::clone(agent),
3307 agent.capture_session_restore_point().await?,
3308 ));
3309 }
3310
3311 for (_, agent, stored) in &existing {
3315 if let Err(error) = agent
3316 .apply_session_restore_unchecked(session_id, stored.clone())
3317 .await
3318 {
3319 drop(staged);
3320 let rollback_errors =
3321 Self::rollback_session_restore_set(None, &existing_before).await;
3322 return Err(Self::restore_failure(error, rollback_errors));
3323 }
3324 }
3325 for (spawned, stored) in &staged {
3326 if let Err(error) = spawned
3327 .agent
3328 .apply_session_restore_unchecked(session_id, stored.clone())
3329 .await
3330 {
3331 drop(staged);
3332 let rollback_errors =
3333 Self::rollback_session_restore_set(None, &existing_before).await;
3334 return Err(Self::restore_failure(error, rollback_errors));
3335 }
3336 }
3337 if let Err(error) = self
3338 .apply_session_restore_unchecked(session_id, target_parent)
3339 .await
3340 {
3341 drop(staged);
3342 let rollback_errors =
3343 Self::rollback_session_restore_set(Some((self, &parent_before)), &existing_before)
3344 .await;
3345 return Err(Self::restore_failure(error, rollback_errors));
3346 }
3347
3348 if let Some(registry) = registry.as_ref()
3349 && let Err(error) = registry
3350 .reconcile(
3351 &target_ids,
3352 staged.into_iter().map(|(spawned, _)| spawned).collect(),
3353 )
3354 .await
3355 {
3356 let rollback_errors =
3357 Self::rollback_session_restore_set(Some((self, &parent_before)), &existing_before)
3358 .await;
3359 return Err(Self::restore_failure(error, rollback_errors));
3360 }
3361
3362 Ok(target_ids.len())
3363 }
3364
3365 pub async fn delete_session(&self, session_id: &str) -> Result<()> {
3366 let storage = self.storage.read().clone();
3367 match storage {
3368 Some(s) => s.delete(session_id).await,
3369 None => Err(AgentError::Config(
3370 "No storage configured. Use with_storage_config() or with_storage() first".into(),
3371 )),
3372 }
3373 }
3374
3375 pub async fn list_sessions(&self) -> Result<Vec<String>> {
3376 let storage = self.storage.read().clone();
3377 match storage {
3378 Some(s) => s.list_sessions().await,
3379 None => Err(AgentError::Config(
3380 "No storage configured. Use with_storage_config() or with_storage() first".into(),
3381 )),
3382 }
3383 }
3384
3385 fn estimate_tokens(&self, text: &str) -> u32 {
3386 (text.len() as f32 / 4.0).ceil() as u32
3387 }
3388
3389 fn estimate_total_tokens(&self, messages: &[ChatMessage]) -> u32 {
3390 messages
3391 .iter()
3392 .map(|m| self.estimate_tokens(&m.content))
3393 .sum()
3394 }
3395
3396 fn native_safe_prefix_at_least(messages: &[ChatMessage], required: usize) -> Result<usize> {
3398 let inspection =
3399 inspect_native_history(messages).map_err(|error| AgentError::LLM(error.to_string()))?;
3400 let has_signed_history = !inspection.exchanges().is_empty();
3401 for count in required.min(messages.len())..=messages.len() {
3402 if has_signed_history
3403 && count < messages.len()
3404 && messages[count].role != ai_agents_core::Role::User
3405 {
3406 continue;
3407 }
3408 if inspection.is_safe_prefix_len(count)
3409 && inspect_native_history(&messages[count..]).is_ok()
3410 {
3411 return Ok(count);
3412 }
3413 }
3414 Err(AgentError::LLM(
3415 "Context limits cannot remove a complete native history prefix".to_string(),
3416 ))
3417 }
3418
3419 fn truncate_context(&self, messages: &mut Vec<ChatMessage>, keep_recent: usize) -> Result<()> {
3421 if messages.len() <= keep_recent + 1 {
3422 return Ok(());
3423 }
3424 let system_msg = messages.remove(0);
3425 let required = messages.len().saturating_sub(keep_recent);
3426 let to_remove = Self::native_safe_prefix_at_least(messages, required)?;
3427 messages.drain(..to_remove);
3428 messages.insert(0, system_msg);
3429 Ok(())
3430 }
3431
3432 fn get_filter(&self, config: &FilterConfig) -> Arc<dyn MessageFilter> {
3433 match config {
3434 FilterConfig::KeepRecent(n) => Arc::new(KeepRecentFilter::new(*n)),
3435 FilterConfig::ByRole { keep_roles } => Arc::new(ByRoleFilter::new(keep_roles.clone())),
3436 FilterConfig::SkipPattern { skip_if_contains } => {
3437 Arc::new(SkipPatternFilter::new(skip_if_contains.clone()))
3438 }
3439 FilterConfig::Custom { name } => {
3440 let filters = self.message_filters.read();
3441 filters
3442 .get(name)
3443 .cloned()
3444 .unwrap_or_else(|| Arc::new(KeepRecentFilter::new(10)))
3445 }
3446 }
3447 }
3448
3449 async fn summarize_context(
3450 &self,
3451 messages: &mut Vec<ChatMessage>,
3452 summarizer_llm: Option<&str>,
3453 max_summary_tokens: u32,
3454 custom_prompt: Option<&str>,
3455 keep_recent: usize,
3456 filter: Option<&FilterConfig>,
3457 ) -> Result<()> {
3458 let system_msg = messages.remove(0);
3459
3460 let required = messages.len().saturating_sub(keep_recent);
3461 if required == 0 {
3462 messages.insert(0, system_msg);
3463 return Ok(());
3464 }
3465 let to_summarize_count = Self::native_safe_prefix_at_least(messages, required)?;
3466
3467 let recent_msgs: Vec<ChatMessage> = messages.drain(to_summarize_count..).collect();
3468 let mut to_summarize = std::mem::take(messages);
3469
3470 if let Some(filter_config) = filter {
3471 let filter = self.get_filter(filter_config);
3472 to_summarize = filter.filter(to_summarize);
3473 }
3474
3475 if to_summarize.is_empty() {
3476 *messages = recent_msgs;
3477 messages.insert(0, system_msg);
3478 return Ok(());
3479 }
3480
3481 let to_summarize = Self::readable_native_messages(to_summarize)?;
3482 let conversation_text = to_summarize
3483 .iter()
3484 .map(|m| format!("{:?}: {}", m.role, m.content))
3485 .collect::<Vec<_>>()
3486 .join("\n");
3487
3488 let default_prompt = format!(
3489 "Summarize the following conversation in under {} tokens, preserving key information:\n\n{}",
3490 max_summary_tokens, conversation_text
3491 );
3492
3493 let summary_prompt = custom_prompt
3494 .map(|p| format!("{}\n\n{}", p, conversation_text))
3495 .unwrap_or(default_prompt);
3496
3497 let summarizer = if let Some(alias) = summarizer_llm {
3498 self.llm_registry
3499 .get(alias)
3500 .map_err(|e| AgentError::Config(e.to_string()))?
3501 } else {
3502 self.llm_registry
3503 .router()
3504 .or_else(|_| self.llm_registry.default())
3505 .map_err(|e| AgentError::Config(e.to_string()))?
3506 };
3507
3508 let summary_msgs = vec![ChatMessage::user(&summary_prompt)];
3509 let response = self
3510 .observe_purpose(
3511 ObservationPurpose::Summarization,
3512 summarizer.complete(&summary_msgs, None),
3513 )
3514 .await?;
3515
3516 let summary_message = ChatMessage::system(format!(
3517 "[Previous conversation summary]\n{}",
3518 response.content
3519 ));
3520
3521 *messages = vec![system_msg, summary_message];
3522 messages.extend(recent_msgs);
3523
3524 debug!(
3525 summarized_count = to_summarize_count,
3526 kept_recent = keep_recent,
3527 "Context summarized"
3528 );
3529
3530 Ok(())
3531 }
3532
3533 fn render_system_prompt(&self) -> Result<String> {
3534 let mut context = self.build_context_with_overlays();
3535
3536 let facts_text = self.format_actor_facts_for_context();
3538 if !facts_text.is_empty() {
3539 context.insert(
3540 "actor_facts".to_string(),
3541 serde_json::Value::String(facts_text),
3542 );
3543 }
3544
3545 if let Some((key, text)) = self.format_relationship_for_context() {
3546 context.insert(key, serde_json::Value::String(text));
3547 }
3548
3549 self.template_renderer
3550 .render(&self.base_system_prompt, &context)
3551 }
3552
3553 fn canonical_unique_tool_ids(&self, ids: &[String]) -> Vec<String> {
3555 let mut seen = HashSet::new();
3556 ids.iter()
3557 .filter_map(|id| self.tools.canonical_id(id))
3558 .filter(|canonical_id| seen.insert(canonical_id.clone()))
3559 .collect()
3560 }
3561
3562 fn get_top_level_tool_ids_for_scope(&self, scope_override: Option<&[String]>) -> Vec<String> {
3564 let Some(declared) = self.declared_tool_ids.as_deref() else {
3565 return Vec::new();
3566 };
3567 let mut effective = self.canonical_unique_tool_ids(declared);
3568 if let Some(scope) = scope_override {
3569 let scope: HashSet<String> =
3570 self.canonical_unique_tool_ids(scope).into_iter().collect();
3571 effective.retain(|canonical_id| scope.contains(canonical_id));
3572 }
3573 effective
3574 }
3575
3576 async fn get_available_tool_ids(&self) -> Result<Vec<String>> {
3578 Ok(self.get_available_tool_ids_snapshot().await?.tool_ids)
3579 }
3580
3581 async fn get_available_tool_ids_snapshot(&self) -> Result<AvailableToolIdsSnapshot> {
3583 let scope_override = self.runtime_control.tool_scope_override.read().clone();
3584 self.get_available_tool_ids_snapshot_for_scope(scope_override.as_deref())
3585 .await
3586 }
3587
3588 async fn get_available_tool_ids_snapshot_for_scope(
3590 &self,
3591 scope_override: Option<&[String]>,
3592 ) -> Result<AvailableToolIdsSnapshot> {
3593 let mut available = self.get_top_level_tool_ids_for_scope(scope_override);
3594 let (state_generation, state_scopes) = self
3595 .state_machine
3596 .as_ref()
3597 .map(|state_machine| {
3598 let (generation, scopes) = state_machine.current_tool_scope_snapshot();
3599 (Some(generation), scopes)
3600 })
3601 .unwrap_or((None, Vec::new()));
3602
3603 if available.is_empty() || state_scopes.is_empty() {
3604 return Ok(AvailableToolIdsSnapshot {
3605 tool_ids: available,
3606 state_generation,
3607 });
3608 }
3609
3610 let eval_ctx = self.build_evaluation_context().await?;
3611 let llm_getter = RegistryLLMGetter {
3612 registry: self.llm_registry.clone(),
3613 };
3614 let evaluator = ConditionEvaluator::new(llm_getter);
3615
3616 for state_scope in state_scopes {
3617 if state_scope.is_empty() {
3618 available.clear();
3619 break;
3620 }
3621
3622 let mut allowed = HashSet::new();
3623 for tool_ref in &state_scope {
3624 let tool_id = tool_ref.id();
3625 let Some(canonical_id) = self.tools.canonical_id(tool_id) else {
3626 continue;
3627 };
3628 let condition_matches = if let Some(condition) = tool_ref.condition() {
3629 match evaluator.evaluate(condition, &eval_ctx).await {
3630 Ok(matches) => matches,
3631 Err(error) => {
3632 warn!(tool = tool_id, error = %error, "Error evaluating tool condition");
3633 false
3634 }
3635 }
3636 } else {
3637 true
3638 };
3639 if condition_matches {
3640 allowed.insert(canonical_id);
3641 } else {
3642 debug!(tool = tool_id, "Tool condition not met, skipping");
3643 }
3644 }
3645 available.retain(|canonical_id| allowed.contains(canonical_id));
3646 if available.is_empty() {
3647 break;
3648 }
3649 }
3650
3651 Ok(AvailableToolIdsSnapshot {
3652 tool_ids: available,
3653 state_generation,
3654 })
3655 }
3656
3657 async fn build_evaluation_context(&self) -> Result<EvaluationContext> {
3658 let context = self.build_context_with_overlays();
3659 let messages = Self::readable_native_messages(self.memory.get_messages(Some(10)).await?)?;
3660 let tool_history = self.tool_call_history.read().clone();
3661
3662 let (state_name, turn_count, previous_state) = if let Some(ref sm) = self.state_machine {
3663 (Some(sm.current()), sm.turn_count(), sm.previous())
3664 } else {
3665 (None, 0, None)
3666 };
3667
3668 Ok(EvaluationContext::default()
3669 .with_context(context)
3670 .with_state(state_name, turn_count, previous_state)
3671 .with_called_tools(tool_history)
3672 .with_messages(messages))
3673 }
3674
3675 fn record_tool_call(&self, tool_id: &str, result: Value) {
3676 self.tool_call_history.write().push(ToolCallRecord {
3677 tool_id: tool_id.to_string(),
3678 result,
3679 timestamp: chrono::Utc::now(),
3680 });
3681 }
3682
3683 async fn get_effective_system_prompt_with_persona_hooks(
3684 &self,
3685 fire_persona_hooks: bool,
3686 include_tool_prompt: bool,
3687 ) -> Result<String> {
3688 let rendered_base = self.render_system_prompt()?;
3689
3690 let persona_prefix = if let Some(ref persona) = self.persona_manager {
3691 let context = self.build_context_with_overlays();
3692 if fire_persona_hooks {
3693 let render_result = persona.render_prompt(&context)?;
3694 for content in &render_result.newly_revealed {
3695 self.hooks.on_secret_revealed(content).await;
3696 }
3697 render_result.prompt
3698 } else {
3699 persona.render_prompt_preview(&context)?
3700 }
3701 } else {
3702 String::new()
3703 };
3704
3705 if let Some(ref sm) = self.state_machine
3706 && let Some(state_def) = sm.current_definition()
3707 {
3708 let state_prompt = if let Some(ref prompt) = state_def.prompt {
3709 let context = self.build_context_with_overlays();
3710 self.template_renderer.render_with_state(
3711 prompt,
3712 &context,
3713 &sm.current(),
3714 sm.previous().as_deref(),
3715 sm.turn_count(),
3716 state_def.max_turns,
3717 )?
3718 } else {
3719 String::new()
3720 };
3721
3722 let combined = match state_def.prompt_mode {
3723 PromptMode::Append => {
3724 if state_prompt.is_empty() {
3725 rendered_base
3726 } else {
3727 format!(
3728 "{}\n\n[Current State: {}]\n{}",
3729 rendered_base,
3730 sm.current(),
3731 state_prompt
3732 )
3733 }
3734 }
3735 PromptMode::Replace => {
3736 if state_prompt.is_empty() {
3737 rendered_base
3738 } else {
3739 state_prompt
3740 }
3741 }
3742 PromptMode::Prepend => {
3743 if state_prompt.is_empty() {
3744 rendered_base
3745 } else {
3746 format!("{}\n\n{}", state_prompt, rendered_base)
3747 }
3748 }
3749 };
3750
3751 let with_persona = if persona_prefix.is_empty() {
3753 combined
3754 } else {
3755 format!("{}\n\n{}", persona_prefix, combined)
3756 };
3757
3758 if include_tool_prompt {
3759 let available_tool_ids = self.get_available_tool_ids().await?;
3760 if !available_tool_ids.is_empty() {
3761 let tools_prompt = self.tools.generate_scoped_prompt_with_mode(
3762 &available_tool_ids,
3763 None,
3764 self.parallel_tools.enabled,
3765 self.runtime_config.tool_schema_prompt_mode,
3766 );
3767 if !tools_prompt.is_empty() {
3768 return Ok(format!("{}\n\n{}", with_persona, tools_prompt));
3769 }
3770 }
3771 }
3772 return Ok(with_persona);
3773 }
3774
3775 let with_persona = if persona_prefix.is_empty() {
3777 rendered_base
3778 } else {
3779 format!("{}\n\n{}", persona_prefix, rendered_base)
3780 };
3781
3782 if include_tool_prompt {
3783 let available_tool_ids = self.get_available_tool_ids().await?;
3784 let tools_prompt = self.tools.generate_scoped_prompt_with_mode(
3785 &available_tool_ids,
3786 None,
3787 self.parallel_tools.enabled,
3788 self.runtime_config.tool_schema_prompt_mode,
3789 );
3790 if !tools_prompt.is_empty() {
3791 return Ok(format!("{}\n\n{}", with_persona, tools_prompt));
3792 }
3793 }
3794 Ok(with_persona)
3795 }
3796
3797 fn get_state_llm(&self) -> Result<Arc<dyn LLMProvider>> {
3798 if let Some(ref sm) = self.state_machine
3799 && let Some(state_def) = sm.current_definition()
3800 && let Some(ref llm_alias) = state_def.llm
3801 {
3802 return self
3803 .llm_registry
3804 .get(llm_alias)
3805 .map_err(|e| AgentError::Config(e.to_string()));
3806 }
3807 self.llm_registry
3808 .default()
3809 .map_err(|e| AgentError::Config(e.to_string()))
3810 }
3811
3812 fn get_effective_reasoning_config(&self) -> ReasoningConfig {
3813 if let Some(ref sm) = self.state_machine
3814 && let Some(state_def) = sm.current_definition()
3815 && let Some(ref state_reasoning) = state_def.reasoning
3816 {
3817 return state_reasoning.clone();
3818 }
3819 self.reasoning_config.clone()
3820 }
3821
3822 fn get_effective_reflection_config(&self) -> ReflectionConfig {
3823 if let Some(ref sm) = self.state_machine
3824 && let Some(state_def) = sm.current_definition()
3825 && let Some(ref state_reflection) = state_def.reflection
3826 {
3827 return state_reflection.clone();
3828 }
3829 self.reflection_config.clone()
3830 }
3831
3832 fn get_skill_reasoning_config(&self, skill: &SkillDefinition) -> ReasoningConfig {
3833 skill
3834 .reasoning
3835 .clone()
3836 .unwrap_or_else(|| self.get_effective_reasoning_config())
3837 }
3838
3839 fn get_skill_reflection_config(&self, skill: &SkillDefinition) -> ReflectionConfig {
3840 skill
3841 .reflection
3842 .clone()
3843 .unwrap_or_else(|| self.get_effective_reflection_config())
3844 }
3845
3846 async fn build_disambiguation_context(&self) -> Result<DisambiguationContext> {
3847 let recent_messages: Vec<String> =
3848 Self::readable_native_messages(self.memory.get_messages(Some(5)).await?)?
3849 .iter()
3850 .rev()
3851 .map(|m| format!("{:?}: {}", m.role, m.content))
3852 .collect();
3853
3854 let current_state = self.current_state().map(|s| s.to_string());
3855
3856 let state_prompt: Option<String> = self
3859 .state_machine
3860 .as_ref()
3861 .and_then(|sm| sm.current_definition())
3862 .and_then(|def| def.prompt.clone());
3863
3864 let available_tools: Vec<String> = self
3865 .get_available_tool_ids()
3866 .await
3867 .unwrap_or_else(|_| self.tools.list_ids());
3868
3869 let available_skills: Vec<String> = self.skills.iter().map(|s| s.id.clone()).collect();
3870
3871 let mut user_context = self.build_context_with_overlays();
3872 user_context.remove(DISAMBIGUATION_STATE_GENERATION_KEY);
3873 if let Some(state_generation) = self
3874 .state_machine
3875 .as_ref()
3876 .map(|state_machine| state_machine.generation())
3877 {
3878 user_context.insert(
3879 DISAMBIGUATION_STATE_GENERATION_KEY.to_string(),
3880 serde_json::json!(state_generation),
3881 );
3882 }
3883
3884 let available_intents: Vec<String> = if let Some(ref sm) = self.state_machine {
3886 sm.current_definition()
3887 .map(|def| {
3888 def.transitions
3889 .iter()
3890 .filter_map(|t| t.intent.clone())
3891 .collect()
3892 })
3893 .unwrap_or_default()
3894 } else {
3895 Vec::new()
3896 };
3897
3898 Ok(DisambiguationContext::from_agent_state(
3899 recent_messages,
3900 current_state,
3901 state_prompt,
3902 available_tools,
3903 available_skills,
3904 available_intents,
3905 user_context,
3906 ))
3907 }
3908
3909 fn get_available_skills(&self) -> Vec<&SkillDefinition> {
3910 if let Some(ref sm) = self.state_machine
3911 && let Some(state_def) = sm.current_definition()
3912 {
3913 let parent_def = sm.get_parent_definition();
3914 let effective_skills = state_def.get_effective_skills(parent_def.as_ref());
3915 if !effective_skills.is_empty() {
3916 return self
3917 .skills
3918 .iter()
3919 .filter(|s| effective_skills.contains(&&s.id))
3920 .collect();
3921 }
3922 }
3923 self.skills.iter().collect()
3924 }
3925
3926 async fn build_messages(&self) -> Result<Vec<ChatMessage>> {
3927 self.build_messages_internal(true, None, true).await
3928 }
3929
3930 async fn build_messages_for_draft(&self, user_message: &str) -> Result<Vec<ChatMessage>> {
3931 self.build_messages_internal(false, Some(user_message), true)
3932 .await
3933 }
3934
3935 async fn build_messages_internal(
3936 &self,
3937 fire_persona_hooks: bool,
3938 ephemeral_user_message: Option<&str>,
3939 include_tool_prompt: bool,
3940 ) -> Result<Vec<ChatMessage>> {
3941 let system_prompt = self
3942 .get_effective_system_prompt_with_persona_hooks(fire_persona_hooks, include_tool_prompt)
3943 .await?;
3944 let mut messages = vec![ChatMessage::system(&system_prompt)];
3945
3946 let context = self.memory.get_context().await?;
3947 let history = if let Some(ref budget) = self.memory_token_budget {
3948 context.to_llm_messages_with_allocation(&budget.allocation)
3949 } else {
3950 context.to_llm_messages()
3951 };
3952 messages.extend(history);
3953 if let Some(user_message) = ephemeral_user_message {
3954 messages.push(ChatMessage::user(user_message));
3955 }
3956
3957 let total_tokens = self.estimate_total_tokens(&messages);
3958
3959 if total_tokens > self.max_context_tokens {
3960 debug!(
3961 total = total_tokens,
3962 limit = self.max_context_tokens,
3963 "Context overflow"
3964 );
3965
3966 match &self.recovery_manager.config().llm.on_context_overflow {
3967 ContextOverflowAction::Error => {
3968 return Err(AgentError::LLM(format!(
3969 "Context overflow: {} tokens > {} limit",
3970 total_tokens, self.max_context_tokens
3971 )));
3972 }
3973 ContextOverflowAction::Truncate { keep_recent } => {
3974 self.truncate_context(&mut messages, *keep_recent)?;
3975 }
3976 ContextOverflowAction::Summarize {
3977 summarizer_llm,
3978 max_summary_tokens,
3979 custom_prompt,
3980 keep_recent,
3981 filter,
3982 } => {
3983 self.summarize_context(
3984 &mut messages,
3985 summarizer_llm.as_deref(),
3986 *max_summary_tokens,
3987 custom_prompt.as_deref(),
3988 *keep_recent,
3989 filter.as_ref(),
3990 )
3991 .await?;
3992 }
3993 }
3994 }
3995
3996 self.validate_active_native_history(&messages, true)?;
3997 Ok(messages)
3998 }
3999
4000 async fn main_tool_protocol(
4001 &self,
4002 llm: &dyn LLMProvider,
4003 ephemeral_new_turn: bool,
4004 ) -> Result<MainToolProtocol> {
4005 let mut choice = llm.configured_tool_choice();
4006 if matches!(choice.as_ref(), Some(ToolChoice::None)) {
4007 return Ok(MainToolProtocol {
4008 choice,
4009 tool_ids: Vec::new(),
4010 definitions: Vec::new(),
4011 });
4012 }
4013
4014 let mut tool_ids = self.get_available_tool_ids().await?;
4015 tool_ids.sort();
4016 tool_ids.dedup();
4017 if let Some(ToolChoice::Specific(expected)) = choice.as_ref() {
4018 let canonical = self.tools.canonical_id(expected).ok_or_else(|| {
4019 AgentError::Config(format!(
4020 "specific tool choice '{expected}' is not registered"
4021 ))
4022 })?;
4023 if canonical != *expected {
4024 return Err(AgentError::Config(format!(
4025 "specific tool choice must use canonical ID '{canonical}', not '{expected}'"
4026 )));
4027 }
4028 if !tool_ids.iter().any(|tool_id| tool_id == expected) {
4029 return Err(AgentError::Config(format!(
4030 "specific tool choice '{expected}' is outside the effective tool grant"
4031 )));
4032 }
4033 }
4034 if matches!(
4035 choice.as_ref(),
4036 Some(ToolChoice::Required | ToolChoice::Specific(_))
4037 ) && tool_ids.is_empty()
4038 {
4039 return Err(AgentError::Config(
4040 "required tool choice has no tool inside the effective grant".to_string(),
4041 ));
4042 }
4043 if !ephemeral_new_turn
4044 && let Some(configured_choice) = choice.as_ref()
4045 && matches!(
4046 configured_choice,
4047 ToolChoice::Required | ToolChoice::Specific(_)
4048 )
4049 && self
4050 .tool_choice_satisfied_in_current_turn(configured_choice, &tool_ids)
4051 .await?
4052 {
4053 choice = Some(ToolChoice::Auto);
4054 }
4055 if let Some(ToolChoice::Specific(expected)) = choice.as_ref() {
4056 tool_ids.retain(|tool_id| tool_id == expected);
4057 }
4058
4059 let definitions = tool_ids
4060 .iter()
4061 .map(|tool_id| {
4062 let tool = self.tools.get(tool_id).ok_or_else(|| {
4063 AgentError::Config(format!(
4064 "effective tool '{tool_id}' disappeared before provider exposure"
4065 ))
4066 })?;
4067 Ok(LLMToolDefinition {
4068 name: tool_id.clone(),
4069 description: tool.description().to_string(),
4070 input_schema: tool.input_schema(),
4071 })
4072 })
4073 .collect::<Result<Vec<_>>>()?;
4074
4075 Ok(MainToolProtocol {
4079 choice,
4080 tool_ids,
4081 definitions,
4082 })
4083 }
4084
4085 async fn tool_choice_satisfied_in_current_turn(
4086 &self,
4087 choice: &ToolChoice,
4088 effective_tool_ids: &[String],
4089 ) -> Result<bool> {
4090 let messages = self.memory.get_messages(None).await?;
4091 let mut saw_tool_result = false;
4092 for message in messages.iter().rev() {
4093 match message.role {
4094 ai_agents_core::Role::Tool | ai_agents_core::Role::Function => {
4095 saw_tool_result = true;
4096 }
4097 ai_agents_core::Role::Assistant if saw_tool_result => {
4098 let Some(calls) = self.parse_tool_calls(&message.content)? else {
4099 continue;
4100 };
4101 let calls_are_effective = !calls.is_empty()
4102 && calls.iter().all(|call| {
4103 self.tools
4104 .canonical_id(&call.name)
4105 .is_some_and(|canonical| effective_tool_ids.contains(&canonical))
4106 });
4107 return Ok(calls_are_effective
4108 && match choice {
4109 ToolChoice::Required => true,
4110 ToolChoice::Specific(expected) => calls.iter().all(|call| {
4111 self.tools.canonical_id(&call.name).as_deref()
4112 == Some(expected.as_str())
4113 }),
4114 _ => false,
4115 });
4116 }
4117 ai_agents_core::Role::User => return Ok(false),
4118 _ => {}
4119 }
4120 }
4121 Ok(false)
4122 }
4123
4124 fn provider_can_use_native_tools(
4125 &self,
4126 llm: &dyn LLMProvider,
4127 protocol: &MainToolProtocol,
4128 ) -> bool {
4129 let Some(choice) = protocol.choice.as_ref() else {
4130 return false;
4131 };
4132 if matches!(choice, ToolChoice::None) || protocol.definitions.is_empty() {
4133 return false;
4134 }
4135 llm.supports_tool_choice(choice)
4136 && protocol.definitions.iter().all(|definition| {
4137 !definition.name.is_empty()
4138 && definition.name.len() <= 64
4139 && definition
4140 .name
4141 .bytes()
4142 .all(|byte| byte.is_ascii_alphanumeric() || matches!(byte, b'_' | b'-'))
4143 })
4144 }
4145
4146 fn prompt_messages_for_tool_protocol(
4147 &self,
4148 messages: &[ChatMessage],
4149 protocol: &MainToolProtocol,
4150 corrective: bool,
4151 ) -> Vec<ChatMessage> {
4152 let mut messages = messages.to_vec();
4153 let Some(choice) = protocol.choice.as_ref() else {
4154 return messages;
4155 };
4156 if matches!(choice, ToolChoice::None) || protocol.tool_ids.is_empty() {
4157 return messages;
4158 }
4159
4160 let mut tool_prompt = self.tools.generate_scoped_prompt_with_mode(
4161 &protocol.tool_ids,
4162 None,
4163 self.parallel_tools.enabled,
4164 self.runtime_config.tool_schema_prompt_mode,
4165 );
4166 match choice {
4167 ToolChoice::Required => tool_prompt.push_str(
4168 "\n\nYou must call at least one listed tool before giving a final answer.",
4169 ),
4170 ToolChoice::Specific(tool_id) => tool_prompt.push_str(&format!(
4171 "\n\nYou must call the '{tool_id}' tool before giving a final answer."
4172 )),
4173 ToolChoice::Auto => {}
4174 ToolChoice::None => return messages,
4175 _ => return messages,
4176 }
4177 if let Some(system) = messages
4178 .iter_mut()
4179 .find(|message| message.role == ai_agents_core::Role::System)
4180 {
4181 system.content.push_str("\n\n");
4182 system.content.push_str(&tool_prompt);
4183 } else {
4184 messages.insert(0, ChatMessage::system(tool_prompt));
4185 }
4186 if corrective {
4187 let instruction = match choice {
4188 ToolChoice::Required => {
4189 "Your previous response did not call a required tool. Call at least one listed tool now and return only the JSON tool call."
4190 }
4191 ToolChoice::Specific(tool_id) => {
4192 messages.push(ChatMessage::user(format!(
4193 "Your previous response did not call the required '{tool_id}' tool. Call it now and return only the JSON tool call."
4194 )));
4195 return messages;
4196 }
4197 _ => return messages,
4198 };
4199 messages.push(ChatMessage::user(instruction));
4200 }
4201 messages
4202 }
4203
4204 async fn invoke_main_provider(
4205 &self,
4206 llm: Arc<dyn LLMProvider>,
4207 messages: &[ChatMessage],
4208 protocol: &MainToolProtocol,
4209 corrective: bool,
4210 ) -> std::result::Result<MainProviderResponse, LLMError> {
4211 let use_native = self.provider_can_use_native_tools(llm.as_ref(), protocol);
4212 let response = if use_native {
4213 let request = LLMToolRequest {
4214 tools: protocol.definitions.clone(),
4215 choice: protocol
4216 .choice
4217 .clone()
4218 .expect("native tool requests require an explicit choice"),
4219 };
4220 self.observe_purpose(
4221 ObservationPurpose::MainResponse,
4222 llm.complete_with_tools(messages, None, &request),
4223 )
4224 .await?
4225 } else {
4226 let prompt_messages =
4227 self.prompt_messages_for_tool_protocol(messages, protocol, corrective);
4228 self.observe_purpose(
4229 ObservationPurpose::MainResponse,
4230 llm.complete(&prompt_messages, None),
4231 )
4232 .await?
4233 };
4234 Ok(MainProviderResponse {
4235 response,
4236 used_native_tools: use_native,
4237 })
4238 }
4239
4240 async fn complete_main_attempt_with_recovery(
4241 &self,
4242 llm: Arc<dyn LLMProvider>,
4243 messages: &[ChatMessage],
4244 protocol: &MainToolProtocol,
4245 corrective: bool,
4246 ) -> Result<MainProviderResponse> {
4247 let primary_result = self
4249 .recovery_manager
4250 .with_llm_retry(
4251 "llm_call",
4252 None,
4253 || {
4254 let llm = Arc::clone(&llm);
4255 async move {
4256 self.invoke_main_provider(llm, messages, protocol, corrective)
4257 .await
4258 }
4259 },
4260 |error| llm.is_terminal_error(error),
4261 )
4262 .await;
4263
4264 match primary_result {
4265 Ok(response) => Ok(response),
4266 Err(ai_agents_recovery::RetryFailure::Terminal { error, .. }) => {
4267 Err(AgentError::LLM(error.to_string()))
4268 }
4269 Err(failure) => {
4270 let primary_error = AgentError::LLM(failure.into_error().to_string());
4271 match &self.recovery_manager.config().llm.on_failure {
4272 LLMFailureAction::FallbackLlm { fallback_llm } => {
4273 let fallback = self.llm_registry.get(fallback_llm).map_err(|error| {
4274 AgentError::Config(format!(
4275 "Fallback LLM '{fallback_llm}' not found: {error}"
4276 ))
4277 })?;
4278 self.invoke_main_provider(fallback, messages, protocol, corrective)
4279 .await
4280 .map_err(|error| AgentError::LLM(error.to_string()))
4281 }
4282 LLMFailureAction::FallbackResponse { message } => {
4283 if matches!(
4284 protocol.choice.as_ref(),
4285 Some(ToolChoice::Required | ToolChoice::Specific(_))
4286 ) {
4287 Err(AgentError::LLM(format!(
4288 "Required tool selection failed and cannot be satisfied by a static fallback response: {primary_error}"
4289 )))
4290 } else {
4291 Ok(MainProviderResponse {
4292 response: LLMResponse::new(message.clone(), FinishReason::Stop),
4293 used_native_tools: false,
4294 })
4295 }
4296 }
4297 LLMFailureAction::Error => Err(primary_error),
4298 }
4299 }
4300 }
4301 }
4302
4303 fn normalize_main_provider_response(
4304 &self,
4305 mut response: LLMResponse,
4306 protocol: &MainToolProtocol,
4307 ) -> Result<(LLMResponse, bool)> {
4308 let provider_state = response
4309 .take_provider_state()
4310 .map_err(|error| AgentError::LLM(error.to_string()))?;
4311 let native_calls = response
4312 .tool_calls()
4313 .map_err(|error| AgentError::LLM(error.to_string()))?;
4314 let calls = match native_calls {
4315 Some(calls) => {
4316 response.content = encode_native_tool_call_markers(&calls, provider_state.as_ref())
4317 .map_err(|error| AgentError::LLM(error.to_string()))?;
4318 Some(calls)
4319 }
4320 None if provider_state.is_some() => {
4321 return Err(AgentError::LLM(
4322 "Provider returned replay state without native tool calls".to_string(),
4323 ));
4324 }
4325 None if !matches!(protocol.choice.as_ref(), Some(ToolChoice::None)) => {
4326 self.parse_tool_calls(response.content.trim())?
4327 }
4328 None => None,
4329 };
4330
4331 if protocol.choice.is_some()
4332 && let Some(calls) = calls.as_ref()
4333 && calls.iter().any(|call| {
4334 self.tools
4335 .canonical_id(&call.name)
4336 .is_none_or(|canonical| !protocol.tool_ids.contains(&canonical))
4337 })
4338 {
4339 return Err(AgentError::LLM(
4340 "Provider returned a tool call outside the effective grant".to_string(),
4341 ));
4342 }
4343
4344 let compliant = match protocol.choice.as_ref() {
4345 Some(ToolChoice::Required) => calls.as_ref().is_some_and(|calls| !calls.is_empty()),
4346 Some(ToolChoice::Specific(expected)) => calls.as_ref().is_some_and(|calls| {
4347 !calls.is_empty()
4348 && calls.iter().all(|call| {
4349 self.tools.canonical_id(&call.name).as_deref() == Some(expected.as_str())
4350 })
4351 }),
4352 _ => true,
4353 };
4354 Ok((response, compliant))
4355 }
4356
4357 async fn complete_main_llm_with_recovery(
4358 &self,
4359 llm: Arc<dyn LLMProvider>,
4360 messages: &[ChatMessage],
4361 protocol: &MainToolProtocol,
4362 ) -> Result<LLMResponse> {
4363 let first = self
4364 .complete_main_attempt_with_recovery(Arc::clone(&llm), messages, protocol, false)
4365 .await?;
4366 let (response, compliant) =
4367 self.normalize_main_provider_response(first.response, protocol)?;
4368 if compliant {
4369 return Ok(response);
4370 }
4371 if first.used_native_tools {
4372 return Err(AgentError::LLM(
4373 "Provider returned no compliant native call for required tool choice".to_string(),
4374 ));
4375 }
4376
4377 let corrected = self
4378 .complete_main_attempt_with_recovery(llm, messages, protocol, true)
4379 .await?;
4380 let (response, compliant) =
4381 self.normalize_main_provider_response(corrected.response, protocol)?;
4382 if compliant {
4383 return Ok(response);
4384 }
4385 Err(AgentError::LLM(
4386 "Provider returned no compliant tool call after one corrective retry".to_string(),
4387 ))
4388 }
4389
4390 async fn open_main_stream_with_recovery(
4397 &self,
4398 llm: Arc<dyn LLMProvider>,
4399 messages: &[ChatMessage],
4400 protocol: &MainToolProtocol,
4401 ) -> Result<MainStreamSource> {
4402 debug_assert!(
4403 protocol.choice.is_none(),
4404 "streaming raw path must not run with explicit tool choice"
4405 );
4406 let primary = self
4407 .recovery_manager
4408 .with_llm_retry(
4409 "llm_stream_open",
4410 None,
4411 || {
4412 let llm = Arc::clone(&llm);
4413 async move {
4414 self.observe_purpose(
4415 ObservationPurpose::MainResponse,
4416 llm.complete_stream(messages, None),
4417 )
4418 .await
4419 }
4420 },
4421 |error| llm.is_terminal_error(error),
4422 )
4423 .await;
4424
4425 match primary {
4426 Ok(stream) => Ok(MainStreamSource::Stream(stream)),
4427 Err(ai_agents_recovery::RetryFailure::Terminal { error, .. }) => {
4428 Err(AgentError::LLM(error.to_string()))
4429 }
4430 Err(failure) => {
4431 let primary_error = AgentError::LLM(failure.into_error().to_string());
4432 match &self.recovery_manager.config().llm.on_failure {
4433 LLMFailureAction::FallbackLlm { fallback_llm } => {
4434 let fallback = self.llm_registry.get(fallback_llm).map_err(|error| {
4435 AgentError::Config(format!(
4436 "Fallback LLM '{fallback_llm}' not found: {error}"
4437 ))
4438 })?;
4439 if fallback.supports(LLMFeature::Streaming) {
4440 let stream = self
4441 .observe_purpose(
4442 ObservationPurpose::MainResponse,
4443 fallback.complete_stream(messages, None),
4444 )
4445 .await
4446 .map_err(|error| AgentError::LLM(error.to_string()))?;
4447 Ok(MainStreamSource::Stream(stream))
4448 } else {
4449 let response = self
4450 .observe_purpose(
4451 ObservationPurpose::MainResponse,
4452 fallback.complete(messages, None),
4453 )
4454 .await
4455 .map_err(|error| AgentError::LLM(error.to_string()))?;
4456 Ok(MainStreamSource::StaticResponse(response.content))
4457 }
4458 }
4459 LLMFailureAction::FallbackResponse { message } => {
4460 Ok(MainStreamSource::StaticResponse(message.clone()))
4461 }
4462 LLMFailureAction::Error => Err(primary_error),
4463 }
4464 }
4465 }
4466 }
4467
4468 fn main_stream_must_buffer(
4477 &self,
4478 reasoning_mode: &ReasoningMode,
4479 protocol: &MainToolProtocol,
4480 ) -> bool {
4481 protocol.choice.is_some()
4482 || self.get_effective_reflection_config().requires_evaluation()
4483 || matches!(reasoning_mode, ReasoningMode::CoT | ReasoningMode::React)
4484 }
4485
4486 fn is_native_tool_call_content(content: &str) -> Result<bool> {
4488 decode_native_tool_call_markers(content)
4489 .map(|batch| batch.is_some())
4490 .map_err(|error| AgentError::LLM(error.to_string()))
4491 }
4492
4493 fn tool_result_message(
4495 tool_call: &ToolCall,
4496 output: &str,
4497 native_tool_call: bool,
4498 ) -> Result<ChatMessage> {
4499 if !native_tool_call {
4500 return Ok(ChatMessage::function(&tool_call.name, output));
4501 }
4502 let output = serde_json::from_str::<serde_json::Value>(output)
4503 .unwrap_or_else(|_| serde_json::Value::String(output.to_string()));
4504 let content = encode_native_tool_result_marker(tool_call, output)
4505 .map_err(|error| AgentError::LLM(error.to_string()))?;
4506 Ok(ChatMessage::function(&tool_call.name, content))
4507 }
4508
4509 fn remember_active_native_exchange(&self, content: &str) -> Result<()> {
4511 let Some(batch) = decode_native_tool_call_markers(content)
4512 .map_err(|error| AgentError::LLM(error.to_string()))?
4513 else {
4514 return Ok(());
4515 };
4516 let Some(state) = batch.provider_state() else {
4517 return Ok(());
4518 };
4519 let expected = ActiveNativeExchange {
4520 exchange_id: state.exchange_id().to_string(),
4521 call_ids: batch.calls().iter().map(|call| call.id.clone()).collect(),
4522 };
4523 let mut active = self.active_native_exchanges.write();
4524 if let Some(existing) = active
4525 .iter()
4526 .find(|existing| existing.exchange_id == expected.exchange_id)
4527 {
4528 if existing.call_ids != expected.call_ids {
4529 return Err(AgentError::LLM(format!(
4530 "Active native exchange '{}' changed its call identities",
4531 expected.exchange_id
4532 )));
4533 }
4534 } else {
4535 active.push(expected);
4536 }
4537 Ok(())
4538 }
4539
4540 fn validate_active_native_history(
4542 &self,
4543 messages: &[ChatMessage],
4544 require_complete: bool,
4545 ) -> Result<()> {
4546 let expected = self.active_native_exchanges.read().clone();
4547 if expected.is_empty() {
4548 return Ok(());
4549 }
4550 let inspection =
4551 inspect_native_history(messages).map_err(|error| AgentError::LLM(error.to_string()))?;
4552 let expected_count = expected.len();
4553 for (index, expected) in expected.iter().enumerate() {
4554 let Some(exchange) = inspection
4555 .exchanges()
4556 .iter()
4557 .find(|exchange| exchange.state().exchange_id() == expected.exchange_id)
4558 else {
4559 return Err(AgentError::LLM(format!(
4560 "Active native exchange '{}' was removed before provider continuation",
4561 expected.exchange_id
4562 )));
4563 };
4564 let must_be_complete = require_complete || index + 1 < expected_count;
4565 if exchange.call_ids() != expected.call_ids
4566 || (must_be_complete && !exchange.is_complete())
4567 {
4568 return Err(AgentError::LLM(format!(
4569 "Active native exchange '{}' is incomplete before provider continuation",
4570 expected.exchange_id
4571 )));
4572 }
4573 }
4574 Ok(())
4575 }
4576
4577 async fn remember_committed_native_exchange(&self, content: &str) -> Result<()> {
4579 self.remember_active_native_exchange(content)?;
4580 if !self.active_native_exchanges.read().is_empty() {
4581 let messages = self.memory.get_messages(None).await?;
4582 self.validate_active_native_history(&messages, false)?;
4583 }
4584 Ok(())
4585 }
4586
4587 fn readable_native_messages(mut messages: Vec<ChatMessage>) -> Result<Vec<ChatMessage>> {
4589 for message in &mut messages {
4590 if matches!(
4591 message.role,
4592 ai_agents_core::Role::Assistant
4593 | ai_agents_core::Role::Tool
4594 | ai_agents_core::Role::Function
4595 ) {
4596 message.content = native_readable_projection(&message.content)
4597 .map_err(|error| AgentError::LLM(error.to_string()))?;
4598 }
4599 }
4600 Ok(messages)
4601 }
4602
4603 fn parse_main_tool_calls(
4605 &self,
4606 content: &str,
4607 protocol: &MainToolProtocol,
4608 ) -> Result<Option<Vec<ToolCall>>> {
4609 if matches!(protocol.choice.as_ref(), Some(ToolChoice::None)) {
4610 Ok(None)
4611 } else {
4612 self.parse_tool_calls(content)
4613 }
4614 }
4615
4616 fn parse_tool_calls(&self, content: &str) -> Result<Option<Vec<ToolCall>>> {
4618 if let Some(batch) = decode_native_tool_call_markers(content)
4619 .map_err(|error| AgentError::LLM(error.to_string()))?
4620 {
4621 return Ok(Some(batch.into_parts().0));
4622 }
4623 if let Ok(parsed) = serde_json::from_str::<serde_json::Value>(content) {
4625 if let Some(arr) = parsed.as_array() {
4627 let calls: Vec<ToolCall> = arr
4628 .iter()
4629 .filter_map(|v| self.extract_tool_call_from_value(v))
4630 .collect();
4631 if !calls.is_empty() {
4632 return Ok(Some(calls));
4633 }
4634 }
4635 if let Some(tool_call) = self.extract_tool_call_from_value(&parsed) {
4637 return Ok(Some(vec![tool_call]));
4638 }
4639 }
4640
4641 if let Some(json_str) = self.extract_json_from_content(content)
4643 && let Ok(parsed) = serde_json::from_str::<serde_json::Value>(&json_str)
4644 {
4645 if let Some(arr) = parsed.as_array() {
4647 let calls: Vec<ToolCall> = arr
4648 .iter()
4649 .filter_map(|v| self.extract_tool_call_from_value(v))
4650 .collect();
4651 if !calls.is_empty() {
4652 return Ok(Some(calls));
4653 }
4654 }
4655 if let Some(tool_call) = self.extract_tool_call_from_value(&parsed) {
4657 return Ok(Some(vec![tool_call]));
4658 }
4659 }
4660
4661 Ok(None)
4662 }
4663
4664 fn extract_tool_call_from_value(&self, parsed: &serde_json::Value) -> Option<ToolCall> {
4665 if let Some(tool_name) = parsed.get("tool").and_then(|v| v.as_str()) {
4666 let arguments = parsed
4667 .get("arguments")
4668 .cloned()
4669 .unwrap_or(serde_json::json!({}));
4670 return Some(ToolCall {
4671 id: parsed
4672 .get("id")
4673 .and_then(|value| value.as_str())
4674 .filter(|id| !id.is_empty())
4675 .map(str::to_string)
4676 .unwrap_or_else(|| uuid::Uuid::new_v4().to_string()),
4677 name: tool_name.to_string(),
4678 arguments,
4679 });
4680 }
4681 None
4682 }
4683
4684 fn extract_json_from_content(&self, content: &str) -> Option<String> {
4686 if let Some(result) = self.extract_json_array_from_content(content) {
4688 return Some(result);
4689 }
4690 self.extract_json_object_from_content(content)
4691 }
4692
4693 fn extract_json_array_from_content(&self, content: &str) -> Option<String> {
4695 let start = content.find('[')?;
4696 let content_from_start = &content[start..];
4697
4698 let mut depth = 0;
4699 let mut end = 0;
4700 for (i, ch) in content_from_start.char_indices() {
4701 match ch {
4702 '[' => depth += 1,
4703 ']' => {
4704 depth -= 1;
4705 if depth == 0 {
4706 end = i + 1;
4707 break;
4708 }
4709 }
4710 _ => {}
4711 }
4712 }
4713
4714 if end > 0 {
4715 let json_str = &content_from_start[..end];
4716 if json_str.contains("\"tool\"") {
4718 return Some(json_str.to_string());
4719 }
4720 }
4721
4722 None
4723 }
4724
4725 fn extract_json_object_from_content(&self, content: &str) -> Option<String> {
4727 let start = content.find('{')?;
4728 let content_from_start = &content[start..];
4729
4730 let mut depth = 0;
4732 let mut end = 0;
4733 for (i, ch) in content_from_start.char_indices() {
4734 match ch {
4735 '{' => depth += 1,
4736 '}' => {
4737 depth -= 1;
4738 if depth == 0 {
4739 end = i + 1;
4740 break;
4741 }
4742 }
4743 _ => {}
4744 }
4745 }
4746
4747 if end > 0 {
4748 let json_str = &content_from_start[..end];
4749 if json_str.contains("\"tool\"") {
4751 return Some(json_str.to_string());
4752 }
4753 }
4754
4755 None
4756 }
4757
4758 #[allow(clippy::too_many_arguments)]
4762 fn record_from_parts(
4763 &self,
4764 request: &ToolExecutionRequest,
4765 canonical_id: String,
4766 executed_arguments: Value,
4767 started_at: chrono::DateTime<chrono::Utc>,
4768 start: Instant,
4769 executed: bool,
4770 success: bool,
4771 output: String,
4772 metadata: HashMap<String, Value>,
4773 policy: ToolPolicyDecisionRecord,
4774 approval: Option<ToolApprovalRecord>,
4775 timed_out: bool,
4776 output_truncated: bool,
4777 ) -> ToolExecutionRecord {
4778 let versions = ToolDecisionVersions {
4779 policy: self.active_tool_security().policy_version(),
4780 registry: self.tools.version(),
4781 runtime_control: self.runtime_control.version.load(Ordering::SeqCst),
4782 state: self
4783 .state_machine
4784 .as_ref()
4785 .map(|state_machine| state_machine.generation()),
4786 };
4787 self.record_from_parts_at(
4788 request,
4789 canonical_id,
4790 executed_arguments,
4791 started_at,
4792 start,
4793 executed,
4794 success,
4795 output,
4796 metadata,
4797 policy,
4798 approval,
4799 timed_out,
4800 output_truncated,
4801 versions,
4802 )
4803 }
4804
4805 #[allow(clippy::too_many_arguments)]
4807 fn record_from_parts_at(
4808 &self,
4809 request: &ToolExecutionRequest,
4810 canonical_id: String,
4811 executed_arguments: Value,
4812 started_at: chrono::DateTime<chrono::Utc>,
4813 start: Instant,
4814 executed: bool,
4815 success: bool,
4816 output: String,
4817 metadata: HashMap<String, Value>,
4818 policy: ToolPolicyDecisionRecord,
4819 approval: Option<ToolApprovalRecord>,
4820 timed_out: bool,
4821 output_truncated: bool,
4822 versions: ToolDecisionVersions,
4823 ) -> ToolExecutionRecord {
4824 ToolExecutionRecord {
4825 call_id: request.call_id.clone(),
4826 requested_name: request.requested_name.clone(),
4827 canonical_id,
4828 source: request.source.clone(),
4829 arguments: request.arguments.clone(),
4830 executed_arguments,
4831 policy_version: versions.policy,
4832 registry_version: versions.registry,
4833 runtime_config_version: versions.runtime_control,
4834 executed,
4835 success,
4836 output,
4837 metadata,
4838 policy,
4839 approval,
4840 started_at,
4841 duration_ms: start.elapsed().as_millis() as u64,
4842 timed_out,
4843 cancelled: false,
4844 cancellation_reason: None,
4845 output_truncated,
4846 }
4847 }
4848
4849 async fn finish_tool_record(&self, record: &ToolExecutionRecord) {
4851 let result = ToolResult {
4852 success: record.success,
4853 output: record.model_output_string(),
4854 metadata: if record.metadata.is_empty() {
4855 None
4856 } else {
4857 Some(record.metadata.clone())
4858 },
4859 };
4860 self.hooks
4861 .on_tool_complete(&record.canonical_id, &result, record.duration_ms)
4862 .await;
4863 self.hooks.on_tool_execution_record(record).await;
4864 self.record_tool_call(&record.canonical_id, record.model_output_value());
4865 if !record.success {
4866 self.hooks
4867 .on_error(&AgentError::Tool(record.output.clone()))
4868 .await;
4869 }
4870 }
4871
4872 async fn finish_tool_record_after_resource_guards(
4874 &self,
4875 resource_guards: ToolResourceGuards,
4876 record: &ToolExecutionRecord,
4877 ) {
4878 drop(resource_guards);
4879 self.finish_tool_record(record).await;
4880 }
4881
4882 fn validated_tool_timeout(timeout_ms: u64) -> Result<ValidatedToolTimeout> {
4886 if timeout_ms > MAX_TOOL_TIMEOUT_MS {
4887 return Err(AgentError::Config(format!(
4888 "effective tool timeout_ms must be no greater than {MAX_TOOL_TIMEOUT_MS} milliseconds"
4889 )));
4890 }
4891 let timer = Duration::from_millis(timeout_ms);
4892 let deadline_delta = chrono::Duration::from_std(timer).map_err(|_| {
4893 AgentError::Config(format!(
4894 "effective tool timeout_ms cannot be represented as a UTC deadline: {timeout_ms}"
4895 ))
4896 })?;
4897 Ok(ValidatedToolTimeout {
4898 timer,
4899 deadline_delta,
4900 })
4901 }
4902
4903 fn effective_tool_limits(
4907 security_engine: &ToolSecurityEngine,
4908 canonical_id: &str,
4909 safety: &ToolSafetyMetadata,
4910 classification: &ToolCallClassification,
4911 recovery_timeout_ms: Option<u64>,
4912 ) -> Result<(ToolExecutionLimits, ValidatedToolTimeout)> {
4913 if let Some(timeout_ms) = classification.timeout_ms {
4914 Self::validated_tool_timeout(timeout_ms)?;
4915 }
4916 if let Some(timeout_ms) = recovery_timeout_ms {
4917 Self::validated_tool_timeout(timeout_ms)?;
4918 }
4919
4920 let mut limits = security_engine.effective_limits(canonical_id, safety, classification);
4921 if let Some(recovery_timeout_ms) = recovery_timeout_ms {
4922 limits.timeout_ms = Some(limits.timeout_ms.map_or(recovery_timeout_ms, |timeout_ms| {
4923 timeout_ms.min(recovery_timeout_ms)
4924 }));
4925 }
4926 let timeout_ms = limits
4927 .timeout_ms
4928 .unwrap_or_else(|| security_engine.get_tool_timeout(canonical_id));
4929 let timeout = Self::validated_tool_timeout(timeout_ms)?;
4930 Ok((limits, timeout))
4931 }
4932
4933 async fn execute_resolved_tool_once(
4935 &self,
4936 tool: Arc<dyn ai_agents_core::Tool>,
4937 args: Value,
4938 mut ctx: ToolExecutionContext,
4939 timeout: ValidatedToolTimeout,
4940 ) -> Result<(ToolResult, bool, bool, bool)> {
4941 if self.runtime_control.emergency_deny.load(Ordering::SeqCst) {
4942 return Ok((
4943 ToolResult::error("Tool execution cancelled by runtime control"),
4944 false,
4945 true,
4946 false,
4947 ));
4948 }
4949 ctx.deadline = Some(
4954 chrono::Utc::now()
4955 .checked_add_signed(timeout.deadline_delta)
4956 .ok_or_else(|| {
4957 AgentError::Config(
4958 "effective tool timeout_ms exceeds the current UTC deadline range"
4959 .to_string(),
4960 )
4961 })?,
4962 );
4963 let invoked = Arc::new(AtomicBool::new(false));
4967 let invoked_by_future = Arc::clone(&invoked);
4968 let actor_context = current_turn_actor_context();
4969 let future = async move {
4970 invoked_by_future.store(true, Ordering::SeqCst);
4971 if let Some(actor_context) = actor_context {
4972 scope_actor_context(actor_context, tool.execute(args, ctx)).await
4973 } else {
4974 tool.execute(args, ctx).await
4975 }
4976 };
4977 tokio::pin!(future);
4978 let timer = tokio::time::sleep(timeout.timer);
4979 tokio::pin!(timer);
4980 let mut cancel_tick = tokio::time::interval(std::time::Duration::from_millis(50));
4981
4982 loop {
4983 tokio::select! {
4984 result = &mut future => return Ok((result, false, false, true)),
4985 _ = &mut timer => {
4986 return Ok((
4987 ToolResult::error("Tool execution timed out"),
4988 true,
4989 false,
4990 invoked.load(Ordering::SeqCst),
4991 ));
4992 }
4993 _ = cancel_tick.tick() => {
4994 if self.runtime_control.emergency_deny.load(Ordering::SeqCst) {
4995 return Ok((
4996 ToolResult::error("Tool execution cancelled by runtime control"),
4997 false,
4998 true,
4999 invoked.load(Ordering::SeqCst),
5000 ));
5001 }
5002 }
5003 }
5004 }
5005 }
5006
5007 fn truncate_tool_output(output: String, max_chars: Option<usize>) -> (String, bool) {
5009 let Some(max_chars) = max_chars else {
5010 return (output, false);
5011 };
5012 let mut chars = output.chars();
5013 let truncated: String = chars.by_ref().take(max_chars).collect();
5014 if chars.next().is_some() {
5015 (truncated, true)
5016 } else {
5017 (output, false)
5018 }
5019 }
5020
5021 async fn acquire_tool_resource_locks(&self, keys: &[String]) -> Option<ToolResourceGuards> {
5023 let locks = {
5024 let mut table = self.resource_locks.write();
5025 table.retain(|_, lock| lock.strong_count() > 0);
5026 keys.iter()
5027 .map(|key| {
5028 if let Some(lock) = table.get(key).and_then(Weak::upgrade) {
5029 lock
5030 } else {
5031 let lock = Arc::new(tokio::sync::Mutex::new(()));
5032 table.insert(key.clone(), Arc::downgrade(&lock));
5033 lock
5034 }
5035 })
5036 .collect::<Vec<_>>()
5037 };
5038 let mut resource_guards = ToolResourceGuards {
5039 guards: Vec::with_capacity(locks.len()),
5040 locks: Arc::clone(&self.resource_locks),
5041 };
5042 let mut locks = locks.into_iter();
5043 while let Some(lock) = locks.next() {
5044 let mut lock = Box::pin(lock.lock_owned());
5045 loop {
5046 tokio::select! {
5047 guard = &mut lock => {
5048 resource_guards.guards.push(guard);
5049 break;
5050 }
5051 _ = tokio::time::sleep(std::time::Duration::from_millis(10)) => {
5052 if self.runtime_control.emergency_deny.load(Ordering::SeqCst) {
5053 drop(lock);
5054 drop(locks);
5055 drop(resource_guards);
5056 return None;
5057 }
5058 }
5059 }
5060 }
5061 }
5062 Some(resource_guards)
5063 }
5064
5065 async fn run_tool_with_retries(
5069 &self,
5070 canonical_id: &str,
5071 tool: Arc<dyn ai_agents_core::Tool>,
5072 args: Value,
5073 ctx: ToolExecutionContext,
5074 timeout: ValidatedToolTimeout,
5075 max_retries: u32,
5076 ) -> Result<(ToolResult, bool, bool, bool)> {
5077 let max_retries = if ctx.classification.safely_retryable {
5078 max_retries
5079 } else {
5080 0
5081 };
5082 let mut attempts = 0;
5083 let mut invoked = false;
5084 loop {
5085 let (result, timed_out, cancelled, attempt_invoked) = self
5086 .execute_resolved_tool_once(tool.clone(), args.clone(), ctx.clone(), timeout)
5087 .await?;
5088 invoked |= attempt_invoked;
5089 if result.success || timed_out || cancelled || attempts >= max_retries {
5090 return Ok((result, timed_out, cancelled, invoked));
5091 }
5092 attempts += 1;
5093 warn!(tool = %canonical_id, attempt = attempts, error = %result.output, "Retrying failed tool call");
5094 }
5095 }
5096
5097 fn host_tool_unavailability(&self, canonical_id: &str) -> Option<(&'static str, &'static str)> {
5099 match canonical_id {
5100 "command" if !self.tools.command_runner_available() => Some((
5101 "Command runner is unavailable",
5102 "command runner is unavailable",
5103 )),
5104 "diagnostics" if !self.tools.diagnostics_available() => Some((
5105 "Diagnostics provider is unavailable",
5106 "diagnostics provider is unavailable",
5107 )),
5108 "web_search" if !self.tools.web_search_available() => Some((
5109 "Web search provider is unavailable",
5110 "web search provider is unavailable",
5111 )),
5112 _ => None,
5113 }
5114 }
5115
5116 fn execute_tool_record(
5118 &self,
5119 request: ToolExecutionRequest,
5120 ) -> Pin<Box<dyn Future<Output = Result<ToolExecutionRecord>> + Send + '_>> {
5121 Box::pin(self.execute_tool_record_inner(request, ToolFallbackState::default()))
5122 }
5123
5124 async fn execute_tool_record_inner(
5128 &self,
5129 request: ToolExecutionRequest,
5130 fallback_state: ToolFallbackState,
5131 ) -> Result<ToolExecutionRecord> {
5132 let started_at = chrono::Utc::now();
5133 let start = Instant::now();
5134 info!(tool = %request.requested_name, args = %request.arguments, "Executing tool");
5135
5136 if self.runtime_control.emergency_deny.load(Ordering::SeqCst) {
5137 let record = self.record_from_parts(
5138 &request,
5139 request.requested_name.clone(),
5140 request.arguments.clone(),
5141 started_at,
5142 start,
5143 false,
5144 false,
5145 "Tool execution is disabled by runtime control".to_string(),
5146 HashMap::new(),
5147 ToolPolicyDecisionRecord::deny("runtime emergency deny is enabled"),
5148 None,
5149 false,
5150 false,
5151 );
5152 self.finish_tool_record(&record).await;
5153 return Ok(record);
5154 }
5155
5156 let Some(resolved) = self.tools.resolve(&request.requested_name) else {
5157 let record = self.record_from_parts(
5158 &request,
5159 request.requested_name.clone(),
5160 request.arguments.clone(),
5161 started_at,
5162 start,
5163 false,
5164 false,
5165 format!("Tool '{}' is unavailable", request.requested_name),
5166 HashMap::new(),
5167 ToolPolicyDecisionRecord::unavailable(format!(
5168 "Tool '{}' is not registered",
5169 request.requested_name
5170 )),
5171 None,
5172 false,
5173 false,
5174 );
5175 self.finish_tool_record(&record).await;
5176 return Ok(record);
5177 };
5178
5179 let canonical_id = resolved.identity.canonical_id.clone();
5180
5181 let initial_scope_snapshot = self.get_available_tool_ids_snapshot().await?;
5182 if !initial_scope_snapshot
5183 .tool_ids
5184 .iter()
5185 .any(|id| id == &canonical_id)
5186 {
5187 let record = self.record_from_parts(
5188 &request,
5189 canonical_id.clone(),
5190 request.arguments.clone(),
5191 started_at,
5192 start,
5193 false,
5194 false,
5195 format!(
5196 "Tool '{}' is not available in the current scope",
5197 canonical_id
5198 ),
5199 HashMap::new(),
5200 ToolPolicyDecisionRecord::deny(format!(
5201 "Tool '{}' is not granted by the current top-level and state tool scope",
5202 canonical_id
5203 )),
5204 None,
5205 false,
5206 false,
5207 );
5208 self.finish_tool_record(&record).await;
5209 return Ok(record);
5210 }
5211
5212 let approval_control_snapshot = self.runtime_safety_snapshot();
5213 let security_engine = approval_control_snapshot.tool_security.clone();
5214 if let Some(reason) = fallback_state.rejection_reason(&canonical_id) {
5215 let mut metadata = HashMap::new();
5216 metadata.insert(
5217 "fallback_chain".to_string(),
5218 serde_json::to_value(&fallback_state.visited_canonical_ids).unwrap_or(Value::Null),
5219 );
5220 let record = self.record_from_parts(
5221 &request,
5222 canonical_id,
5223 request.arguments.clone(),
5224 started_at,
5225 start,
5226 false,
5227 false,
5228 format!("Denied: {reason}"),
5229 metadata,
5230 ToolPolicyDecisionRecord::deny(reason),
5231 None,
5232 false,
5233 false,
5234 );
5235 self.finish_tool_record(&record).await;
5236 return Ok(record);
5237 }
5238 let admitted_canonical_id = canonical_id.clone();
5239 let fallback_state = fallback_state.with_current(canonical_id.clone());
5240 let bindings = resolved.tool.policy_bindings();
5241 let mut executed_arguments = security_engine.prepare_tool_arguments_with_bindings(
5242 &canonical_id,
5243 &request.arguments,
5244 &bindings,
5245 );
5246 let mut metadata = HashMap::new();
5247 let safety = resolved.tool.safety_metadata();
5248 let classification = resolved.tool.classify_call(&executed_arguments);
5249 let initial_recovery_timeout_ms = self.recovery_manager.get_tool_timeout(&canonical_id);
5250 let (limits, _) = Self::effective_tool_limits(
5251 &security_engine,
5252 &canonical_id,
5253 &safety,
5254 &classification,
5255 initial_recovery_timeout_ms,
5256 )?;
5257 self.hooks
5258 .on_tool_start(&canonical_id, &executed_arguments)
5259 .await;
5260 metadata.insert(
5261 "classification".to_string(),
5262 serde_json::to_value(&classification).unwrap_or(Value::Null),
5263 );
5264 metadata.insert(
5265 "effective_limits".to_string(),
5266 serde_json::to_value(&limits).unwrap_or(Value::Null),
5267 );
5268 let policy_snapshot = security_engine.policy_snapshot(&canonical_id);
5269 if !policy_snapshot.is_null() {
5270 metadata.insert("policy_snapshot".to_string(), policy_snapshot.clone());
5271 }
5272
5273 let mut approval_record = Some(ToolApprovalRecord {
5274 status: ToolApprovalStatus::NotRequired,
5275 reason: None,
5276 modified_arguments: None,
5277 });
5278
5279 let mut security_result = security_engine
5280 .validate_tool_execution_with_bindings(&canonical_id, &executed_arguments, &bindings)
5281 .await?;
5282 if (security_result.is_allowed()
5287 || matches!(
5288 &security_result,
5289 SecurityCheckResult::RequireConfirmation { .. }
5290 ))
5291 && let Some((output, reason)) = self.host_tool_unavailability(&canonical_id)
5292 {
5293 let record = self.record_from_parts(
5294 &request,
5295 canonical_id,
5296 executed_arguments,
5297 started_at,
5298 start,
5299 false,
5300 false,
5301 output.to_string(),
5302 metadata,
5303 ToolPolicyDecisionRecord::unavailable(reason),
5304 Some(ToolApprovalRecord {
5305 status: ToolApprovalStatus::Unavailable,
5306 reason: Some(reason.to_string()),
5307 modified_arguments: None,
5308 }),
5309 false,
5310 false,
5311 );
5312 self.finish_tool_record(&record).await;
5313 return Ok(record);
5314 }
5315 match &security_result {
5316 SecurityCheckResult::Allow => {}
5317 SecurityCheckResult::Warn { message } => {
5318 warn!(tool = %canonical_id, message = %message, "Tool security warning");
5319 }
5320 SecurityCheckResult::Block { reason } => {
5321 let record = self.record_from_parts(
5322 &request,
5323 canonical_id,
5324 executed_arguments,
5325 started_at,
5326 start,
5327 false,
5328 false,
5329 format!("Denied: {}", reason),
5330 metadata,
5331 ToolPolicyDecisionRecord::deny(reason.clone()),
5332 approval_record,
5333 false,
5334 false,
5335 );
5336 self.finish_tool_record(&record).await;
5337 return Ok(record);
5338 }
5339 SecurityCheckResult::Unavailable { reason } => {
5340 let record = self.record_from_parts(
5341 &request,
5342 canonical_id,
5343 executed_arguments,
5344 started_at,
5345 start,
5346 false,
5347 false,
5348 format!("Unavailable: {}", reason),
5349 metadata,
5350 ToolPolicyDecisionRecord::unavailable(reason.clone()),
5351 approval_record,
5352 false,
5353 false,
5354 );
5355 self.finish_tool_record(&record).await;
5356 return Ok(record);
5357 }
5358 SecurityCheckResult::RequireConfirmation { message } => {
5359 if self.hitl_engine.is_none() {
5360 approval_record = Some(ToolApprovalRecord {
5361 status: ToolApprovalStatus::Unavailable,
5362 reason: Some("No HITL engine configured".to_string()),
5363 modified_arguments: None,
5364 });
5365 let record = self.record_from_parts(
5366 &request,
5367 canonical_id,
5368 executed_arguments,
5369 started_at,
5370 start,
5371 false,
5372 false,
5373 format!("Approval unavailable: {}", message),
5374 metadata,
5375 ToolPolicyDecisionRecord::approval(message.clone()),
5376 approval_record,
5377 false,
5378 false,
5379 );
5380 self.finish_tool_record(&record).await;
5381 return Ok(record);
5382 }
5383
5384 let check_result = HITLCheckResult::required(
5385 ApprovalTrigger::tool(&canonical_id, executed_arguments.clone()),
5386 HashMap::new(),
5387 message.clone(),
5388 None,
5389 );
5390 match self.request_hitl_approval(check_result).await? {
5391 ApprovalResult::Approved => {
5392 merge_approved_record(&mut approval_record);
5393 }
5394 ApprovalResult::Modified { changes } => {
5395 if let Some(obj) = executed_arguments.as_object_mut() {
5396 for (key, value) in changes {
5397 obj.insert(key, value);
5398 }
5399 }
5400 security_result = security_engine
5401 .validate_tool_execution_with_bindings(
5402 &canonical_id,
5403 &executed_arguments,
5404 &bindings,
5405 )
5406 .await?;
5407 if !matches!(
5408 security_result,
5409 SecurityCheckResult::Allow
5410 | SecurityCheckResult::Warn { .. }
5411 | SecurityCheckResult::RequireConfirmation { .. }
5412 ) {
5413 let reason = security_result
5414 .reason()
5415 .unwrap_or("modified arguments failed policy")
5416 .to_string();
5417 let record = self.record_from_parts(
5418 &request,
5419 canonical_id,
5420 executed_arguments.clone(),
5421 started_at,
5422 start,
5423 false,
5424 false,
5425 reason.clone(),
5426 metadata,
5427 ToolPolicyDecisionRecord::deny(reason),
5428 Some(ToolApprovalRecord {
5429 status: ToolApprovalStatus::Modified,
5430 reason: None,
5431 modified_arguments: Some(executed_arguments),
5432 }),
5433 false,
5434 false,
5435 );
5436 self.finish_tool_record(&record).await;
5437 return Ok(record);
5438 }
5439 approval_record = Some(ToolApprovalRecord {
5440 status: ToolApprovalStatus::Modified,
5441 reason: None,
5442 modified_arguments: Some(executed_arguments.clone()),
5443 });
5444 }
5445 ApprovalResult::Rejected { reason } => {
5446 let reason = reason.unwrap_or_else(|| "rejected".to_string());
5447 approval_record = Some(ToolApprovalRecord {
5448 status: ToolApprovalStatus::Rejected,
5449 reason: Some(reason.clone()),
5450 modified_arguments: None,
5451 });
5452 let record = self.record_from_parts(
5453 &request,
5454 canonical_id,
5455 executed_arguments,
5456 started_at,
5457 start,
5458 false,
5459 false,
5460 format!("Approval rejected: {}", reason),
5461 metadata,
5462 ToolPolicyDecisionRecord::approval(reason),
5463 approval_record,
5464 false,
5465 false,
5466 );
5467 self.finish_tool_record(&record).await;
5468 return Ok(record);
5469 }
5470 ApprovalResult::Timeout => {
5471 approval_record = Some(ToolApprovalRecord {
5472 status: ToolApprovalStatus::Timeout,
5473 reason: Some("approval timeout".to_string()),
5474 modified_arguments: None,
5475 });
5476 let record = self.record_from_parts(
5477 &request,
5478 canonical_id,
5479 executed_arguments,
5480 started_at,
5481 start,
5482 false,
5483 false,
5484 "Approval timed out".to_string(),
5485 metadata,
5486 ToolPolicyDecisionRecord::approval("approval timeout"),
5487 approval_record,
5488 false,
5489 false,
5490 );
5491 self.finish_tool_record(&record).await;
5492 return Ok(record);
5493 }
5494 }
5495 }
5496 }
5497
5498 if approval_record
5499 .as_ref()
5500 .is_some_and(|record| matches!(record.status, ToolApprovalStatus::NotRequired))
5501 && let Some(message) =
5502 security_engine.classification_approval_message(&canonical_id, &classification)
5503 {
5504 if self.hitl_engine.is_none() {
5505 approval_record = Some(ToolApprovalRecord {
5506 status: ToolApprovalStatus::Unavailable,
5507 reason: Some("No HITL engine configured".to_string()),
5508 modified_arguments: None,
5509 });
5510 let record = self.record_from_parts(
5511 &request,
5512 canonical_id,
5513 executed_arguments,
5514 started_at,
5515 start,
5516 false,
5517 false,
5518 format!("Approval unavailable: {}", message),
5519 metadata,
5520 ToolPolicyDecisionRecord::approval(message),
5521 approval_record,
5522 false,
5523 false,
5524 );
5525 self.finish_tool_record(&record).await;
5526 return Ok(record);
5527 }
5528 let check_result = HITLCheckResult::required(
5529 ApprovalTrigger::tool(&canonical_id, executed_arguments.clone()),
5530 HashMap::new(),
5531 message.clone(),
5532 None,
5533 );
5534 match self.request_hitl_approval(check_result).await? {
5535 ApprovalResult::Approved => {
5536 merge_approved_record(&mut approval_record);
5537 }
5538 ApprovalResult::Modified { changes } => {
5539 if let Some(obj) = executed_arguments.as_object_mut() {
5540 for (key, value) in changes {
5541 obj.insert(key, value);
5542 }
5543 }
5544 let modified_security = security_engine
5545 .validate_tool_execution_with_bindings(
5546 &canonical_id,
5547 &executed_arguments,
5548 &bindings,
5549 )
5550 .await?;
5551 if !matches!(
5552 modified_security,
5553 SecurityCheckResult::Allow | SecurityCheckResult::Warn { .. }
5554 ) {
5555 let reason = modified_security
5556 .reason()
5557 .unwrap_or("modified arguments failed policy")
5558 .to_string();
5559 let record = self.record_from_parts(
5560 &request,
5561 canonical_id,
5562 executed_arguments.clone(),
5563 started_at,
5564 start,
5565 false,
5566 false,
5567 reason.clone(),
5568 metadata,
5569 ToolPolicyDecisionRecord::deny(reason),
5570 Some(ToolApprovalRecord {
5571 status: ToolApprovalStatus::Modified,
5572 reason: None,
5573 modified_arguments: Some(executed_arguments),
5574 }),
5575 false,
5576 false,
5577 );
5578 self.finish_tool_record(&record).await;
5579 return Ok(record);
5580 }
5581 approval_record = Some(ToolApprovalRecord {
5582 status: ToolApprovalStatus::Modified,
5583 reason: None,
5584 modified_arguments: Some(executed_arguments.clone()),
5585 });
5586 }
5587 ApprovalResult::Rejected { reason } => {
5588 let reason = reason.unwrap_or_else(|| "rejected".to_string());
5589 let record = self.record_from_parts(
5590 &request,
5591 canonical_id,
5592 executed_arguments,
5593 started_at,
5594 start,
5595 false,
5596 false,
5597 format!("Approval rejected: {}", reason),
5598 metadata,
5599 ToolPolicyDecisionRecord::approval(reason.clone()),
5600 Some(ToolApprovalRecord {
5601 status: ToolApprovalStatus::Rejected,
5602 reason: Some(reason),
5603 modified_arguments: None,
5604 }),
5605 false,
5606 false,
5607 );
5608 self.finish_tool_record(&record).await;
5609 return Ok(record);
5610 }
5611 ApprovalResult::Timeout => {
5612 let record = self.record_from_parts(
5613 &request,
5614 canonical_id,
5615 executed_arguments,
5616 started_at,
5617 start,
5618 false,
5619 false,
5620 "Approval timed out".to_string(),
5621 metadata,
5622 ToolPolicyDecisionRecord::approval("approval timeout"),
5623 Some(ToolApprovalRecord {
5624 status: ToolApprovalStatus::Timeout,
5625 reason: Some("approval timeout".to_string()),
5626 modified_arguments: None,
5627 }),
5628 false,
5629 false,
5630 );
5631 self.finish_tool_record(&record).await;
5632 return Ok(record);
5633 }
5634 }
5635 }
5636
5637 let hitl_lang_ctx = self.build_hitl_language_context();
5638 if let Some(ref hitl_engine) = self.hitl_engine {
5639 let check_result = self
5640 .observe_purpose(
5641 ObservationPurpose::HitlLocalization,
5642 hitl_engine.check_tool_with_localization(
5643 &canonical_id,
5644 &executed_arguments,
5645 &hitl_lang_ctx,
5646 self.approval_handler.as_ref(),
5647 Some(&self.llm_registry),
5648 ),
5649 )
5650 .await?;
5651 if check_result.is_required() {
5652 match self.request_hitl_approval(check_result).await? {
5653 ApprovalResult::Approved => {
5654 merge_approved_record(&mut approval_record);
5655 }
5656 ApprovalResult::Modified { changes } => {
5657 if let Some(obj) = executed_arguments.as_object_mut() {
5658 for (key, value) in changes {
5659 obj.insert(key, value);
5660 }
5661 }
5662 let modified_security = security_engine
5663 .validate_tool_execution_with_bindings(
5664 &canonical_id,
5665 &executed_arguments,
5666 &bindings,
5667 )
5668 .await?;
5669 if !matches!(
5670 modified_security,
5671 SecurityCheckResult::Allow | SecurityCheckResult::Warn { .. }
5672 ) {
5673 let reason = modified_security
5674 .reason()
5675 .unwrap_or("modified arguments failed policy")
5676 .to_string();
5677 let record = self.record_from_parts(
5678 &request,
5679 canonical_id,
5680 executed_arguments.clone(),
5681 started_at,
5682 start,
5683 false,
5684 false,
5685 reason.clone(),
5686 metadata,
5687 ToolPolicyDecisionRecord::deny(reason),
5688 Some(ToolApprovalRecord {
5689 status: ToolApprovalStatus::Modified,
5690 reason: None,
5691 modified_arguments: Some(executed_arguments),
5692 }),
5693 false,
5694 false,
5695 );
5696 self.finish_tool_record(&record).await;
5697 return Ok(record);
5698 }
5699 approval_record = Some(ToolApprovalRecord {
5700 status: ToolApprovalStatus::Modified,
5701 reason: None,
5702 modified_arguments: Some(executed_arguments.clone()),
5703 });
5704 }
5705 ApprovalResult::Rejected { reason } => {
5706 let reason = reason.unwrap_or_else(|| "rejected".to_string());
5707 let record = self.record_from_parts(
5708 &request,
5709 canonical_id,
5710 executed_arguments,
5711 started_at,
5712 start,
5713 false,
5714 false,
5715 format!("Approval rejected: {}", reason),
5716 metadata,
5717 ToolPolicyDecisionRecord::approval(reason.clone()),
5718 Some(ToolApprovalRecord {
5719 status: ToolApprovalStatus::Rejected,
5720 reason: Some(reason),
5721 modified_arguments: None,
5722 }),
5723 false,
5724 false,
5725 );
5726 self.finish_tool_record(&record).await;
5727 return Ok(record);
5728 }
5729 ApprovalResult::Timeout => {
5730 let record = self.record_from_parts(
5731 &request,
5732 canonical_id,
5733 executed_arguments,
5734 started_at,
5735 start,
5736 false,
5737 false,
5738 "Approval timed out".to_string(),
5739 metadata,
5740 ToolPolicyDecisionRecord::approval("approval timeout"),
5741 Some(ToolApprovalRecord {
5742 status: ToolApprovalStatus::Timeout,
5743 reason: Some("approval timeout".to_string()),
5744 modified_arguments: None,
5745 }),
5746 false,
5747 false,
5748 );
5749 self.finish_tool_record(&record).await;
5750 return Ok(record);
5751 }
5752 }
5753 }
5754
5755 let condition_check = self
5756 .observe_purpose(
5757 ObservationPurpose::HitlLocalization,
5758 hitl_engine.check_conditions_with_localization(
5759 &executed_arguments,
5760 &hitl_lang_ctx,
5761 self.approval_handler.as_ref(),
5762 Some(&self.llm_registry),
5763 ),
5764 )
5765 .await?;
5766 if condition_check.is_required() {
5767 match self.request_hitl_approval(condition_check).await? {
5768 ApprovalResult::Approved => {
5769 merge_approved_record(&mut approval_record);
5770 }
5771 ApprovalResult::Modified { changes } => {
5772 if let Some(obj) = executed_arguments.as_object_mut() {
5773 for (key, value) in changes {
5774 obj.insert(key, value);
5775 }
5776 }
5777 let modified_security = security_engine
5778 .validate_tool_execution_with_bindings(
5779 &canonical_id,
5780 &executed_arguments,
5781 &bindings,
5782 )
5783 .await?;
5784 if !matches!(
5785 modified_security,
5786 SecurityCheckResult::Allow | SecurityCheckResult::Warn { .. }
5787 ) {
5788 let reason = modified_security
5789 .reason()
5790 .unwrap_or("modified arguments failed policy")
5791 .to_string();
5792 let record = self.record_from_parts(
5793 &request,
5794 canonical_id,
5795 executed_arguments,
5796 started_at,
5797 start,
5798 false,
5799 false,
5800 reason.clone(),
5801 metadata,
5802 ToolPolicyDecisionRecord::deny(reason),
5803 approval_record,
5804 false,
5805 false,
5806 );
5807 self.finish_tool_record(&record).await;
5808 return Ok(record);
5809 }
5810 approval_record = Some(ToolApprovalRecord {
5811 status: ToolApprovalStatus::Modified,
5812 reason: None,
5813 modified_arguments: Some(executed_arguments.clone()),
5814 });
5815 }
5816 ApprovalResult::Rejected { reason } => {
5817 let reason = reason.unwrap_or_else(|| "rejected".to_string());
5818 let record = self.record_from_parts(
5819 &request,
5820 canonical_id,
5821 executed_arguments,
5822 started_at,
5823 start,
5824 false,
5825 false,
5826 format!("Approval rejected: {}", reason),
5827 metadata,
5828 ToolPolicyDecisionRecord::approval(reason.clone()),
5829 Some(ToolApprovalRecord {
5830 status: ToolApprovalStatus::Rejected,
5831 reason: Some(reason),
5832 modified_arguments: None,
5833 }),
5834 false,
5835 false,
5836 );
5837 self.finish_tool_record(&record).await;
5838 return Ok(record);
5839 }
5840 ApprovalResult::Timeout => {
5841 let record = self.record_from_parts(
5842 &request,
5843 canonical_id,
5844 executed_arguments,
5845 started_at,
5846 start,
5847 false,
5848 false,
5849 "Approval timed out".to_string(),
5850 metadata,
5851 ToolPolicyDecisionRecord::approval("approval timeout"),
5852 Some(ToolApprovalRecord {
5853 status: ToolApprovalStatus::Timeout,
5854 reason: Some("approval timeout".to_string()),
5855 modified_arguments: None,
5856 }),
5857 false,
5858 false,
5859 );
5860 self.finish_tool_record(&record).await;
5861 return Ok(record);
5862 }
5863 }
5864 }
5865 }
5866
5867 executed_arguments = security_engine.prepare_tool_arguments_with_bindings(
5872 &canonical_id,
5873 &executed_arguments,
5874 &bindings,
5875 );
5876 if let Some(record) = approval_record.as_mut()
5877 && matches!(record.status, ToolApprovalStatus::Modified)
5878 {
5879 record.modified_arguments = Some(executed_arguments.clone());
5880 }
5881 let binding_security_result = security_engine
5882 .validate_tool_execution_with_bindings(&canonical_id, &executed_arguments, &bindings)
5883 .await?;
5884 let approval_confirmation_required = matches!(
5885 binding_security_result,
5886 SecurityCheckResult::RequireConfirmation { .. }
5887 ) || security_engine
5888 .classification_approval_message(
5889 &canonical_id,
5890 &resolved.tool.classify_call(&executed_arguments),
5891 )
5892 .is_some();
5893 let approval_binding = approval_record.as_ref().and_then(|record| {
5894 matches!(
5895 record.status,
5896 ToolApprovalStatus::Approved | ToolApprovalStatus::Modified
5897 )
5898 .then(|| ToolApprovalBinding {
5899 canonical_id: canonical_id.clone(),
5900 arguments: executed_arguments.clone(),
5901 confirmation_required: approval_confirmation_required,
5902 policy_version: security_engine.policy_version(),
5903 runtime_control_version: approval_control_snapshot.version,
5904 state_generation: initial_scope_snapshot.state_generation,
5905 reviewed_tool: Arc::clone(&resolved.tool),
5906 })
5907 });
5908
5909 let control_snapshot = self.runtime_safety_snapshot();
5914 let resolved = self.tools.resolve(&request.requested_name);
5915 let registry_version = self.tools.version();
5916 let mut versions = ToolDecisionVersions {
5917 policy: control_snapshot.tool_security.policy_version(),
5918 registry: registry_version,
5919 runtime_control: control_snapshot.version,
5920 state: None,
5921 };
5922 metadata.insert(
5923 "runtime_scope_snapshot".to_string(),
5924 serde_json::to_value(&control_snapshot.tool_scope_override).unwrap_or(Value::Null),
5925 );
5926 let resolved = match resolved {
5927 Some(resolved) => resolved,
5928 None => {
5929 let reason = format!(
5930 "Tool '{}' became unavailable after approval",
5931 request.requested_name
5932 );
5933 let record = self.record_from_parts_at(
5934 &request,
5935 request.requested_name.clone(),
5936 executed_arguments,
5937 started_at,
5938 start,
5939 false,
5940 false,
5941 reason.clone(),
5942 metadata,
5943 ToolPolicyDecisionRecord::unavailable(reason),
5944 approval_record,
5945 false,
5946 false,
5947 versions,
5948 );
5949 self.finish_tool_record(&record).await;
5950 return Ok(record);
5951 }
5952 };
5953
5954 let canonical_id = resolved.identity.canonical_id.clone();
5955 if let Some(reason) =
5956 fallback_state.final_rejection_reason(&admitted_canonical_id, &canonical_id)
5957 {
5958 metadata.insert(
5962 "fallback_chain".to_string(),
5963 serde_json::to_value(&fallback_state.visited_canonical_ids).unwrap_or(Value::Null),
5964 );
5965 metadata.insert(
5966 "final_resolved_canonical_id".to_string(),
5967 Value::String(canonical_id),
5968 );
5969 let record = self.record_from_parts_at(
5970 &request,
5971 admitted_canonical_id,
5972 executed_arguments,
5973 started_at,
5974 start,
5975 false,
5976 false,
5977 format!("Denied: {reason}"),
5978 metadata,
5979 ToolPolicyDecisionRecord::deny(reason),
5980 approval_record,
5981 false,
5982 false,
5983 versions,
5984 );
5985 self.finish_tool_record(&record).await;
5986 return Ok(record);
5987 }
5988 let bindings = resolved.tool.policy_bindings();
5989 let final_arguments = control_snapshot
5990 .tool_security
5991 .prepare_tool_arguments_with_bindings(&canonical_id, &executed_arguments, &bindings);
5992 if let Some(record) = approval_record.as_mut()
5993 && matches!(record.status, ToolApprovalStatus::Modified)
5994 {
5995 record.modified_arguments = Some(final_arguments.clone());
5996 }
5997 let classification = resolved.tool.classify_call(&final_arguments);
5998 let safety = resolved.tool.safety_metadata();
5999 let security_engine = control_snapshot.tool_security;
6000 let tool_config = self.recovery_manager.get_tool_config(&canonical_id).clone();
6001 let recovery_timeout_ms = self.recovery_manager.get_tool_timeout(&canonical_id);
6002 metadata.insert(
6003 "classification".to_string(),
6004 serde_json::to_value(&classification).unwrap_or(Value::Null),
6005 );
6006 let (limits, timeout) = match Self::effective_tool_limits(
6010 &security_engine,
6011 &canonical_id,
6012 &safety,
6013 &classification,
6014 recovery_timeout_ms,
6015 ) {
6016 Ok(effective) => effective,
6017 Err(error) => {
6018 let reason = error.to_string();
6019 metadata.insert(
6020 "configuration_error".to_string(),
6021 Value::String(reason.clone()),
6022 );
6023 let record = self.record_from_parts_at(
6024 &request,
6025 canonical_id,
6026 final_arguments,
6027 started_at,
6028 start,
6029 false,
6030 false,
6031 format!("Denied: {reason}"),
6032 metadata,
6033 ToolPolicyDecisionRecord::deny(reason),
6034 approval_record,
6035 false,
6036 false,
6037 versions,
6038 );
6039 self.finish_tool_record(&record).await;
6040 return Ok(record);
6041 }
6042 };
6043 let policy_snapshot = security_engine.policy_snapshot(&canonical_id);
6044 let resource_lock_keys =
6045 tool_resource_lock_keys(&canonical_id, &final_arguments, &bindings, &classification);
6046 metadata.insert(
6047 "effective_limits".to_string(),
6048 serde_json::to_value(&limits).unwrap_or(Value::Null),
6049 );
6050 metadata.insert(
6051 "resource_lock_keys".to_string(),
6052 serde_json::to_value(&resource_lock_keys).unwrap_or(Value::Null),
6053 );
6054 if policy_snapshot.is_null() {
6055 metadata.remove("policy_snapshot");
6056 } else {
6057 metadata.insert("policy_snapshot".to_string(), policy_snapshot.clone());
6058 }
6059
6060 let final_denial = |canonical_id: String,
6061 output: String,
6062 policy: ToolPolicyDecisionRecord,
6063 metadata: HashMap<String, Value>,
6064 decision_versions: ToolDecisionVersions| {
6065 self.record_from_parts_at(
6066 &request,
6067 canonical_id,
6068 final_arguments.clone(),
6069 started_at,
6070 start,
6071 false,
6072 false,
6073 output,
6074 metadata,
6075 policy,
6076 approval_record.clone(),
6077 false,
6078 false,
6079 decision_versions,
6080 )
6081 };
6082
6083 if control_snapshot.emergency_deny {
6084 let reason = "Tool execution is disabled by runtime control".to_string();
6085 let record = final_denial(
6086 canonical_id,
6087 reason.clone(),
6088 ToolPolicyDecisionRecord::deny(reason),
6089 metadata,
6090 versions,
6091 );
6092 self.finish_tool_record(&record).await;
6093 return Ok(record);
6094 }
6095
6096 let available_snapshot = self
6101 .get_available_tool_ids_snapshot_for_scope(
6102 control_snapshot.tool_scope_override.as_deref(),
6103 )
6104 .await?;
6105 versions.state = available_snapshot.state_generation;
6106 metadata.insert(
6107 "available_tool_ids_snapshot".to_string(),
6108 serde_json::to_value(&available_snapshot.tool_ids).unwrap_or(Value::Null),
6109 );
6110 metadata.insert(
6111 "state_generation_snapshot".to_string(),
6112 serde_json::to_value(available_snapshot.state_generation).unwrap_or(Value::Null),
6113 );
6114 if !available_snapshot
6115 .tool_ids
6116 .iter()
6117 .any(|tool_id| tool_id == &canonical_id)
6118 {
6119 let reason = format!(
6120 "Tool '{}' is not available in the final runtime scope",
6121 canonical_id
6122 );
6123 let record = final_denial(
6124 canonical_id,
6125 reason.clone(),
6126 ToolPolicyDecisionRecord::deny(reason),
6127 metadata,
6128 versions,
6129 );
6130 self.finish_tool_record(&record).await;
6131 return Ok(record);
6132 }
6133
6134 let final_security_result = security_engine
6139 .validate_tool_execution_with_bindings(&canonical_id, &final_arguments, &bindings)
6140 .await?;
6141 match &final_security_result {
6142 SecurityCheckResult::Block { reason } => {
6143 let record = final_denial(
6144 canonical_id,
6145 format!("Denied: {}", reason),
6146 ToolPolicyDecisionRecord::deny(reason.clone()),
6147 metadata,
6148 versions,
6149 );
6150 self.finish_tool_record(&record).await;
6151 return Ok(record);
6152 }
6153 SecurityCheckResult::Unavailable { reason } => {
6154 let record = final_denial(
6155 canonical_id,
6156 format!("Unavailable: {}", reason),
6157 ToolPolicyDecisionRecord::unavailable(reason.clone()),
6158 metadata,
6159 versions,
6160 );
6161 self.finish_tool_record(&record).await;
6162 return Ok(record);
6163 }
6164 SecurityCheckResult::Warn { message } => {
6165 warn!(tool = %canonical_id, message = %message, "Tool security warning after approval");
6166 }
6167 SecurityCheckResult::Allow | SecurityCheckResult::RequireConfirmation { .. } => {}
6168 }
6169 let final_confirmation_required = matches!(
6170 final_security_result,
6171 SecurityCheckResult::RequireConfirmation { .. }
6172 ) || security_engine
6173 .classification_approval_message(&canonical_id, &classification)
6174 .is_some();
6175 let stale_approval = approval_binding.as_ref().is_some_and(|binding| {
6176 binding.is_stale(
6177 &canonical_id,
6178 &final_arguments,
6179 final_confirmation_required,
6180 versions,
6181 &resolved.tool,
6182 )
6183 });
6184 if stale_approval {
6185 let reason = "Approval became stale before final admission".to_string();
6186 let record = final_denial(
6187 canonical_id,
6188 reason.clone(),
6189 ToolPolicyDecisionRecord::deny(reason),
6190 metadata,
6191 versions,
6192 );
6193 self.finish_tool_record(&record).await;
6194 return Ok(record);
6195 }
6196 if final_confirmation_required && approval_binding.is_none() {
6197 let reason = "Final policy requires fresh approval".to_string();
6198 let record = final_denial(
6199 canonical_id,
6200 reason.clone(),
6201 ToolPolicyDecisionRecord::approval(reason),
6202 metadata,
6203 versions,
6204 );
6205 self.finish_tool_record(&record).await;
6206 return Ok(record);
6207 }
6208
6209 if let Some((_, reason)) = self.host_tool_unavailability(&canonical_id) {
6210 let record = final_denial(
6211 canonical_id,
6212 reason.to_string(),
6213 ToolPolicyDecisionRecord::unavailable(reason),
6214 metadata,
6215 versions,
6216 );
6217 self.finish_tool_record(&record).await;
6218 return Ok(record);
6219 }
6220
6221 let Some(resource_guards) = self.acquire_tool_resource_locks(&resource_lock_keys).await
6226 else {
6227 let reason = "Tool execution cancelled while waiting for resource locks".to_string();
6231 let mut record = final_denial(
6232 canonical_id,
6233 reason.clone(),
6234 ToolPolicyDecisionRecord::deny(reason),
6235 metadata,
6236 versions,
6237 );
6238 record.cancelled = true;
6239 record.cancellation_reason = Some("runtime control cancellation".to_string());
6240 self.finish_tool_record(&record).await;
6241 return Ok(record);
6242 };
6243
6244 let admission = self.admit_tool_execution(
6249 versions.runtime_control,
6250 versions.policy,
6251 versions.state,
6252 &canonical_id,
6253 );
6254 if !matches!(admission, SecurityCheckResult::Allow) {
6255 let latest_control = self.runtime_safety_snapshot();
6256 let reason = admission
6257 .reason()
6258 .unwrap_or("tool admission was denied")
6259 .to_string();
6260 let policy = if admission.is_unavailable() {
6261 ToolPolicyDecisionRecord::unavailable(reason.clone())
6262 } else {
6263 ToolPolicyDecisionRecord::deny(reason.clone())
6264 };
6265 let record = self.record_from_parts_at(
6266 &request,
6267 canonical_id,
6268 final_arguments,
6269 started_at,
6270 start,
6271 false,
6272 false,
6273 reason,
6274 metadata,
6275 policy,
6276 approval_record,
6277 false,
6278 false,
6279 ToolDecisionVersions {
6280 policy: latest_control.tool_security.policy_version(),
6281 registry: versions.registry,
6282 runtime_control: latest_control.version,
6283 state: self
6284 .state_machine
6285 .as_ref()
6286 .map(|state_machine| state_machine.generation()),
6287 },
6288 );
6289 self.finish_tool_record_after_resource_guards(resource_guards, &record)
6290 .await;
6291 return Ok(record);
6292 }
6293 let executed_arguments = final_arguments;
6294
6295 let turn_actor = current_turn_actor_context();
6296 let actor = ToolActorContext {
6297 actor_id: turn_actor
6298 .as_ref()
6299 .and_then(|context| context.effective_actor_id().map(str::to_string))
6300 .or_else(|| self.actor_id()),
6301 origin_actor_id: turn_actor
6302 .as_ref()
6303 .and_then(|context| context.origin_actor_id.clone()),
6304 sender_agent_id: turn_actor
6305 .as_ref()
6306 .and_then(|context| context.sender_agent_id.clone()),
6307 };
6308 let tool_context = ToolExecutionContext {
6309 requested_name: request.requested_name.clone(),
6310 canonical_id: canonical_id.clone(),
6311 display_name: resolved.identity.display_name.clone(),
6312 provider_id: resolved.identity.provider_id.clone(),
6313 registry_version: versions.registry,
6314 policy_version: versions.policy,
6315 runtime_control_version: versions.runtime_control,
6316 call_id: request.call_id.clone(),
6317 source: request.source.clone(),
6318 actor,
6319 cancellation: ToolCancellationToken::new(
6320 Arc::clone(&self.runtime_control.emergency_deny),
6321 Some("runtime control cancellation".to_string()),
6322 ),
6323 started_at,
6324 deadline: None,
6325 permission: ToolPolicyDecisionRecord::allow(),
6326 approval: approval_record.clone(),
6327 classification: classification.clone(),
6328 safety,
6329 limits: limits.clone(),
6330 policy_snapshot,
6331 custom_config: security_engine.custom_config(&canonical_id),
6332 };
6333 let (mut result, timed_out, cancelled, invoked) = self
6334 .run_tool_with_retries(
6335 &canonical_id,
6336 resolved.tool.clone(),
6337 executed_arguments.clone(),
6338 tool_context,
6339 timeout,
6340 tool_config.max_retries,
6341 )
6342 .await?;
6343
6344 let fallback_tool = if !result.success && !cancelled {
6348 match &tool_config.on_failure {
6349 ToolFailureAction::Skip => {
6350 result = ToolResult::ok(format!(
6351 "{{\"skipped\": true, \"reason\": \"Tool '{}' was skipped after failure\"}}",
6352 canonical_id
6353 ));
6354 None
6355 }
6356 ToolFailureAction::Fallback { fallback_tool } => Some(fallback_tool.clone()),
6357 ToolFailureAction::ReportError => None,
6358 }
6359 } else {
6360 None
6361 };
6362
6363 let output_cap = limits.max_output_chars;
6364 let (output, output_truncated) =
6365 Self::truncate_tool_output(result.output.clone(), output_cap);
6366 if let Some(result_metadata) = result.metadata {
6367 metadata.extend(result_metadata);
6368 }
6369 let mut record = self.record_from_parts_at(
6370 &request,
6371 canonical_id,
6372 executed_arguments,
6373 started_at,
6374 start,
6375 invoked,
6376 result.success,
6377 output,
6378 metadata,
6379 ToolPolicyDecisionRecord::allow(),
6380 approval_record,
6381 timed_out,
6382 output_truncated,
6383 versions,
6384 );
6385 record.cancelled = cancelled;
6386 if cancelled {
6387 record.cancellation_reason = Some("runtime control cancellation".to_string());
6388 }
6389 if let Some(fallback_tool) = fallback_tool {
6390 let fallback_arguments = record.executed_arguments.clone();
6391 let original_tool = record.canonical_id.clone();
6392 self.finish_tool_record_after_resource_guards(resource_guards, &record)
6396 .await;
6397 let fallback_request = ToolExecutionRequest::new(
6398 request.call_id.clone(),
6399 fallback_tool,
6400 fallback_arguments,
6401 ToolCallSource::Fallback { original_tool },
6402 );
6403 return Box::pin(self.execute_tool_record_inner(fallback_request, fallback_state))
6404 .await;
6405 }
6406 self.finish_tool_record_after_resource_guards(resource_guards, &record)
6407 .await;
6408 Ok(record)
6409 }
6410
6411 #[instrument(skip(self, tool_call), fields(tool = %tool_call.name))]
6412 async fn execute_tool_smart(&self, tool_call: &ToolCall) -> Result<String> {
6413 let record = self
6414 .execute_tool_record(ToolExecutionRequest::new(
6415 tool_call.id.clone(),
6416 tool_call.name.clone(),
6417 tool_call.arguments.clone(),
6418 ToolCallSource::Model,
6419 ))
6420 .await?;
6421 if record.success {
6422 Ok(record.model_output_string())
6423 } else if matches!(record.policy.outcome, PermissionOutcome::RequiresApproval) {
6424 Err(AgentError::HITLRejected(record.model_output_string()))
6425 } else {
6426 Err(AgentError::Tool(record.model_output_string()))
6427 }
6428 }
6429
6430 async fn select_skill_candidate(&self, input: &str) -> Result<Option<SkillCandidate>> {
6436 let Some(ref router) = self.skill_router else {
6437 return Ok(None);
6438 };
6439 let available_skills = self.get_available_skills();
6440 if available_skills.is_empty() {
6441 return Ok(None);
6442 }
6443 let skill_ids: Vec<&str> = available_skills.iter().map(|s| s.id.as_str()).collect();
6444 let Some(skill_id) = self
6445 .observe_purpose(
6446 ObservationPurpose::SkillRouting,
6447 router.select_skill_filtered(input, &skill_ids),
6448 )
6449 .await?
6450 else {
6451 return Ok(None);
6452 };
6453 let skill = router
6454 .get_skill(&skill_id)
6455 .cloned()
6456 .ok_or_else(|| AgentError::Skill(format!("Skill not found: {}", skill_id)))?;
6457 info!(skill_id = %skill_id, "Skill selected");
6458 Ok(Some(SkillCandidate::new(skill_id, skill)))
6459 }
6460
6461 async fn commit_skill_candidate_route_result(
6466 &self,
6467 candidate: SkillCandidate,
6468 input: &str,
6469 ) -> Result<SkillRouteResult> {
6470 let skill_id = candidate.skill_id;
6471 let skill = candidate.skill;
6472 let expected_state_generation = self
6473 .state_machine
6474 .as_ref()
6475 .map(|state_machine| state_machine.generation());
6476 let expected_disambiguation_epoch = self.disambiguation_epoch.load(Ordering::SeqCst);
6477 if let Some(ref skill_disambig) = skill.disambiguation
6478 && skill_disambig.enabled.unwrap_or(false)
6479 && let Some(ref disambiguator) = self.disambiguation_manager
6480 {
6481 let context = self.build_disambiguation_context().await?;
6482 let state_override = self
6483 .state_machine
6484 .as_ref()
6485 .and_then(|sm| sm.current_definition())
6486 .and_then(|def| def.disambiguation.clone());
6487
6488 let disambiguation_result = self
6489 .observe_purpose(
6490 ObservationPurpose::DisambiguationDetection,
6491 disambiguator.process_input_with_override(
6492 input,
6493 &context,
6494 state_override.as_ref(),
6495 Some(skill_disambig),
6496 ),
6497 )
6498 .await?;
6499 let current_state_generation = self
6500 .state_machine
6501 .as_ref()
6502 .map(|state_machine| state_machine.generation());
6503 if current_state_generation != expected_state_generation
6504 || self.disambiguation_epoch.load(Ordering::SeqCst) != expected_disambiguation_epoch
6505 {
6506 disambiguator.clear_pending().await;
6507 *self.pending_skill_id.write() = None;
6508 return Err(AgentError::Other(
6509 "State or reset ownership changed during skill disambiguation".to_string(),
6510 ));
6511 }
6512 match disambiguation_result {
6513 DisambiguationResult::Clear => {
6514 debug!(skill_id = %skill_id, "Skill disambiguation: clear");
6515 }
6516 DisambiguationResult::NeedsClarification {
6517 question,
6518 detection,
6519 } => {
6520 let admission = self
6521 .admit_disambiguation_redispatch(
6522 expected_disambiguation_epoch,
6523 expected_state_generation,
6524 )
6525 .await?;
6526 let awaiting_confirmation = disambiguator.has_pending_confirmation().await;
6527 info!(
6528 skill_id = %skill_id,
6529 ambiguity_type = ?detection.ambiguity_type,
6530 confidence = detection.confidence,
6531 "Skill requires clarification before execution"
6532 );
6533 *self.pending_skill_id.write() = Some(skill_id.clone());
6534 let response = AgentResponse::new(&question.question).with_metadata(
6535 "disambiguation",
6536 serde_json::json!({
6537 "status": if awaiting_confirmation { "awaiting_confirmation" } else { "awaiting_clarification" },
6538 "skill_id": skill_id,
6539 "options": question.options,
6540 "clarifying": question.clarifying,
6541 "detection": {
6542 "type": detection.ambiguity_type,
6543 "confidence": detection.confidence,
6544 "what_is_unclear": detection.what_is_unclear,
6545 }
6546 }),
6547 );
6548 drop(admission);
6549 return Ok(SkillRouteResult::NeedsClarification {
6550 response,
6551 ownership: Some(DisambiguationOwnership {
6552 epoch: expected_disambiguation_epoch,
6553 state_generation: expected_state_generation,
6554 }),
6555 });
6556 }
6557 DisambiguationResult::Clarified { enriched_input, .. } => {
6558 info!(skill_id = %skill_id, enriched = %enriched_input, "Skill disambiguation clarified");
6559 let admission = self
6560 .admit_disambiguation_redispatch(
6561 expected_disambiguation_epoch,
6562 expected_state_generation,
6563 )
6564 .await?;
6565 drop(admission);
6566 let content = self.execute_skill(&skill, &enriched_input).await?;
6567 return Ok(SkillRouteResult::Response { skill_id, content });
6568 }
6569 DisambiguationResult::ProceedWithBestGuess { enriched_input } => {
6570 info!(skill_id = %skill_id, "Skill disambiguation best guess");
6571 let admission = self
6572 .admit_disambiguation_redispatch(
6573 expected_disambiguation_epoch,
6574 expected_state_generation,
6575 )
6576 .await?;
6577 drop(admission);
6578 let content = self.execute_skill(&skill, &enriched_input).await?;
6579 return Ok(SkillRouteResult::Response { skill_id, content });
6580 }
6581 DisambiguationResult::GiveUp { reason } => {
6582 warn!(skill_id = %skill_id, reason = %reason, "Skill disambiguation gave up");
6583 let apology = self
6584 .generate_localized_apology(
6585 "Generate a brief, polite apology saying you couldn't understand the request. Be concise.",
6586 &reason,
6587 )
6588 .await
6589 .unwrap_or_else(|_| {
6590 format!("I'm sorry, I couldn't understand your request: {}", reason)
6591 });
6592 return Ok(SkillRouteResult::NeedsClarification {
6593 response: AgentResponse::new(&apology),
6594 ownership: None,
6595 });
6596 }
6597 DisambiguationResult::Escalate { reason } => {
6598 info!(skill_id = %skill_id, reason = %reason, "Skill disambiguation escalating");
6599 let apology = self
6600 .generate_localized_apology(
6601 "Explain briefly that you're transferring the user to a human agent for help.",
6602 &reason,
6603 )
6604 .await
6605 .unwrap_or_else(|_| {
6606 format!("I need human assistance to help with your request: {}", reason)
6607 });
6608 return Ok(SkillRouteResult::NeedsClarification {
6609 response: AgentResponse::new(&apology),
6610 ownership: None,
6611 });
6612 }
6613 DisambiguationResult::Abandoned { .. } => {
6614 debug!(skill_id = %skill_id, "Skill disambiguation abandoned");
6615 return Ok(SkillRouteResult::NoMatch);
6616 }
6617 }
6618 }
6619 let admission = self
6620 .admit_disambiguation_redispatch(
6621 expected_disambiguation_epoch,
6622 expected_state_generation,
6623 )
6624 .await?;
6625 drop(admission);
6626 let content = self.execute_skill(&skill, input).await?;
6627 Ok(SkillRouteResult::Response { skill_id, content })
6628 }
6629
6630 async fn try_skill_route(&self, input: &str) -> Result<SkillRouteResult> {
6632 if let Some(candidate) = self.select_skill_candidate(input).await? {
6633 self.commit_skill_candidate_route_result(candidate, input)
6634 .await
6635 } else {
6636 Ok(SkillRouteResult::NoMatch)
6637 }
6638 }
6639
6640 fn skill_clarification_needs_memory_record(response: &AgentResponse) -> bool {
6643 response
6644 .metadata
6645 .as_ref()
6646 .and_then(|m| m.get("disambiguation"))
6647 .and_then(|d| d.get("status"))
6648 .and_then(|s| s.as_str())
6649 == Some("awaiting_clarification")
6650 }
6651
6652 async fn commit_winning_skill_candidate(
6659 &self,
6660 candidate: SkillCandidate,
6661 processed_input: &str,
6662 input_context: &HashMap<String, Value>,
6663 ) -> Result<Option<AgentResponse>> {
6664 self.commit_root_user_message(processed_input).await?;
6665 match self
6666 .commit_skill_candidate_route_result(candidate, processed_input)
6667 .await?
6668 {
6669 SkillRouteResult::Response { skill_id, content } => self
6670 .handle_skill_response(processed_input, &skill_id, content, input_context)
6671 .await
6672 .map(Some),
6673 SkillRouteResult::NeedsClarification {
6674 response,
6675 ownership,
6676 } => {
6677 let admission = self
6678 .admit_optional_disambiguation_ownership(ownership)
6679 .await?;
6680 if Self::skill_clarification_needs_memory_record(&response) {
6681 self.memory
6682 .add_message(ChatMessage::assistant(&response.content))
6683 .await?;
6684 }
6685 drop(admission);
6686 self.finish_turn_if_root(&response).await?;
6687 Ok(Some(response))
6688 }
6689 SkillRouteResult::NoMatch => Ok(None),
6690 }
6691 }
6692
6693 async fn execute_skill(&self, skill: &SkillDefinition, input: &str) -> Result<String> {
6695 if let Some(ref executor) = self.skill_executor {
6696 let skill_reasoning = self.get_skill_reasoning_config(skill);
6697 let skill_reflection = self.get_skill_reflection_config(skill);
6698
6699 debug!(
6700 skill_id = %skill.id,
6701 reasoning_mode = ?skill_reasoning.mode,
6702 reflection_enabled = ?skill_reflection.enabled,
6703 "Skill reasoning/reflection config"
6704 );
6705
6706 let response = self
6707 .observe_purpose(
6708 ObservationPurpose::SkillPrompt,
6709 executor.execute_with_invoker(skill, input, serde_json::json!({}), self),
6710 )
6711 .await?;
6712
6713 if skill_reflection.requires_evaluation() && skill_reflection.is_enabled() {
6714 let should_reflect = self
6715 .should_reflect_with_config(input, &response, &skill_reflection)
6716 .await?;
6717 if should_reflect {
6718 let evaluated = self
6719 .evaluate_and_retry_with_config(input, response, &skill_reflection)
6720 .await?;
6721 return Ok(evaluated);
6722 }
6723 }
6724
6725 return Ok(response);
6726 }
6727 Err(AgentError::Skill(
6728 "No skill executor configured".to_string(),
6729 ))
6730 }
6731
6732 async fn execute_skill_by_id(&self, skill_id: &str, input: &str) -> Result<String> {
6735 let skill = self
6736 .skill_router
6737 .as_ref()
6738 .and_then(|r| r.get_skill(skill_id).cloned())
6739 .ok_or_else(|| AgentError::Skill(format!("Skill not found: {}", skill_id)))?;
6740 self.execute_skill(&skill, input).await
6741 }
6742
6743 async fn should_reflect_with_config(
6744 &self,
6745 input: &str,
6746 response: &str,
6747 config: &ReflectionConfig,
6748 ) -> Result<bool> {
6749 if !config.requires_evaluation() {
6750 return Ok(false);
6751 }
6752
6753 if config.is_enabled() {
6754 return Ok(true);
6755 }
6756
6757 let evaluator_llm = config
6758 .evaluator_llm
6759 .as_ref()
6760 .and_then(|alias| self.llm_registry.get(alias).ok())
6761 .or_else(|| self.llm_registry.router().ok())
6762 .or_else(|| self.llm_registry.default().ok());
6763
6764 let Some(llm) = evaluator_llm else {
6765 return Ok(false);
6766 };
6767
6768 let response_preview: String = response.chars().take(500).collect();
6769 let prompt = format!(
6770 r#"Should this response be evaluated for quality? Consider if it's a complex or important response.
6771
6772User query: "{}"
6773Response: "{}"
6774
6775Answer YES or NO only."#,
6776 input, response_preview
6777 );
6778
6779 let messages = vec![ChatMessage::user(&prompt)];
6780 let result = self
6781 .observe_purpose(
6782 ObservationPurpose::ReflectionDecision,
6783 llm.complete(&messages, None),
6784 )
6785 .await;
6786
6787 match result {
6788 Ok(resp) => Ok(resp.content.trim().to_uppercase().contains("YES")),
6789 Err(_) => Ok(false),
6790 }
6791 }
6792
6793 async fn evaluate_and_retry_with_config(
6794 &self,
6795 input: &str,
6796 mut response: String,
6797 config: &ReflectionConfig,
6798 ) -> Result<String> {
6799 let llm = self.get_state_llm()?;
6800 let mut attempts = 0u32;
6801 let max_retries = config.max_retries;
6802
6803 loop {
6804 let evaluation = self
6805 .evaluate_response_with_config(input, &response, config)
6806 .await?;
6807
6808 if evaluation.passed || attempts >= max_retries {
6809 info!(
6810 passed = evaluation.passed,
6811 confidence = evaluation.confidence,
6812 attempts = attempts + 1,
6813 "Skill reflection evaluation complete"
6814 );
6815 return Ok(response);
6816 }
6817
6818 debug!(
6819 attempt = attempts + 1,
6820 failed_criteria = evaluation.failed_criteria().count(),
6821 "Skill response did not meet criteria, retrying"
6822 );
6823
6824 let feedback: Vec<String> = evaluation
6825 .failed_criteria()
6826 .map(|c| format!("- {}", c.criterion))
6827 .collect();
6828
6829 let retry_prompt = format!(
6830 "Your previous response did not meet these criteria:\n{}\n\nPlease provide an improved response to: {}",
6831 feedback.join("\n"),
6832 input
6833 );
6834
6835 let messages = vec![ChatMessage::user(&retry_prompt)];
6836 let retry_response = self
6837 .observe_purpose(
6838 ObservationPurpose::ReflectionEvaluation,
6839 llm.complete(&messages, None),
6840 )
6841 .await
6842 .map_err(|e| AgentError::LLM(e.to_string()))?;
6843
6844 response = retry_response.content.trim().to_string();
6845 attempts += 1;
6846 }
6847 }
6848
6849 async fn evaluate_response_with_config(
6850 &self,
6851 input: &str,
6852 response: &str,
6853 config: &ReflectionConfig,
6854 ) -> Result<EvaluationResult> {
6855 let evaluator_llm = config
6856 .evaluator_llm
6857 .as_ref()
6858 .and_then(|alias| self.llm_registry.get(alias).ok())
6859 .or_else(|| self.llm_registry.router().ok())
6860 .or_else(|| self.llm_registry.default().ok())
6861 .ok_or_else(|| AgentError::Config("No LLM available for evaluation".into()))?;
6862
6863 let criteria = &config.criteria;
6864 let criteria_list = criteria
6865 .iter()
6866 .enumerate()
6867 .map(|(i, c)| format!("{}. {}", i + 1, c))
6868 .collect::<Vec<_>>()
6869 .join("\n");
6870
6871 let prompt = format!(
6872 r#"Evaluate this response against the criteria.
6873
6874User query: "{}"
6875
6876Response to evaluate: "{}"
6877
6878Criteria:
6879{}
6880
6881For each criterion, respond with:
6882- criterion number
6883- PASS or FAIL
6884- brief reason
6885
6886Then provide overall confidence (0.0 to 1.0) and whether it passes overall.
6887
6888Format:
68891. PASS/FAIL - reason
68902. PASS/FAIL - reason
6891...
6892CONFIDENCE: 0.X
6893OVERALL: PASS/FAIL"#,
6894 input, response, criteria_list
6895 );
6896
6897 let messages = vec![ChatMessage::user(&prompt)];
6898 let eval_response = self
6899 .observe_purpose(
6900 ObservationPurpose::ReflectionEvaluation,
6901 evaluator_llm.complete(&messages, None),
6902 )
6903 .await
6904 .map_err(|e| AgentError::LLM(format!("Evaluation failed: {}", e)))?;
6905
6906 let content = eval_response.content.to_uppercase();
6907 let llm_pass = content.contains("OVERALL: PASS");
6908
6909 let confidence = content
6910 .lines()
6911 .find(|l| l.contains("CONFIDENCE:"))
6912 .and_then(|l| {
6913 l.split(':')
6914 .nth(1)
6915 .and_then(|v| v.trim().parse::<f32>().ok())
6916 })
6917 .unwrap_or(if llm_pass { 0.8 } else { 0.4 });
6918
6919 let overall_pass = llm_pass && confidence >= config.pass_threshold;
6922
6923 let mut criteria_results = Vec::new();
6924 for (i, criterion) in criteria.iter().enumerate() {
6925 let line_marker = format!("{}.", i + 1);
6926 let passed = eval_response
6927 .content
6928 .lines()
6929 .find(|l| l.contains(&line_marker))
6930 .map(|l| l.to_uppercase().contains("PASS"))
6931 .unwrap_or(overall_pass);
6932
6933 if passed {
6934 criteria_results.push(CriterionResult::pass(criterion));
6935 } else {
6936 criteria_results.push(CriterionResult::fail(criterion, "Did not meet criterion"));
6937 }
6938 }
6939
6940 Ok(EvaluationResult::new(overall_pass, confidence).with_criteria(criteria_results))
6941 }
6942
6943 async fn process_input(&self, input: &str) -> Result<ProcessData> {
6945 if let Some(processor) = self.get_state_process_processor() {
6946 let purpose = observation_purpose_for_process(processor.input_purpose_hint());
6947 return self
6948 .observe_purpose(purpose, processor.process_input(input))
6949 .await;
6950 }
6951 if let Some(ref processor) = self.process_processor {
6952 let purpose = observation_purpose_for_process(processor.input_purpose_hint());
6953 self.observe_purpose(purpose, processor.process_input(input))
6954 .await
6955 } else {
6956 Ok(ProcessData::new(input))
6957 }
6958 }
6959
6960 async fn process_output(
6962 &self,
6963 output: &str,
6964 input_context: &std::collections::HashMap<String, serde_json::Value>,
6965 ) -> Result<ProcessData> {
6966 if let Some(processor) = self.get_state_process_processor() {
6967 let purpose = observation_purpose_for_process(processor.output_purpose_hint());
6968 return self
6969 .observe_purpose(purpose, processor.process_output(output, input_context))
6970 .await;
6971 }
6972 if let Some(ref processor) = self.process_processor {
6973 let purpose = observation_purpose_for_process(processor.output_purpose_hint());
6974 self.observe_purpose(purpose, processor.process_output(output, input_context))
6975 .await
6976 } else {
6977 Ok(ProcessData::new(output))
6978 }
6979 }
6980
6981 fn get_state_process_processor(&self) -> Option<ProcessProcessor> {
6983 let sm = self.state_machine.as_ref()?;
6984 let def = sm.current_definition()?;
6985 let config = def.process.as_ref()?;
6986 let mut processor = ProcessProcessor::new(config.clone());
6987 if let Some(ref registry) = Some(self.llm_registry.clone()) {
6988 processor = processor.with_llm_registry(registry.clone());
6989 }
6990 processor = processor.with_stage_observer(Arc::new(ObservabilityProcessStageObserver));
6991 Some(processor)
6992 }
6993
6994 async fn check_turn_timeout(&self) -> Result<()> {
6996 let Some(ref sm) = self.state_machine else {
6997 return Ok(());
6998 };
6999 let Some(timeout_state) = sm.check_timeout() else {
7000 return Ok(());
7001 };
7002 let claim_admission = self.disambiguation_admission.write().await;
7003 if sm.check_timeout().as_deref() != Some(timeout_state.as_str()) {
7004 return Ok(());
7005 }
7006 let Some(reservation) = self.reserve_state_transition() else {
7007 return Ok(());
7008 };
7009 let from_state = sm.current();
7010 let expected_state_generation = sm.generation();
7011 let expected_disambiguation_epoch = self.disambiguation_epoch.load(Ordering::SeqCst);
7012 let history_before = sm.history();
7013 drop(claim_admission);
7014
7015 self.execute_state_exit_actions(&from_state).await;
7016
7017 let admission = self.disambiguation_admission.write().await;
7018 if sm.current() != from_state
7019 || sm.generation() != expected_state_generation
7020 || self.disambiguation_epoch.load(Ordering::SeqCst) != expected_disambiguation_epoch
7021 || sm.check_timeout().as_deref() != Some(timeout_state.as_str())
7022 {
7023 return Ok(());
7024 }
7025 sm.transition_to(&timeout_state, "max_turns exceeded")?;
7026 self.invalidate_pending_confirmation("state_timeout").await;
7027 let entered = sm.current();
7028 let is_reentry = Self::state_was_previously_entered(&entered, &from_state, &history_before);
7029 drop(admission);
7030
7031 self.execute_state_enter_actions(&entered, is_reentry).await;
7032 drop(reservation);
7033 info!(to = %entered, "Timeout transition");
7034 Ok(())
7035 }
7036
7037 fn increment_turn(&self) {
7038 if let Some(ref sm) = self.state_machine {
7039 sm.increment_turn();
7040 }
7041 }
7042
7043 fn transitions_available_for_commit(&self) -> Option<(Vec<Transition>, String)> {
7044 let sm = self.state_machine.as_ref()?;
7045 let current = sm.current();
7046 let transitions: Vec<_> = sm
7047 .auto_transitions()
7048 .into_iter()
7049 .filter(|t| match t.cooldown_turns {
7050 Some(cd) if cd > 0 => {
7051 let resolved = sm.config().resolve_full_path(¤t, &t.to);
7052 !sm.is_on_cooldown(&resolved, cd)
7053 }
7054 _ => true,
7055 })
7056 .collect();
7057 Some((transitions, current))
7058 }
7059
7060 fn transition_reason(transition: &Transition) -> String {
7061 if transition.when.is_empty() {
7062 "guard condition met".to_string()
7063 } else {
7064 transition.when.clone()
7065 }
7066 }
7067
7068 fn build_transition_context(
7070 &self,
7071 user_message: &str,
7072 response: &str,
7073 current_state: &str,
7074 staged: Option<&HashMap<String, Value>>,
7075 ) -> TransitionContext {
7076 let context_map = staged
7077 .map(|writes| self.build_context_with_staged(writes))
7078 .unwrap_or_else(|| self.build_context_with_overlays());
7079 TransitionContext::new(user_message, response, current_state).with_context(context_map)
7080 }
7081
7082 async fn select_transition_candidate(
7084 &self,
7085 user_message: &str,
7086 response: &str,
7087 ) -> Result<Option<TransitionCandidate>> {
7088 let Some((transitions, current_state)) = self.transitions_available_for_commit() else {
7089 return Ok(None);
7090 };
7091 let transitions: Vec<Transition> = transitions
7092 .into_iter()
7093 .filter(|transition| matches!(transition.timing, TransitionTiming::PostResponse))
7094 .collect();
7095 if transitions.is_empty() {
7096 return Ok(None);
7097 }
7098 let Some(evaluator) = self.transition_evaluator.as_ref() else {
7099 return Ok(None);
7100 };
7101 let context = self.build_transition_context(user_message, response, ¤t_state, None);
7102 let selected = self
7103 .observe_purpose(
7104 ObservationPurpose::StateTransitionEvaluation,
7105 evaluator.select_transition(&transitions, &context),
7106 )
7107 .await?;
7108 Ok(selected.map(|index| {
7109 let transition = transitions[index].clone();
7110 TransitionCandidate::new(
7111 current_state,
7112 transition.clone(),
7113 Self::transition_reason(&transition),
7114 )
7115 }))
7116 }
7117
7118 fn select_deterministic_transition_candidate(
7120 &self,
7121 user_message: &str,
7122 current_state: &str,
7123 transitions: &[Transition],
7124 staged: &HashMap<String, Value>,
7125 ) -> Option<TransitionCandidate> {
7126 let context = self.build_transition_context(user_message, "", current_state, Some(staged));
7127
7128 for transition in transitions {
7129 if let Some(guard) = transition.guard.as_ref()
7130 && evaluate_guard(guard, &context)
7131 {
7132 return Some(TransitionCandidate::new(
7133 current_state,
7134 transition.clone(),
7135 Self::transition_reason(transition),
7136 ));
7137 }
7138 }
7139
7140 let resolved_intent = context
7141 .context
7142 .get("resolved_intent")
7143 .and_then(Value::as_str)
7144 .filter(|value| !value.is_empty());
7145 if let Some(resolved_intent) = resolved_intent {
7146 for transition in transitions {
7147 if transition.intent.as_deref() == Some(resolved_intent) {
7148 return Some(TransitionCandidate::new(
7149 current_state,
7150 transition.clone(),
7151 Self::transition_reason(transition),
7152 ));
7153 }
7154 }
7155 }
7156
7157 None
7158 }
7159
7160 async fn commit_transition_candidate(&self, candidate: &TransitionCandidate) -> Result<bool> {
7162 self.commit_transition_target(&candidate.from_state, candidate.target(), &candidate.reason)
7163 .await
7164 }
7165
7166 async fn approve_transition_target(&self, from_state: &str, target: &str) -> Result<bool> {
7168 let approved = self.check_state_hitl(Some(from_state), target).await?;
7169 if !approved {
7170 info!(to = %target, "State transition rejected by HITL");
7171 }
7172 Ok(approved)
7173 }
7174
7175 async fn apply_transition_target(
7177 &self,
7178 from_state: &str,
7179 target: &str,
7180 reason: &str,
7181 staged: Option<&HashMap<String, Value>>,
7182 ) -> Result<bool> {
7183 let Some(ref sm) = self.state_machine else {
7184 return Ok(false);
7185 };
7186 let claim_admission = self.disambiguation_admission.write().await;
7187 if sm.current() != from_state {
7188 return Ok(false);
7189 }
7190 let Some(reservation) = self.reserve_state_transition() else {
7191 return Ok(false);
7192 };
7193 let expected_state_generation = sm.generation();
7194 let expected_disambiguation_epoch = self.disambiguation_epoch.load(Ordering::SeqCst);
7195 let history_before = sm.history();
7196 drop(claim_admission);
7197
7198 self.execute_state_exit_actions(from_state).await;
7199
7200 let admission = self.disambiguation_admission.write().await;
7201 if sm.current() != from_state
7202 || sm.generation() != expected_state_generation
7203 || self.disambiguation_epoch.load(Ordering::SeqCst) != expected_disambiguation_epoch
7204 {
7205 return Ok(false);
7206 }
7207 sm.transition_to(target, reason)?;
7208 self.invalidate_pending_confirmation("state_transition")
7209 .await;
7210 sm.reset_no_transition();
7211 if let Some(staged) = staged {
7212 self.commit_staged_context_writes(staged);
7213 }
7214 let entered = sm.current();
7215 let is_reentry = Self::state_was_previously_entered(&entered, from_state, &history_before);
7216 drop(admission);
7217
7218 self.execute_state_enter_actions(&entered, is_reentry).await;
7219 drop(reservation);
7220 self.hooks
7221 .on_state_transition(Some(from_state), &entered, reason)
7222 .await;
7223 info!(from = %from_state, to = %entered, "State transition");
7224 Ok(true)
7225 }
7226
7227 async fn commit_transition_target(
7229 &self,
7230 from_state: &str,
7231 target: &str,
7232 reason: &str,
7233 ) -> Result<bool> {
7234 if !self.approve_transition_target(from_state, target).await? {
7235 return Ok(false);
7236 }
7237 self.apply_transition_target(from_state, target, reason, None)
7238 .await
7239 }
7240
7241 async fn apply_pre_response_transition_candidate(
7243 &self,
7244 candidate: &TransitionCandidate,
7245 staged: &HashMap<String, Value>,
7246 processed_input: &str,
7247 ) -> Result<bool> {
7248 self.commit_root_user_message(processed_input).await?;
7249 self.apply_transition_target(
7250 &candidate.from_state,
7251 candidate.target(),
7252 &candidate.reason,
7253 Some(staged),
7254 )
7255 .await
7256 }
7257
7258 async fn commit_pre_response_transition_candidate(
7260 &self,
7261 candidate: &TransitionCandidate,
7262 staged: &HashMap<String, Value>,
7263 processed_input: &str,
7264 ) -> Result<bool> {
7265 if !self
7266 .approve_transition_target(&candidate.from_state, candidate.target())
7267 .await?
7268 {
7269 return Ok(false);
7270 }
7271 self.apply_pre_response_transition_candidate(candidate, staged, processed_input)
7272 .await
7273 }
7274
7275 async fn handle_transition_miss(&self, current_state: &str) -> Result<bool> {
7277 let Some(ref sm) = self.state_machine else {
7278 return Ok(false);
7279 };
7280 sm.increment_no_transition();
7281 let Some(fallback) = sm.check_fallback() else {
7282 return Ok(false);
7283 };
7284 self.commit_transition_target(current_state, &fallback, "fallback after no transitions")
7285 .await
7286 }
7287
7288 async fn evaluate_transitions(&self, user_message: &str, response: &str) -> Result<bool> {
7290 let Some((transitions, current_state)) = self.transitions_available_for_commit() else {
7291 return Ok(false);
7292 };
7293 if transitions.is_empty() {
7294 return Ok(false);
7295 }
7296 if let Some(candidate) = self
7297 .select_transition_candidate(user_message, response)
7298 .await?
7299 {
7300 return self.commit_transition_candidate(&candidate).await;
7301 }
7302 self.handle_transition_miss(¤t_state).await
7303 }
7304
7305 async fn try_pre_response_transition(
7307 &self,
7308 processed_input: &str,
7309 ) -> Result<Option<AgentResponse>> {
7310 let optimization = &self.runtime_config.optimization;
7311 if !optimization.enabled || !optimization.pre_response_deterministic_transitions {
7312 return Ok(None);
7313 }
7314 let Some((transitions, current_state)) = self.transitions_available_for_commit() else {
7315 return Ok(None);
7316 };
7317 let eligible: Vec<Transition> = transitions
7318 .into_iter()
7319 .filter(|transition| !transition.requires_response)
7320 .filter(|transition| matches!(transition.timing, TransitionTiming::PreResponse))
7321 .collect();
7322 if eligible.is_empty() {
7323 return Ok(None);
7324 }
7325
7326 let empty_staged = HashMap::new();
7327 let mut extracted_staged: Option<HashMap<String, Value>> = None;
7328 let mut selected: Option<(TransitionCandidate, HashMap<String, Value>)> = None;
7329
7330 for transition in &eligible {
7331 let use_extractors = optimization.pre_response_extractors || transition.run_extractors;
7332 let staged_for_eval = if use_extractors {
7333 if extracted_staged.is_none() {
7334 extracted_staged =
7335 Some(self.run_context_extractors_staged(processed_input).await);
7336 }
7337 extracted_staged.as_ref().unwrap_or(&empty_staged)
7338 } else {
7339 &empty_staged
7340 };
7341
7342 if let Some(candidate) = self.select_deterministic_transition_candidate(
7343 processed_input,
7344 ¤t_state,
7345 std::slice::from_ref(transition),
7346 staged_for_eval,
7347 ) {
7348 let staged_for_commit = if use_extractors {
7349 staged_for_eval.clone()
7350 } else {
7351 HashMap::new()
7352 };
7353 selected = Some((candidate, staged_for_commit));
7354 break;
7355 }
7356 }
7357
7358 let Some((candidate, staged)) = selected else {
7359 return Ok(None);
7360 };
7361
7362 if !self
7363 .commit_pre_response_transition_candidate(&candidate, &staged, processed_input)
7364 .await?
7365 {
7366 return Ok(None);
7367 }
7368 self.redispatch_current_state(processed_input)
7369 .await
7370 .map(Some)
7371 }
7372
7373 async fn try_speculative_branches(
7378 &self,
7379 processed_input: &str,
7380 input_context: &HashMap<String, Value>,
7381 ) -> Result<Option<AgentResponse>> {
7382 let optimization = &self.runtime_config.optimization;
7383 if !optimization.enabled {
7384 return Ok(None);
7385 }
7386
7387 let effective_reasoning_mode = self.get_effective_reasoning_config().mode.clone();
7388 if !matches!(
7389 effective_reasoning_mode,
7390 ReasoningMode::None | ReasoningMode::Auto
7391 ) {
7392 return Ok(None);
7393 }
7394
7395 let mut transition_enabled =
7396 optimization.speculative_state_transitions && self.has_parallel_transition_candidates();
7397 let mut skill_enabled = optimization.speculative_skill_routing
7398 && self.skill_router.is_some()
7399 && self.pending_skill_id.read().is_none();
7400 let mut reasoning_enabled = optimization.speculative_reasoning_auto
7401 && matches!(effective_reasoning_mode, ReasoningMode::Auto);
7402
7403 if matches!(effective_reasoning_mode, ReasoningMode::Auto)
7404 && (!reasoning_enabled || optimization.max_speculative_llm_calls_per_turn < 2)
7405 {
7406 return Ok(None);
7407 }
7408
7409 if !transition_enabled && !skill_enabled && !reasoning_enabled {
7410 return Ok(None);
7411 }
7412
7413 let mut optional_slots = optimization.max_parallel_runtime_tasks.saturating_sub(1);
7414 let mut speculative_call_slots = optimization
7415 .max_speculative_llm_calls_per_turn
7416 .saturating_sub(1);
7417 if reasoning_enabled {
7418 if optional_slots == 0 || speculative_call_slots == 0 {
7419 return Ok(None);
7420 }
7421 optional_slots -= 1;
7422 speculative_call_slots -= 1;
7423 }
7424 if transition_enabled {
7425 if optional_slots == 0 {
7426 transition_enabled = false;
7427 } else {
7428 optional_slots -= 1;
7429 }
7430 }
7431 if skill_enabled && (optional_slots == 0 || speculative_call_slots == 0) {
7432 skill_enabled = false;
7433 }
7434
7435 if !transition_enabled && !skill_enabled && !reasoning_enabled {
7436 return Ok(None);
7437 }
7438
7439 let main_kind = if transition_enabled {
7440 RuntimeOptimizationKind::ParallelStateTransition
7441 } else if skill_enabled {
7442 RuntimeOptimizationKind::SpeculativeSkillRouting
7443 } else {
7444 RuntimeOptimizationKind::SpeculativeReasoningAuto
7445 };
7446 if !self.reserve_active_speculative_llm_call(main_kind) {
7447 return Ok(None);
7448 }
7449
7450 let mut branch_set = ScheduledBranchSet::new(optimization.max_parallel_runtime_tasks)?;
7451 let main_branch = RuntimeBranch::new(
7452 RuntimeTaskPurpose::MainResponse,
7453 main_kind,
7454 RuntimeTaskPriority::Normal,
7455 RuntimeCommitBehavior::FinalResponse,
7456 );
7457 let transition_branch = RuntimeBranch::new(
7458 RuntimeTaskPurpose::StateTransition,
7459 RuntimeOptimizationKind::ParallelStateTransition,
7460 RuntimeTaskPriority::Critical,
7461 RuntimeCommitBehavior::TransitionDecision,
7462 );
7463 let skill_branch = RuntimeBranch::new(
7464 RuntimeTaskPurpose::SkillRouting,
7465 RuntimeOptimizationKind::SpeculativeSkillRouting,
7466 RuntimeTaskPriority::High,
7467 RuntimeCommitBehavior::SkillSelection,
7468 );
7469 let reasoning_branch = RuntimeBranch::new(
7470 RuntimeTaskPurpose::ReasoningJudge,
7471 RuntimeOptimizationKind::SpeculativeReasoningAuto,
7472 RuntimeTaskPriority::Normal,
7473 RuntimeCommitBehavior::ReasoningDecision,
7474 );
7475 let main_id = main_branch.branch_id();
7476 let transition_id = transition_branch.branch_id();
7477 let skill_id = skill_branch.branch_id();
7478 let reasoning_id = reasoning_branch.branch_id();
7479
7480 let main_id_for_future = main_id.clone();
7481 if !branch_set.schedule(
7482 main_branch,
7483 Box::pin(async move {
7484 match crate::optimization::observability::with_branch_observation(
7485 &main_id_for_future,
7486 main_kind,
7487 RuntimeCommitBehavior::FinalResponse,
7488 self.generate_main_response_draft(processed_input, &ReasoningMode::None),
7489 )
7490 .await
7491 {
7492 Ok(draft) => RuntimeBranchResult::MainDraft(draft),
7493 Err(error) => RuntimeBranchResult::Failed(error),
7494 }
7495 }),
7496 ) {
7497 return Ok(None);
7498 }
7499
7500 if transition_enabled {
7501 let transition_id_for_future = transition_id.clone();
7502 if !branch_set.schedule(
7503 transition_branch,
7504 Box::pin(async move {
7505 match crate::optimization::observability::with_branch_observation(
7506 &transition_id_for_future,
7507 RuntimeOptimizationKind::ParallelStateTransition,
7508 RuntimeCommitBehavior::TransitionDecision,
7509 self.select_parallel_transition_candidate(processed_input),
7510 )
7511 .await
7512 {
7513 Ok(ParallelTransitionSelection::Candidate(candidate)) => {
7514 RuntimeBranchResult::Transition(Some(candidate))
7515 }
7516 Ok(ParallelTransitionSelection::NoMatch) => {
7517 RuntimeBranchResult::Transition(None)
7518 }
7519 Ok(ParallelTransitionSelection::ReservationExhausted) => {
7520 RuntimeBranchResult::Cancelled
7521 }
7522 Err(error) => RuntimeBranchResult::Failed(error),
7523 }
7524 }),
7525 ) {
7526 transition_enabled = false;
7527 }
7528 }
7529
7530 if skill_enabled {
7531 let skill_id_for_future = skill_id.clone();
7532 if !branch_set.schedule(
7533 skill_branch,
7534 Box::pin(async move {
7535 if !self.reserve_active_speculative_llm_call(
7536 RuntimeOptimizationKind::SpeculativeSkillRouting,
7537 ) {
7538 return RuntimeBranchResult::Cancelled;
7539 }
7540 match crate::optimization::observability::with_branch_observation(
7541 &skill_id_for_future,
7542 RuntimeOptimizationKind::SpeculativeSkillRouting,
7543 RuntimeCommitBehavior::SkillSelection,
7544 self.select_skill_candidate(processed_input),
7545 )
7546 .await
7547 {
7548 Ok(candidate) => RuntimeBranchResult::Skill(candidate),
7549 Err(error) => RuntimeBranchResult::Failed(error),
7550 }
7551 }),
7552 ) {
7553 skill_enabled = false;
7554 }
7555 }
7556
7557 if reasoning_enabled {
7558 let reasoning_id_for_future = reasoning_id.clone();
7559 if !branch_set.schedule(
7560 reasoning_branch,
7561 Box::pin(async move {
7562 if !self.reserve_active_speculative_llm_call(
7563 RuntimeOptimizationKind::SpeculativeReasoningAuto,
7564 ) {
7565 return RuntimeBranchResult::Cancelled;
7566 }
7567 match crate::optimization::observability::with_branch_observation(
7568 &reasoning_id_for_future,
7569 RuntimeOptimizationKind::SpeculativeReasoningAuto,
7570 RuntimeCommitBehavior::ReasoningDecision,
7571 self.determine_reasoning_mode_strict(processed_input),
7572 )
7573 .await
7574 {
7575 Ok(mode) => RuntimeBranchResult::Reasoning(mode),
7576 Err(error) => RuntimeBranchResult::Failed(error),
7577 }
7578 }),
7579 ) {
7580 reasoning_enabled = false;
7581 }
7582 }
7583
7584 if matches!(effective_reasoning_mode, ReasoningMode::Auto) && !reasoning_enabled {
7585 self.finalize_pending_branches(branch_set.cancel_pending());
7586 return Ok(None);
7587 }
7588
7589 if !transition_enabled && !skill_enabled && !reasoning_enabled {
7590 self.finalize_pending_branches(branch_set.cancel_pending());
7591 return Ok(None);
7592 }
7593
7594 let mut main_pending = true;
7595 let mut skill_pending = skill_enabled;
7596 let mut reasoning_pending = reasoning_enabled;
7597 let mut transition_finalized = !transition_enabled;
7598 let mut skill_finalized = !skill_enabled && self.skill_router.is_none();
7601 let mut reasoning_finalized = !reasoning_enabled;
7602 let mut main_result: Option<Result<MainResponseDraft>> = None;
7603 let mut transition_candidate: Option<TransitionCandidate> = None;
7604 let mut skill_candidate: Option<SkillCandidate> = None;
7605 let mut reasoning_decision: Option<ReasoningMode> = None;
7606 let mut transition_fallback_required = false;
7607 let mut skill_fallback_required = false;
7608 let mut reasoning_fallback_required = false;
7609
7610 loop {
7611 if let Some(candidate) = transition_candidate.take() {
7612 if self
7613 .approve_transition_target(&candidate.from_state, candidate.target())
7614 .await?
7615 {
7616 self.finalize_pending_branches(branch_set.cancel_pending());
7618 if !main_pending {
7619 self.finalize_branch_loss(
7620 &main_id,
7621 main_kind,
7622 RuntimeCommitBehavior::FinalResponse,
7623 false,
7624 main_result.as_ref().map(|result| result.is_err()),
7625 );
7626 }
7627 if skill_enabled && !skill_pending {
7628 self.finalize_branch_loss(
7629 &skill_id,
7630 RuntimeOptimizationKind::SpeculativeSkillRouting,
7631 RuntimeCommitBehavior::SkillSelection,
7632 false,
7633 Some(false),
7634 );
7635 }
7636 if reasoning_enabled && !reasoning_pending {
7637 self.finalize_branch_loss(
7638 &reasoning_id,
7639 RuntimeOptimizationKind::SpeculativeReasoningAuto,
7640 RuntimeCommitBehavior::ReasoningDecision,
7641 false,
7642 Some(false),
7643 );
7644 }
7645 if !self
7646 .apply_pre_response_transition_candidate(
7647 &candidate,
7648 &HashMap::new(),
7649 processed_input,
7650 )
7651 .await?
7652 {
7653 self.finalize_optional_branch(
7654 &transition_id,
7655 RuntimeOptimizationKind::ParallelStateTransition,
7656 RuntimeCommitBehavior::TransitionDecision,
7657 "discarded",
7658 false,
7659 );
7660 return Ok(None);
7661 }
7662 self.finalize_optional_branch(
7663 &transition_id,
7664 RuntimeOptimizationKind::ParallelStateTransition,
7665 RuntimeCommitBehavior::TransitionDecision,
7666 "committed",
7667 true,
7668 );
7669 return self
7670 .redispatch_current_state(processed_input)
7671 .await
7672 .map(Some);
7673 }
7674 self.finalize_optional_branch(
7675 &transition_id,
7676 RuntimeOptimizationKind::ParallelStateTransition,
7677 RuntimeCommitBehavior::TransitionDecision,
7678 "discarded",
7679 false,
7680 );
7681 transition_finalized = true;
7682 }
7683
7684 if transition_finalized
7693 && !skill_finalized
7694 && !skill_enabled
7695 && self.skill_router.is_some()
7696 {
7697 match self.select_skill_candidate(processed_input).await {
7698 Ok(Some(candidate)) => skill_candidate = Some(candidate),
7699 Ok(None) => {}
7700 Err(error) => {
7701 self.finalize_pending_branches(branch_set.cancel_pending());
7703 return Err(error);
7704 }
7705 }
7706 skill_finalized = true;
7707 }
7708
7709 if transition_finalized && skill_candidate.is_some() {
7710 let candidate = skill_candidate.take().unwrap();
7711 if skill_enabled {
7713 self.finalize_optional_branch(
7714 &skill_id,
7715 RuntimeOptimizationKind::SpeculativeSkillRouting,
7716 RuntimeCommitBehavior::SkillSelection,
7717 "committed",
7718 true,
7719 );
7720 }
7721 if !main_pending {
7722 self.finalize_branch_loss(
7723 &main_id,
7724 main_kind,
7725 RuntimeCommitBehavior::FinalResponse,
7726 false,
7727 main_result.as_ref().map(|result| result.is_err()),
7728 );
7729 }
7730 if reasoning_enabled && !reasoning_pending {
7731 self.finalize_branch_loss(
7732 &reasoning_id,
7733 RuntimeOptimizationKind::SpeculativeReasoningAuto,
7734 RuntimeCommitBehavior::ReasoningDecision,
7735 false,
7736 Some(false),
7737 );
7738 }
7739 self.finalize_pending_branches(branch_set.cancel_pending());
7740 return self
7741 .commit_winning_skill_candidate(candidate, processed_input, input_context)
7742 .await;
7743 }
7744
7745 if transition_finalized
7746 && skill_finalized
7747 && let Some(reasoning_mode) = reasoning_decision.take()
7748 {
7749 if !matches!(reasoning_mode, ReasoningMode::None) {
7750 self.finalize_optional_branch(
7751 &reasoning_id,
7752 RuntimeOptimizationKind::SpeculativeReasoningAuto,
7753 RuntimeCommitBehavior::ReasoningDecision,
7754 "committed",
7755 true,
7756 );
7757 if !main_pending {
7758 self.finalize_branch_loss(
7759 &main_id,
7760 main_kind,
7761 RuntimeCommitBehavior::FinalResponse,
7762 false,
7763 main_result.as_ref().map(|result| result.is_err()),
7764 );
7765 }
7766 self.finalize_pending_branches(branch_set.cancel_pending());
7767 self.commit_root_user_message(processed_input).await?;
7768 return if matches!(reasoning_mode, ReasoningMode::PlanAndExecute) {
7769 self.handle_plan_and_execute(processed_input, input_context, true)
7770 .await
7771 .map(Some)
7772 } else {
7773 self.run_committed_response_loop_with_reasoning(
7774 processed_input,
7775 input_context,
7776 reasoning_mode,
7777 true,
7778 )
7779 .await
7780 .map(Some)
7781 };
7782 }
7783 self.finalize_optional_branch(
7784 &reasoning_id,
7785 RuntimeOptimizationKind::SpeculativeReasoningAuto,
7786 RuntimeCommitBehavior::ReasoningDecision,
7787 "committed",
7788 true,
7789 );
7790 reasoning_finalized = true;
7791 }
7792
7793 if transition_finalized && skill_finalized && reasoning_finalized {
7794 if transition_fallback_required
7795 || skill_fallback_required
7796 || reasoning_fallback_required
7797 {
7798 if !main_pending {
7799 self.finalize_branch_loss(
7800 &main_id,
7801 main_kind,
7802 RuntimeCommitBehavior::FinalResponse,
7803 false,
7804 main_result.as_ref().map(|result| result.is_err()),
7805 );
7806 }
7807 self.finalize_pending_branches(branch_set.cancel_pending());
7808 return Ok(None);
7809 }
7810
7811 if let Some(result) = main_result.take() {
7812 let draft = match result {
7813 Ok(draft) => draft,
7814 Err(error) => {
7815 self.finalize_optional_branch(
7816 &main_id,
7817 main_kind,
7818 RuntimeCommitBehavior::FinalResponse,
7819 "failed",
7820 false,
7821 );
7822 self.finalize_pending_branches(branch_set.cancel_pending());
7823 return Err(error);
7824 }
7825 };
7826 self.finalize_optional_branch(
7827 &main_id,
7828 main_kind,
7829 RuntimeCommitBehavior::FinalResponse,
7830 "committed",
7831 true,
7832 );
7833 self.finalize_pending_branches(branch_set.cancel_pending());
7834 return self
7835 .commit_main_response_draft(
7836 processed_input,
7837 input_context,
7838 draft,
7839 ReasoningMode::None,
7840 reasoning_enabled,
7841 )
7842 .await
7843 .map(Some);
7844 }
7845 }
7846
7847 if branch_set.is_empty() {
7848 return Ok(None);
7849 }
7850
7851 let Some(outcome) = branch_set.next_completed().await else {
7852 return Ok(None);
7853 };
7854 let branch_id = outcome.branch.branch_id();
7855 match outcome.result {
7856 RuntimeBranchResult::MainDraft(draft) => {
7857 main_pending = false;
7858 main_result = Some(Ok(draft));
7859 }
7860 RuntimeBranchResult::Transition(candidate) => {
7861 if let Some(candidate) = candidate {
7862 transition_candidate = Some(candidate);
7863 } else {
7864 self.finalize_optional_branch(
7865 &transition_id,
7866 RuntimeOptimizationKind::ParallelStateTransition,
7867 RuntimeCommitBehavior::TransitionDecision,
7868 "discarded",
7869 false,
7870 );
7871 transition_finalized = true;
7872 }
7873 }
7874 RuntimeBranchResult::Skill(candidate) => {
7875 skill_pending = false;
7876 if let Some(candidate) = candidate {
7877 skill_candidate = Some(candidate);
7878 } else {
7879 self.finalize_optional_branch(
7880 &skill_id,
7881 RuntimeOptimizationKind::SpeculativeSkillRouting,
7882 RuntimeCommitBehavior::SkillSelection,
7883 "discarded",
7884 false,
7885 );
7886 skill_finalized = true;
7887 }
7888 }
7889 RuntimeBranchResult::Reasoning(mode) => {
7890 reasoning_pending = false;
7891 reasoning_decision = Some(mode);
7892 }
7893 RuntimeBranchResult::Failed(error) => {
7894 if branch_id == main_id {
7895 main_pending = false;
7896 main_result = Some(Err(error));
7897 } else if branch_id == transition_id {
7898 self.finalize_optional_branch(
7899 &transition_id,
7900 RuntimeOptimizationKind::ParallelStateTransition,
7901 RuntimeCommitBehavior::TransitionDecision,
7902 "failed",
7903 false,
7904 );
7905 transition_finalized = true;
7906 } else if branch_id == skill_id {
7907 skill_pending = false;
7908 self.finalize_optional_branch(
7909 &skill_id,
7910 RuntimeOptimizationKind::SpeculativeSkillRouting,
7911 RuntimeCommitBehavior::SkillSelection,
7912 "failed",
7913 false,
7914 );
7915 skill_finalized = true;
7916 } else if branch_id == reasoning_id {
7917 reasoning_pending = false;
7918 self.finalize_optional_branch(
7919 &reasoning_id,
7920 RuntimeOptimizationKind::SpeculativeReasoningAuto,
7921 RuntimeCommitBehavior::ReasoningDecision,
7922 "failed",
7923 false,
7924 );
7925 reasoning_finalized = true;
7926 }
7927 }
7928 RuntimeBranchResult::Cancelled => {
7929 self.finalize_optional_branch(
7930 &branch_id,
7931 outcome.branch.optimization,
7932 outcome.branch.commit_behavior,
7933 "cancelled",
7934 false,
7935 );
7936 if branch_id == main_id {
7937 main_pending = false;
7938 main_result =
7939 Some(Err(AgentError::Other("main branch cancelled".to_string())));
7940 } else if branch_id == transition_id {
7941 transition_finalized = true;
7942 transition_fallback_required = true;
7943 } else if branch_id == skill_id {
7944 skill_pending = false;
7945 skill_finalized = true;
7946 skill_fallback_required = true;
7947 } else if branch_id == reasoning_id {
7948 reasoning_pending = false;
7949 reasoning_finalized = true;
7950 reasoning_fallback_required = true;
7951 }
7952 }
7953 }
7954 }
7955 }
7956
7957 fn finalize_pending_branches(&self, branches: Vec<RuntimeBranch>) {
7958 for branch in branches {
7959 self.finalize_optional_branch(
7960 &branch.branch_id(),
7961 branch.optimization,
7962 branch.commit_behavior,
7963 "cancelled",
7964 false,
7965 );
7966 }
7967 }
7968
7969 fn finalize_branch_loss(
7974 &self,
7975 branch_id: &str,
7976 optimization: RuntimeOptimizationKind,
7977 commit_behavior: RuntimeCommitBehavior,
7978 pending: bool,
7979 completed_failed: Option<bool>,
7980 ) {
7981 let status = if pending {
7982 "cancelled"
7983 } else if completed_failed.unwrap_or(false) {
7984 "failed"
7985 } else {
7986 "discarded"
7987 };
7988 self.finalize_optional_branch(branch_id, optimization, commit_behavior, status, false);
7989 }
7990
7991 fn finalize_optional_branch(
7996 &self,
7997 branch_id: &str,
7998 optimization: RuntimeOptimizationKind,
7999 commit_behavior: RuntimeCommitBehavior,
8000 status: &str,
8001 winner: bool,
8002 ) {
8003 crate::optimization::observability::finalize_branch(
8004 self.observability_manager.as_ref(),
8005 branch_id,
8006 status,
8007 winner,
8008 optimization,
8009 commit_behavior,
8010 );
8011 }
8012
8013 fn has_parallel_transition_candidates(&self) -> bool {
8018 self.transitions_available_for_commit()
8019 .map(|(transitions, _)| {
8020 transitions
8021 .iter()
8022 .any(|transition| matches!(transition.timing, TransitionTiming::Parallel))
8023 })
8024 .unwrap_or(false)
8025 }
8026
8027 async fn select_parallel_transition_candidate(
8032 &self,
8033 processed_input: &str,
8034 ) -> Result<ParallelTransitionSelection> {
8035 let Some((transitions, current_state)) = self.transitions_available_for_commit() else {
8036 return Ok(ParallelTransitionSelection::NoMatch);
8037 };
8038 let parallel: Vec<Transition> = transitions
8039 .into_iter()
8040 .filter(|transition| matches!(transition.timing, TransitionTiming::Parallel))
8041 .filter(|transition| !transition.requires_response)
8042 .collect();
8043 if parallel.is_empty() {
8044 return Ok(ParallelTransitionSelection::NoMatch);
8045 }
8046 let empty_staged = HashMap::new();
8047 if let Some(candidate) = self.select_deterministic_transition_candidate(
8048 processed_input,
8049 ¤t_state,
8050 ¶llel,
8051 &empty_staged,
8052 ) {
8053 return Ok(ParallelTransitionSelection::Candidate(candidate));
8054 }
8055 let when_transitions: Vec<(usize, &Transition)> = parallel
8056 .iter()
8057 .enumerate()
8058 .filter(|(_, transition)| !transition.when.trim().is_empty())
8059 .collect();
8060 if when_transitions.is_empty() {
8061 return Ok(ParallelTransitionSelection::NoMatch);
8062 }
8063 let llm = self
8064 .llm_registry
8065 .router()
8066 .or_else(|_| self.llm_registry.default())
8067 .map_err(|e| AgentError::Config(e.to_string()))?;
8068 let conditions = when_transitions
8069 .iter()
8070 .enumerate()
8071 .map(|(display_idx, (_, transition))| {
8072 format!("{}. {}", display_idx + 1, transition.when)
8073 })
8074 .collect::<Vec<_>>()
8075 .join("\n");
8076 if !self
8077 .reserve_active_speculative_llm_call(RuntimeOptimizationKind::ParallelStateTransition)
8078 {
8079 return Ok(ParallelTransitionSelection::ReservationExhausted);
8080 }
8081 let context_preview = self.branch_context_preview();
8082 let prompt = format!(
8083 "Based only on the current user message and context, which transition condition is met?\n\nCurrent state: {}\nUser message: {}\nContext:\n{}\n\nConditions:\n{}\n0. None of the above\n\nReply with ONLY the number (0-{}).",
8084 current_state,
8085 processed_input,
8086 context_preview,
8087 conditions,
8088 when_transitions.len()
8089 );
8090 let response = self
8091 .observe_purpose(
8092 ObservationPurpose::StateTransitionEvaluation,
8093 llm.complete(&[ChatMessage::user(prompt)], None),
8094 )
8095 .await
8096 .map_err(|e| AgentError::LLM(e.to_string()))?;
8097 let choice = response.content.trim().parse::<usize>().unwrap_or(0);
8098 if choice == 0 || choice > when_transitions.len() {
8099 return Ok(ParallelTransitionSelection::NoMatch);
8100 }
8101 let transition = when_transitions[choice - 1].1.clone();
8102 Ok(ParallelTransitionSelection::Candidate(
8103 TransitionCandidate::new(
8104 current_state,
8105 transition.clone(),
8106 Self::transition_reason(&transition),
8107 ),
8108 ))
8109 }
8110
8111 async fn redispatch_current_state(&self, processed_input: &str) -> Result<AgentResponse> {
8113 const MAX_REDISPATCH_DEPTH: u32 = 3;
8114 let current_depth = *self.redispatch_depth.read();
8115 if current_depth >= MAX_REDISPATCH_DEPTH {
8116 warn!(depth = current_depth, "Re-dispatch depth limit reached");
8117 let response = AgentResponse::new("");
8118 self.finish_turn_if_root(&response).await?;
8119 return Ok(response);
8120 }
8121 *self.redispatch_depth.write() += 1;
8122 if let Some(context) = self.active_turn_context.write().as_mut() {
8123 context.enter_redispatch();
8124 }
8125 let result = Box::pin(self.run_loop_internal(processed_input)).await;
8126 *self.redispatch_depth.write() -= 1;
8127 if let Some(context) = self.active_turn_context.write().as_mut() {
8128 context.exit_redispatch();
8129 }
8130 let response = result?;
8131 self.finish_turn_if_root(&response).await?;
8132 Ok(response)
8133 }
8134
8135 async fn finish_turn_if_root(&self, response: &AgentResponse) -> Result<()> {
8137 if *self.redispatch_depth.read() == 0 {
8138 self.post_turn_session_lifecycle().await?;
8139 if let Some(context) = self.active_turn_context.write().as_mut() {
8140 context.mark_post_turn_lifecycle_completed();
8141 }
8142 self.hooks.on_response(response).await;
8143 self.end_root_turn();
8144 }
8145 Ok(())
8146 }
8147
8148 async fn execute_state_exit_actions(&self, state_path: &str) {
8150 if let Some(ref sm) = self.state_machine
8151 && let Some(def) = sm.get_definition(state_path)
8152 && !def.on_exit.is_empty()
8153 {
8154 debug!(state = %state_path, count = def.on_exit.len(), "Executing on_exit actions");
8155 self.execute_state_actions(&def.on_exit).await;
8156 }
8157 }
8158
8159 fn state_was_previously_entered(
8161 state_path: &str,
8162 from_state: &str,
8163 history_before: &[StateTransitionEvent],
8164 ) -> bool {
8165 state_path == from_state
8166 || history_before
8167 .iter()
8168 .any(|event| event.from == state_path || event.to == state_path)
8169 }
8170
8171 async fn execute_state_enter_actions(&self, state_path: &str, is_reentry: bool) {
8173 if let Some(ref sm) = self.state_machine
8174 && let Some(def) = sm.get_definition(state_path)
8175 {
8176 if is_reentry && !def.on_reenter.is_empty() {
8177 debug!(state = %state_path, count = def.on_reenter.len(), "Executing on_reenter actions");
8178 self.execute_state_actions(&def.on_reenter).await;
8179 } else if !def.on_enter.is_empty() {
8180 debug!(state = %state_path, count = def.on_enter.len(), "Executing on_enter actions");
8181 self.execute_state_actions(&def.on_enter).await;
8182 }
8183 }
8184 }
8185
8186 async fn execute_state_actions(&self, actions: &[StateAction]) {
8188 for (action_index, action) in actions.iter().enumerate() {
8189 match action {
8190 StateAction::Tool { tool, args } => {
8191 let raw_args = args.clone().unwrap_or(Value::Object(Default::default()));
8192 let args_value = self.render_action_args(&raw_args);
8193 let state = self.state_machine.as_ref().map(|sm| sm.current());
8194 let request = ToolExecutionRequest::new(
8195 uuid::Uuid::new_v4().to_string(),
8196 tool.clone(),
8197 args_value,
8198 ToolCallSource::StateAction {
8199 state,
8200 action_index,
8201 },
8202 );
8203 match self.execute_tool_record(request).await {
8204 Ok(record) if record.success => {
8205 debug!(tool = %record.canonical_id, "State action: tool executed");
8206 let _ = self.context_manager.set(
8207 "last_tool_result",
8208 serde_json::Value::String(record.model_output_string()),
8209 );
8210 let _ = self.context_manager.set(
8211 "last_tool_record",
8212 serde_json::to_value(record).unwrap_or(Value::Null),
8213 );
8214 }
8215 Ok(record) => {
8216 warn!(tool = %record.canonical_id, error = %record.output, "State action: tool failed");
8217 }
8218 Err(e) => {
8219 warn!(tool = %tool, error = %e, "State action: tool failed")
8220 }
8221 }
8222 }
8223 StateAction::Skill { skill } => {
8224 if let Some(ref executor) = self.skill_executor {
8225 if let Some(def) = self.skills.iter().find(|s| s.id == *skill) {
8226 match executor
8227 .execute_with_invoker(def, "", serde_json::json!({}), self)
8228 .await
8229 {
8230 Ok(_) => debug!(skill = %skill, "State action: skill executed"),
8231 Err(e) => {
8232 warn!(skill = %skill, error = %e, "State action: skill failed")
8233 }
8234 }
8235 } else {
8236 warn!(skill = %skill, "State action: skill not found");
8237 }
8238 }
8239 }
8240 StateAction::SetContext { set_context } => {
8241 for (key, value) in set_context {
8242 if let Err(e) = self.context_manager.set(key, value.clone()) {
8243 warn!(key = %key, error = %e, "State action: set_context failed");
8244 } else {
8245 debug!(key = %key, "State action: context set");
8246 }
8247 }
8248 }
8249 StateAction::Prompt {
8250 prompt,
8251 llm,
8252 store_as,
8253 } => {
8254 let llm_result = if let Some(alias) = llm {
8255 self.llm_registry.get(alias)
8256 } else {
8257 self.llm_registry.default()
8258 };
8259 match llm_result {
8260 Ok(llm_provider) => {
8261 let context = self.build_context_with_overlays();
8263 let rendered_prompt = self
8264 .template_renderer
8265 .render(prompt, &context)
8266 .unwrap_or_else(|_| prompt.clone());
8267 let recent =
8268 self.memory.get_messages(Some(5)).await.unwrap_or_default();
8269 let mut messages: Vec<ChatMessage> = recent;
8270 messages.push(ChatMessage::user(&rendered_prompt));
8271 match self
8272 .observe_purpose(
8273 ObservationPurpose::StateAction,
8274 llm_provider.complete(&messages, None),
8275 )
8276 .await
8277 {
8278 Ok(response) => {
8279 if let Some(key) = store_as {
8280 let _ = self
8281 .context_manager
8282 .set(key, Value::String(response.content));
8283 debug!(key = %key, "State action: prompt result stored");
8284 }
8285 }
8286 Err(e) => {
8287 warn!(error = %e, "State action: prompt LLM call failed");
8288 }
8289 }
8290 }
8291 Err(e) => {
8292 warn!(error = %e, "State action: LLM not found for prompt");
8293 }
8294 }
8295 }
8296 }
8297 }
8298 }
8299
8300 async fn run_context_extractors_staged(&self, user_message: &str) -> HashMap<String, Value> {
8301 let extractors = match &self.state_machine {
8302 Some(sm) => match sm.current_definition() {
8303 Some(def) if !def.extract.is_empty() => def.extract.clone(),
8304 _ => return HashMap::new(),
8305 },
8306 None => return HashMap::new(),
8307 };
8308
8309 let mut staged = HashMap::new();
8310 for extractor in &extractors {
8311 let prompt = if let Some(ref custom) = extractor.llm_extract {
8312 format!(
8313 "User message:\n\"{}\"\n\nInstruction:\n{}",
8314 user_message, custom
8315 )
8316 } else if let Some(ref desc) = extractor.description {
8317 format!(
8318 "From the following message, extract: {}\n\n\
8319 Message: \"{}\"\n\n\
8320 If the information is present, return ONLY the extracted value.\n\
8321 If NOT present, return exactly: __NONE__",
8322 desc, user_message
8323 )
8324 } else {
8325 continue;
8326 };
8327
8328 let llm = match self
8329 .llm_registry
8330 .get(&extractor.llm)
8331 .or_else(|_| self.llm_registry.get("router"))
8332 .or_else(|_| self.llm_registry.get("default"))
8333 {
8334 Ok(llm) => llm,
8335 Err(e) => {
8336 warn!(key = %extractor.key, error = %e, "Extractor LLM not found");
8337 continue;
8338 }
8339 };
8340
8341 let messages = vec![ChatMessage::user(&prompt)];
8342 match self
8343 .observe_purpose(
8344 ObservationPurpose::ContextExtraction,
8345 llm.complete(&messages, None),
8346 )
8347 .await
8348 {
8349 Ok(response) => {
8350 let value = response.content.trim().to_string();
8351 if value != "__NONE__" && !value.is_empty() {
8352 staged.insert(
8353 extractor.key.clone(),
8354 serde_json::Value::String(value.clone()),
8355 );
8356 debug!(key = %extractor.key, value = %value, "Context extracted");
8357 } else if extractor.required {
8358 warn!(key = %extractor.key, "Required extraction returned no value");
8359 }
8360 }
8361 Err(e) => {
8362 warn!(key = %extractor.key, error = %e, "Context extraction LLM call failed");
8363 }
8364 }
8365 }
8366 staged
8367 }
8368
8369 fn commit_staged_context_writes(&self, staged: &HashMap<String, Value>) {
8370 for (key, value) in staged {
8371 if let Err(error) = self.context_manager.update(key, value.clone()) {
8372 warn!(key = %key, error = %error, "staged context write failed");
8373 }
8374 }
8375 }
8376
8377 async fn run_context_extractors(&self, user_message: &str) {
8379 let staged = self.run_context_extractors_staged(user_message).await;
8380 self.commit_staged_context_writes(&staged);
8381 }
8382
8383 async fn check_memory_compression(&self) -> Result<()> {
8384 if self.memory.needs_compression() {
8385 let result = self.memory.compress(None).await?;
8386 if let CompressResult::Compressed {
8387 messages_summarized,
8388 new_summary_length,
8389 tokens_saved,
8390 } = result
8391 {
8392 let event = MemoryCompressEvent::new(
8393 messages_summarized,
8394 tokens_saved,
8395 new_summary_length as u32,
8396 );
8397 self.hooks.on_memory_compress(&event).await;
8398 debug!(
8399 messages = messages_summarized,
8400 tokens_saved = tokens_saved,
8401 "Memory compressed"
8402 );
8403 }
8404 }
8405
8406 self.handle_memory_overflow().await?;
8408 self.check_memory_budget().await;
8409
8410 Ok(())
8411 }
8412
8413 async fn check_memory_budget(&self) {
8414 let Some(ref budget) = self.memory_token_budget else {
8415 return;
8416 };
8417
8418 let context = match self.memory.get_context().await {
8419 Ok(ctx) => ctx,
8420 Err(_) => return,
8421 };
8422
8423 let used_tokens = context.estimated_tokens();
8425 if budget.is_over_warn_threshold(used_tokens) {
8426 let event = MemoryBudgetEvent::new("memory", used_tokens, budget.total);
8427 self.hooks.on_memory_budget_warning(&event).await;
8428 debug!(
8429 used = used_tokens,
8430 total = budget.total,
8431 percent = event.usage_percent,
8432 "Memory budget warning"
8433 );
8434 }
8435
8436 if let Some(ref summary) = context.summary {
8438 let summary_tokens = ai_agents_memory::estimate_tokens(summary);
8439 let summary_budget = budget.allocation.summary;
8440 if summary_budget > 0 {
8441 let warn_threshold =
8442 (summary_budget as f64 * budget.warn_at_percent as f64 / 100.0) as u32;
8443 if summary_tokens >= warn_threshold {
8444 let event = MemoryBudgetEvent::new("summary", summary_tokens, summary_budget);
8445 self.hooks.on_memory_budget_warning(&event).await;
8446 }
8447 }
8448 }
8449
8450 let recent_tokens: u32 = context
8452 .messages
8453 .iter()
8454 .map(ai_agents_memory::estimate_message_tokens)
8455 .sum();
8456 let recent_budget = budget.allocation.recent_messages;
8457 if recent_budget > 0 {
8458 let warn_threshold =
8459 (recent_budget as f64 * budget.warn_at_percent as f64 / 100.0) as u32;
8460 if recent_tokens >= warn_threshold {
8461 let event = MemoryBudgetEvent::new("recent_messages", recent_tokens, recent_budget);
8462 self.hooks.on_memory_budget_warning(&event).await;
8463 }
8464 }
8465
8466 let relationship_budget = budget.allocation.relationships;
8467 if relationship_budget > 0 {
8468 let relationship_tokens = self
8469 .relationship_memory_text()
8470 .map(|text| ai_agents_memory::estimate_tokens(&text))
8471 .unwrap_or(0);
8472 let warn_threshold =
8473 (relationship_budget as f64 * budget.warn_at_percent as f64 / 100.0) as u32;
8474 if relationship_tokens >= warn_threshold {
8475 let event = MemoryBudgetEvent::new(
8476 "relationships",
8477 relationship_tokens,
8478 relationship_budget,
8479 );
8480 self.hooks.on_memory_budget_warning(&event).await;
8481 }
8482 }
8483 }
8484
8485 async fn handle_memory_overflow(&self) -> Result<()> {
8486 let Some(ref budget) = self.memory_token_budget else {
8487 return Ok(());
8488 };
8489
8490 let context = self.memory.get_context().await?;
8491 let used_tokens = context.estimated_tokens();
8492
8493 if used_tokens <= budget.total {
8494 return Ok(());
8495 }
8496
8497 match budget.overflow_strategy {
8498 OverflowStrategy::TruncateOldest => {
8499 let tokens_to_free = used_tokens - budget.total;
8500 let messages_to_evict = self.calculate_eviction_count(tokens_to_free);
8501 if messages_to_evict > 0 {
8502 self.evict_messages(messages_to_evict, EvictionReason::TokenBudgetExceeded)
8503 .await?;
8504 }
8505 }
8506 OverflowStrategy::SummarizeMore => {
8507 let max_attempts = context.total_messages.max(1);
8508 for _ in 0..max_attempts {
8509 match self.memory.compress(None).await? {
8510 CompressResult::Compressed {
8511 messages_summarized,
8512 ..
8513 } if messages_summarized > 0 => {
8514 let context = self.memory.get_context().await?;
8515 if context.estimated_tokens() <= budget.total {
8516 return Ok(());
8517 }
8518 }
8519 _ => break,
8520 }
8521 }
8522 let context = self.memory.get_context().await?;
8523 let used_tokens = context.estimated_tokens();
8524 if used_tokens > budget.total {
8525 return Err(AgentError::MemoryBudgetExceeded {
8526 used: used_tokens,
8527 budget: budget.total,
8528 });
8529 }
8530 }
8531 OverflowStrategy::Error => {
8532 return Err(AgentError::MemoryBudgetExceeded {
8533 used: used_tokens,
8534 budget: budget.total,
8535 });
8536 }
8537 }
8538 Ok(())
8539 }
8540
8541 fn calculate_eviction_count(&self, tokens_to_free: u32) -> usize {
8542 ((tokens_to_free as f64 / 50.0).ceil() as usize).max(1)
8544 }
8545
8546 async fn evict_messages(&self, count: usize, reason: EvictionReason) -> Result<()> {
8547 let evicted = self.memory.evict_oldest(count).await?;
8548 if !evicted.is_empty() {
8549 let event = MemoryEvictEvent {
8550 reason,
8551 messages_evicted: evicted.len(),
8552 importance_scores: vec![],
8553 };
8554 self.hooks.on_memory_evict(&event).await;
8555 debug!(count = evicted.len(), "Messages evicted from memory");
8556 }
8557 Ok(())
8558 }
8559
8560 #[instrument(skip(self, input), fields(agent = %self.info.name))]
8561 async fn determine_reasoning_mode(&self, input: &str) -> Result<ReasoningMode> {
8562 match self.determine_reasoning_mode_strict(input).await {
8563 Ok(mode) => Ok(mode),
8564 Err(_) => Ok(ReasoningMode::None),
8565 }
8566 }
8567
8568 async fn determine_reasoning_mode_strict(&self, input: &str) -> Result<ReasoningMode> {
8569 let effective_config = self.get_effective_reasoning_config();
8570
8571 if !matches!(effective_config.mode, ReasoningMode::Auto) {
8572 return Ok(effective_config.mode.clone());
8573 }
8574
8575 let judge_llm = effective_config
8576 .judge_llm
8577 .as_ref()
8578 .and_then(|alias| self.llm_registry.get(alias).ok())
8579 .or_else(|| self.llm_registry.router().ok())
8580 .or_else(|| self.llm_registry.default().ok());
8581
8582 let Some(llm) = judge_llm else {
8583 return Ok(ReasoningMode::None);
8584 };
8585
8586 let prompt = format!(
8587 r#"Analyze this user request and determine the appropriate reasoning mode.
8588
8589User request: "{}"
8590
8591Choose ONE of these modes:
8592- none: Simple queries, greetings, direct answers (fastest)
8593- cot: Complex analysis, multi-step reasoning, math problems
8594- react: Tasks requiring multiple tool calls with observation
8595- plan_and_execute: Complex multi-step tasks requiring coordination
8596
8597Respond with ONLY the mode name (none, cot, react, or plan_and_execute)."#,
8598 input
8599 );
8600
8601 let messages = vec![ChatMessage::user(&prompt)];
8602 let response = self
8603 .observe_purpose(
8604 ObservationPurpose::ReflectionDecision,
8605 llm.complete(&messages, None),
8606 )
8607 .await
8608 .map_err(|e| AgentError::LLM(e.to_string()))?;
8609
8610 let mode_str = response.content.trim().to_lowercase();
8611 Ok(match mode_str.as_str() {
8612 "cot" => ReasoningMode::CoT,
8613 "react" => ReasoningMode::React,
8614 "plan_and_execute" => ReasoningMode::PlanAndExecute,
8615 _ => ReasoningMode::None,
8616 })
8617 }
8618
8619 async fn should_reflect(&self, input: &str, response: &str) -> Result<bool> {
8620 let effective_config = self.get_effective_reflection_config();
8621
8622 if !effective_config.requires_evaluation() {
8623 return Ok(false);
8624 }
8625
8626 if effective_config.is_enabled() {
8627 return Ok(true);
8628 }
8629
8630 let evaluator_llm = effective_config
8631 .evaluator_llm
8632 .as_ref()
8633 .and_then(|alias| self.llm_registry.get(alias).ok())
8634 .or_else(|| self.llm_registry.router().ok())
8635 .or_else(|| self.llm_registry.default().ok());
8636
8637 let Some(llm) = evaluator_llm else {
8638 return Ok(false);
8639 };
8640
8641 let response_preview: String = response.chars().take(500).collect();
8642 let prompt = format!(
8643 r#"Should this response be evaluated for quality? Consider if it's a complex or important response.
8644
8645User query: "{}"
8646Response: "{}"
8647
8648Answer YES or NO only."#,
8649 input, response_preview
8650 );
8651
8652 let messages = vec![ChatMessage::user(&prompt)];
8653 let result = self
8654 .observe_purpose(
8655 ObservationPurpose::ReflectionDecision,
8656 llm.complete(&messages, None),
8657 )
8658 .await;
8659
8660 match result {
8661 Ok(resp) => Ok(resp.content.trim().to_uppercase().contains("YES")),
8662 Err(_) => Ok(false),
8663 }
8664 }
8665
8666 fn build_cot_system_prompt(&self, base_prompt: &str) -> String {
8667 format!(
8668 "{}\n\n<instruction>\nThink through this step by step before answering:\n1. Understand what is being asked\n2. Break down the problem\n3. Work through each part\n4. Provide your final answer\n\nShow your thinking process, then give your final answer.\n</instruction>",
8669 base_prompt
8670 )
8671 }
8672
8673 fn build_react_system_prompt(&self, base_prompt: &str) -> String {
8674 format!(
8675 "{}\n\n<instruction>\nUse the Reason-Act-Observe pattern:\n1. Thought: Think about what to do\n2. Action: Use a tool if needed\n3. Observation: Analyze the result\n4. Repeat until you have the answer\n\nFormat your response showing Thought/Action/Observation steps.\n</instruction>",
8676 base_prompt
8677 )
8678 }
8679
8680 async fn generate_plan(&self, input: &str) -> Result<Plan> {
8681 let effective = self.get_effective_reasoning_config();
8682 let planning_config = effective.get_planning();
8683
8684 let planner_llm = planning_config
8685 .and_then(|c| c.planner_llm.as_ref())
8686 .and_then(|alias| self.llm_registry.get(alias).ok())
8687 .or_else(|| self.llm_registry.router().ok())
8688 .or_else(|| self.llm_registry.default().ok())
8689 .ok_or_else(|| AgentError::Config("No LLM available for planning".into()))?;
8690
8691 let mut available_tool_ids: Vec<String> = self
8692 .get_available_tool_ids()
8693 .await
8694 .unwrap_or_else(|_| self.tools.list_ids());
8695 let mut available_skills: Vec<String> = self.skills.iter().map(|s| s.id.clone()).collect();
8696
8697 if let Some(config) = planning_config {
8699 if !config.available.tools.is_all() {
8700 available_tool_ids.retain(|t| config.available.tools.allows(t));
8701 }
8702 if !config.available.skills.is_all() {
8703 available_skills.retain(|s| config.available.skills.allows(s));
8704 }
8705 }
8706
8707 let tool_descriptions: Vec<String> = available_tool_ids
8710 .iter()
8711 .filter_map(|id| {
8712 self.tools.get(id).map(|tool| {
8713 let schema = tool.input_schema();
8714 let args_desc = schema
8715 .get("properties")
8716 .and_then(|p| serde_json::to_string(p).ok())
8717 .unwrap_or_else(|| "{}".to_string());
8718 format!(
8719 "- {} ({}): {}\n Arguments: {}",
8720 id,
8721 tool.name(),
8722 tool.description(),
8723 args_desc
8724 )
8725 })
8726 })
8727 .collect();
8728
8729 let tools_section = if tool_descriptions.is_empty() {
8730 "Available tools: none".to_string()
8731 } else {
8732 format!("Available tools:\n{}", tool_descriptions.join("\n"))
8733 };
8734
8735 let skills_section = if available_skills.is_empty() {
8736 "Available skills: none".to_string()
8737 } else {
8738 format!("Available skills: {}", available_skills.join(", "))
8739 };
8740
8741 let prompt = format!(
8742 r#"Create a step-by-step plan to accomplish this goal.
8743
8744Goal: "{}"
8745
8746{}
8747
8748{}
8749
8750Create a plan with clear steps. For each step, specify:
8751- description: What this step accomplishes
8752- action_type: "tool", "skill", "think", or "respond"
8753- action_target: The tool/skill id (if applicable)
8754- args: The arguments object matching the tool's schema (if action_type is "tool")
8755- dependencies: List of step IDs this depends on (empty if none)
8756
8757Respond in JSON format:
8758{{
8759 "steps": [
8760 {{"id": "step1", "description": "...", "action_type": "tool", "action_target": "tool_id", "args": {{"required_field": "value"}}, "dependencies": []}},
8761 {{"id": "step2", "description": "...", "action_type": "think", "action_target": "...", "dependencies": ["step1"]}}
8762 ]
8763}}"#,
8764 input, tools_section, skills_section,
8765 );
8766
8767 let messages = vec![ChatMessage::user(&prompt)];
8768 let response = self
8769 .observe_purpose(
8770 ObservationPurpose::PlanGeneration,
8771 planner_llm.complete(&messages, None),
8772 )
8773 .await
8774 .map_err(|e| AgentError::LLM(format!("Planning failed: {}", e)))?;
8775
8776 let mut plan = Plan::new(input);
8777
8778 if let Some(json_start) = response.content.find('{')
8779 && let Some(json_end) = response.content.rfind('}')
8780 {
8781 let json_str = &response.content[json_start..=json_end];
8782 if let Ok(parsed) = serde_json::from_str::<serde_json::Value>(json_str)
8783 && let Some(steps) = parsed.get("steps").and_then(|s| s.as_array())
8784 {
8785 for step_value in steps {
8786 let id = step_value
8787 .get("id")
8788 .and_then(|v| v.as_str())
8789 .unwrap_or("step");
8790 let desc = step_value
8791 .get("description")
8792 .and_then(|v| v.as_str())
8793 .unwrap_or("");
8794 let action_type = step_value
8795 .get("action_type")
8796 .and_then(|v| v.as_str())
8797 .unwrap_or("think");
8798 let action_target = step_value
8799 .get("action_target")
8800 .and_then(|v| v.as_str())
8801 .unwrap_or("");
8802 let args = step_value
8803 .get("args")
8804 .cloned()
8805 .unwrap_or(serde_json::json!({}));
8806 let deps: Vec<String> = step_value
8807 .get("dependencies")
8808 .and_then(|v| v.as_array())
8809 .map(|arr| {
8810 arr.iter()
8811 .filter_map(|v| v.as_str().map(String::from))
8812 .collect()
8813 })
8814 .unwrap_or_default();
8815
8816 let action = match action_type {
8817 "tool" => PlanAction::tool(action_target, args),
8818 "skill" => PlanAction::skill(action_target),
8819 "respond" => PlanAction::respond(action_target),
8820 _ => PlanAction::think(desc),
8821 };
8822
8823 let step = PlanStep::new(desc, action)
8824 .with_id(id)
8825 .with_dependencies(deps);
8826 plan.add_step(step);
8827 }
8828 }
8829 }
8830
8831 if plan.steps.is_empty() {
8832 plan.add_step(PlanStep::new(
8833 "Process the request",
8834 PlanAction::think(input),
8835 ));
8836 plan.add_step(PlanStep::new(
8837 "Provide response",
8838 PlanAction::respond("Answer based on analysis"),
8839 ));
8840 }
8841
8842 Ok(plan)
8843 }
8844
8845 async fn execute_plan(&self, plan: &mut Plan) -> Result<String> {
8846 let llm = self.get_state_llm()?;
8847 let mut results: HashMap<String, serde_json::Value> = HashMap::new();
8848 let effective = self.get_effective_reasoning_config();
8849 let max_steps = effective.get_planning().map(|c| c.max_steps).unwrap_or(10);
8850
8851 plan.status = PlanStatus::InProgress;
8852
8853 for step_idx in 0..plan.steps.len().min(max_steps as usize) {
8854 let step = &plan.steps[step_idx];
8855
8856 let deps_satisfied = step.dependencies.iter().all(|dep| {
8857 plan.steps
8858 .iter()
8859 .find(|s| &s.id == dep)
8860 .map(|s| s.status.is_completed())
8861 .unwrap_or(false)
8862 });
8863
8864 if !deps_satisfied {
8865 continue;
8866 }
8867
8868 plan.steps[step_idx].mark_running();
8869
8870 let result = match &plan.steps[step_idx].action {
8871 PlanAction::Tool { tool, args } => {
8872 let has_dep_results = plan.steps[step_idx]
8878 .dependencies
8879 .iter()
8880 .any(|dep| results.contains_key(dep));
8881
8882 let final_args = if has_dep_results {
8883 let dep_context: String = plan.steps[step_idx]
8884 .dependencies
8885 .iter()
8886 .filter_map(|dep| results.get(dep).map(|r| format!("{}: {}", dep, r)))
8887 .collect::<Vec<_>>()
8888 .join("\n");
8889
8890 let tool_schema = self
8891 .tools
8892 .get(tool)
8893 .map(|t| {
8894 let schema = t.input_schema();
8895 let props = schema
8896 .get("properties")
8897 .and_then(|p| serde_json::to_string(p).ok())
8898 .unwrap_or_else(|| "{}".to_string());
8899 format!(
8900 "{}: {}\nArguments schema: {}",
8901 t.id(),
8902 t.description(),
8903 props
8904 )
8905 })
8906 .unwrap_or_default();
8907
8908 let step_desc = &plan.steps[step_idx].description;
8909 let arg_prompt = format!(
8910 "Generate the JSON arguments for a tool call.\n\n\
8911 Tool: {}\n\n\
8912 Task: {}\n\n\
8913 Previous step results:\n{}\n\n\
8914 Planner's draft arguments: {}\n\n\
8915 Produce ONLY a valid JSON object with the correct argument values.\n\
8916 Use actual values from the previous step results, not template references.",
8917 tool_schema,
8918 step_desc,
8919 dep_context,
8920 serde_json::to_string(args).unwrap_or_default()
8921 );
8922 let messages = vec![ChatMessage::user(&arg_prompt)];
8923 match self
8924 .observe_purpose(
8925 ObservationPurpose::PlanStep,
8926 llm.complete(&messages, None),
8927 )
8928 .await
8929 {
8930 Ok(resp) => {
8931 let content = resp.content.trim();
8932 let json_start = content.find('{');
8934 let json_end = content.rfind('}');
8935 if let (Some(start), Some(end)) = (json_start, json_end) {
8936 serde_json::from_str(&content[start..=end])
8937 .unwrap_or_else(|_| args.clone())
8938 } else {
8939 args.clone()
8940 }
8941 }
8942 Err(_) => args.clone(),
8943 }
8944 } else {
8945 args.clone()
8946 };
8947
8948 let request = ToolExecutionRequest::new(
8949 uuid::Uuid::new_v4().to_string(),
8950 tool.clone(),
8951 final_args,
8952 ToolCallSource::Plan {
8953 step_index: step_idx,
8954 },
8955 );
8956 match self.execute_tool_record(request).await {
8957 Ok(record) if record.success => {
8958 serde_json::json!({ "output": record.model_output_string() })
8959 }
8960 Ok(record) => {
8961 plan.steps[step_idx].mark_failed(record.model_output_string());
8962 continue;
8963 }
8964 Err(e) => {
8965 plan.steps[step_idx].mark_failed(e.to_string());
8966 continue;
8967 }
8968 }
8969 }
8970 PlanAction::Skill { skill } => {
8971 if let Some(skill_def) = self.skills.iter().find(|s| &s.id == skill) {
8972 if let Some(ref executor) = self.skill_executor {
8973 match executor
8974 .execute_with_invoker(skill_def, "", serde_json::json!({}), self)
8975 .await
8976 {
8977 Ok(output) => serde_json::json!({ "output": output }),
8978 Err(e) => {
8979 plan.steps[step_idx].mark_failed(e.to_string());
8980 continue;
8981 }
8982 }
8983 } else {
8984 serde_json::json!({ "output": "Skill executor not available" })
8985 }
8986 } else {
8987 plan.steps[step_idx].mark_failed("Skill not found");
8988 continue;
8989 }
8990 }
8991 PlanAction::Think { prompt } => {
8992 let context: String = results
8993 .iter()
8994 .map(|(k, v)| format!("{}: {}", k, v))
8995 .collect::<Vec<_>>()
8996 .join("\n");
8997
8998 let think_prompt = format!("Context:\n{}\n\nTask: {}", context, prompt);
8999 let messages = vec![ChatMessage::user(&think_prompt)];
9000
9001 match self
9002 .observe_purpose(
9003 ObservationPurpose::PlanStep,
9004 llm.complete(&messages, None),
9005 )
9006 .await
9007 {
9008 Ok(resp) => serde_json::json!({ "output": resp.content }),
9009 Err(e) => {
9010 plan.steps[step_idx].mark_failed(e.to_string());
9011 continue;
9012 }
9013 }
9014 }
9015 PlanAction::Respond { template } => {
9016 let context: String = results
9017 .iter()
9018 .map(|(k, v)| format!("{}: {}", k, v))
9019 .collect::<Vec<_>>()
9020 .join("\n");
9021
9022 let respond_prompt = format!(
9023 "Based on this context:\n{}\n\nGenerate a response following this template/instruction: {}",
9024 context, template
9025 );
9026 let messages = vec![ChatMessage::user(&respond_prompt)];
9027
9028 match self
9029 .observe_purpose(
9030 ObservationPurpose::PlanStep,
9031 llm.complete(&messages, None),
9032 )
9033 .await
9034 {
9035 Ok(resp) => serde_json::json!({ "output": resp.content }),
9036 Err(e) => {
9037 plan.steps[step_idx].mark_failed(e.to_string());
9038 continue;
9039 }
9040 }
9041 }
9042 };
9043
9044 results.insert(plan.steps[step_idx].id.clone(), result.clone());
9045 plan.steps[step_idx].mark_completed(Some(result));
9046 }
9047
9048 let has_failures = plan.steps.iter().any(|s| s.status.is_failed());
9050 if has_failures {
9051 let failed_ids: Vec<String> = plan
9052 .steps
9053 .iter()
9054 .filter(|s| s.status.is_failed())
9055 .map(|s| s.id.clone())
9056 .collect();
9057 plan.status = PlanStatus::Failed {
9058 error: format!("Steps failed: {}", failed_ids.join(", ")),
9059 };
9060 } else {
9061 plan.status = PlanStatus::Completed;
9062 }
9063
9064 let all_outputs: Vec<String> = plan
9066 .steps
9067 .iter()
9068 .filter(|s| s.status.is_completed())
9069 .filter_map(|s| {
9070 s.result
9071 .as_ref()
9072 .and_then(|r| r.get("output"))
9073 .and_then(|o| o.as_str())
9074 .map(|o| format!("{}: {}", s.description, o))
9075 })
9076 .collect();
9077
9078 if all_outputs.is_empty() {
9079 return Ok("Plan execution completed but produced no results.".to_string());
9080 }
9081
9082 if all_outputs.len() == 1 {
9083 return Ok(all_outputs.into_iter().next().unwrap());
9084 }
9085
9086 let context = all_outputs.join("\n\n");
9088 let prompt = format!(
9089 "You completed a multi-step plan for: \"{}\"\n\nStep results:\n{}\n\nProvide a coherent final response that synthesizes these results.",
9090 plan.goal, context
9091 );
9092 let messages = vec![ChatMessage::user(&prompt)];
9093 match self
9094 .observe_purpose(ObservationPurpose::PlanStep, llm.complete(&messages, None))
9095 .await
9096 {
9097 Ok(resp) => Ok(resp.content.trim().to_string()),
9098 Err(_) => Ok(context),
9099 }
9100 }
9101
9102 async fn evaluate_response(&self, input: &str, response: &str) -> Result<EvaluationResult> {
9103 let effective_config = self.get_effective_reflection_config();
9104 self.evaluate_response_with_config(input, response, &effective_config)
9105 .await
9106 }
9107
9108 fn extract_thinking(&self, content: &str) -> (Option<String>, String) {
9109 if let Some(start) = content.find("<thinking>")
9110 && let Some(end) = content.find("</thinking>")
9111 {
9112 let thinking = content[start + 10..end].trim().to_string();
9113 let answer = content[end + 11..].trim().to_string();
9114 return (Some(thinking), answer);
9115 }
9116 (None, content.to_string())
9117 }
9118
9119 fn format_response_with_thinking(&self, thinking: Option<&str>, answer: &str) -> String {
9120 match self.get_effective_reasoning_config().output {
9121 ReasoningOutput::Hidden => answer.to_string(),
9122 ReasoningOutput::Visible => {
9123 if let Some(t) = thinking {
9124 format!("Thinking:\n{}\n\nAnswer:\n{}", t, answer)
9125 } else {
9126 answer.to_string()
9127 }
9128 }
9129 ReasoningOutput::Tagged => {
9130 if let Some(t) = thinking {
9131 format!("<thinking>{}</thinking>\n{}", t, answer)
9132 } else {
9133 answer.to_string()
9134 }
9135 }
9136 }
9137 }
9138
9139 fn disambiguation_question_response(
9142 question: &ClarificationQuestion,
9143 detection: &AmbiguityDetectionResult,
9144 awaiting_confirmation: bool,
9145 ) -> AgentResponse {
9146 let status = if awaiting_confirmation {
9147 "awaiting_confirmation"
9148 } else {
9149 "awaiting_clarification"
9150 };
9151 AgentResponse::new(&question.question).with_metadata(
9152 "disambiguation",
9153 serde_json::json!({
9154 "status": status,
9155 "options": question.options,
9156 "clarifying": question.clarifying,
9157 "detection": {
9158 "type": detection.ambiguity_type,
9159 "confidence": detection.confidence,
9160 "what_is_unclear": detection.what_is_unclear,
9161 }
9162 }),
9163 )
9164 }
9165
9166 async fn resolve_disambiguation(&self, input: &str) -> Result<DisambiguationDispatch> {
9178 let Some(ref disambiguator) = self.disambiguation_manager else {
9179 return Ok(DisambiguationDispatch::Proceed(input.to_string()));
9180 };
9181 let disambiguation_context = self.build_disambiguation_context().await?;
9182
9183 let state_override = self
9185 .state_machine
9186 .as_ref()
9187 .and_then(|sm| sm.current_definition())
9188 .and_then(|def| def.disambiguation.clone());
9189
9190 let state_generation = self
9191 .state_machine
9192 .as_ref()
9193 .map(|state_machine| state_machine.generation());
9194 let disambiguation_epoch = self.disambiguation_epoch.load(Ordering::SeqCst);
9195 let mut disambiguation_result = self
9196 .observe_purpose(
9197 ObservationPurpose::DisambiguationDetection,
9198 disambiguator.process_input_with_override(
9199 input,
9200 &disambiguation_context,
9201 state_override.as_ref(),
9202 None,
9203 ),
9204 )
9205 .await?;
9206 let current_state_generation = self
9207 .state_machine
9208 .as_ref()
9209 .map(|state_machine| state_machine.generation());
9210 if current_state_generation != state_generation
9211 || self.disambiguation_epoch.load(Ordering::SeqCst) != disambiguation_epoch
9212 {
9213 disambiguator.clear_pending().await;
9214 *self.pending_skill_id.write() = None;
9215 disambiguation_result = DisambiguationResult::Abandoned { new_input: None };
9216 info!(
9217 confirmation_event = "invalidated",
9218 invalidation_reason = "state_generation_changed",
9219 "Disambiguation result invalidated before redispatch"
9220 );
9221 }
9222 match disambiguation_result {
9223 DisambiguationResult::Clear => {
9224 debug!("Input is clear, proceeding normally");
9225 Ok(DisambiguationDispatch::Proceed(input.to_string()))
9226 }
9227 DisambiguationResult::NeedsClarification {
9228 question,
9229 detection,
9230 } => {
9231 let admission = match self
9232 .admit_disambiguation_redispatch(disambiguation_epoch, state_generation)
9233 .await
9234 {
9235 Ok(admission) => admission,
9236 Err(error) => {
9237 *self.pending_skill_id.write() = None;
9238 return Err(error);
9239 }
9240 };
9241 let awaiting_confirmation = disambiguator.has_pending_confirmation().await;
9242 info!(
9243 ambiguity_type = ?detection.ambiguity_type,
9244 confidence = detection.confidence,
9245 "Input requires clarification"
9246 );
9247
9248 self.commit_root_user_message(input).await?;
9251 self.memory
9252 .add_message(ChatMessage::assistant(&question.question))
9253 .await?;
9254
9255 let response = Self::disambiguation_question_response(
9256 &question,
9257 &detection,
9258 awaiting_confirmation,
9259 );
9260 drop(admission);
9261 self.finish_turn_if_root(&response).await?;
9262 Ok(DisambiguationDispatch::Terminal(response))
9263 }
9264 DisambiguationResult::Clarified {
9265 enriched_input,
9266 resolved,
9267 ..
9268 } => {
9269 let admission = match self
9270 .admit_disambiguation_redispatch(disambiguation_epoch, state_generation)
9271 .await
9272 {
9273 Ok(admission) => admission,
9274 Err(error) => {
9275 *self.pending_skill_id.write() = None;
9276 return Err(error);
9277 }
9278 };
9279 info!(
9280 resolved_count = resolved.len(),
9281 enriched = %enriched_input,
9282 "Input clarified, injecting resolved intent into context"
9283 );
9284
9285 for (key, value) in &resolved {
9288 let context_key = format!("disambiguation.{}", key);
9289 let _ = self.context_manager.set(&context_key, value.clone());
9290 }
9291
9292 if let Some(intent) = resolved.get("intent") {
9293 let _ = self.context_manager.set("resolved_intent", intent.clone());
9294 }
9295
9296 let _ = self
9297 .context_manager
9298 .set("disambiguation.resolved", serde_json::Value::Bool(true));
9299
9300 let skill_id = self.pending_skill_id.read().clone();
9304 drop(admission);
9305 if let Some(skill_id) = skill_id {
9306 info!(skill_id = %skill_id, "Re-checking skill disambiguation on clarified input");
9307 return Ok(DisambiguationDispatch::RecheckSkill {
9308 skill_id,
9309 enriched_input,
9310 disambiguation_epoch,
9311 state_generation,
9312 });
9313 }
9314 Ok(DisambiguationDispatch::Proceed(enriched_input))
9315 }
9316 DisambiguationResult::ProceedWithBestGuess { enriched_input } => {
9317 info!("Proceeding with best guess interpretation");
9318
9319 let skill_id = self.pending_skill_id.read().clone();
9321 if let Some(skill_id) = skill_id {
9322 info!(skill_id = %skill_id, "Re-checking skill disambiguation on best-guess input");
9323 return Ok(DisambiguationDispatch::RecheckSkill {
9324 skill_id,
9325 enriched_input,
9326 disambiguation_epoch,
9327 state_generation,
9328 });
9329 }
9330 Ok(DisambiguationDispatch::Proceed(enriched_input))
9331 }
9332 DisambiguationResult::GiveUp { reason } => {
9333 *self.pending_skill_id.write() = None;
9334 warn!(reason = %reason, "Disambiguation gave up");
9335 let apology = self
9336 .generate_localized_apology(
9337 "Generate a brief, polite apology saying you couldn't understand the request. Be concise.",
9338 &reason,
9339 )
9340 .await
9341 .unwrap_or_else(|_| {
9342 format!("I'm sorry, I couldn't understand your request: {}", reason)
9343 });
9344 let response = AgentResponse::new(&apology);
9345 self.finish_turn_if_root(&response).await?;
9346 Ok(DisambiguationDispatch::Terminal(response))
9347 }
9348 DisambiguationResult::Escalate { reason } => {
9349 *self.pending_skill_id.write() = None;
9350 info!(reason = %reason, "Escalating to human");
9351 if let Some(ref hitl) = self.hitl_engine {
9352 let trigger =
9353 ApprovalTrigger::condition("disambiguation_escalation", reason.clone());
9354 let mut context_map = HashMap::new();
9355 context_map.insert("original_input".to_string(), serde_json::json!(input));
9356 context_map.insert("reason".to_string(), serde_json::json!(&reason));
9357 let check_result = HITLCheckResult::required(
9358 trigger,
9359 context_map,
9360 format!("User request needs human assistance: {}", reason),
9361 Some(hitl.config().default_timeout_seconds),
9362 );
9363 let result = self.request_hitl_approval(check_result).await?;
9364 if matches!(
9365 result,
9366 ApprovalResult::Approved | ApprovalResult::Modified { .. }
9367 ) {
9368 return Ok(DisambiguationDispatch::Proceed(input.to_string()));
9370 }
9371 }
9372 let apology = self
9373 .generate_localized_apology(
9374 "Explain briefly that you're transferring the user to a human agent for help.",
9375 &reason,
9376 )
9377 .await
9378 .unwrap_or_else(|_| {
9379 format!("I need human assistance to help with your request: {}", reason)
9380 });
9381 let response = AgentResponse::new(&apology);
9382 self.finish_turn_if_root(&response).await?;
9383 Ok(DisambiguationDispatch::Terminal(response))
9384 }
9385 DisambiguationResult::Abandoned { new_input } => {
9386 *self.pending_skill_id.write() = None;
9387
9388 info!(
9389 has_new_input = new_input.is_some(),
9390 "Clarification abandoned by user"
9391 );
9392
9393 self.commit_root_user_message(input).await?;
9394
9395 match new_input {
9396 Some(fresh_input) => {
9397 Ok(DisambiguationDispatch::Proceed(fresh_input))
9400 }
9401 None => {
9402 let ack = self
9404 .generate_localized_apology(
9405 "The user changed their mind about their previous request. \
9406 Generate a brief, friendly acknowledgment (e.g. 'OK, no problem. What else can I help with?'). \
9407 Do NOT apologize excessively. Be concise.",
9408 "User abandoned clarification",
9409 )
9410 .await
9411 .unwrap_or_else(|_| {
9412 "OK, no problem. What else can I help with?".to_string()
9413 });
9414
9415 self.memory
9416 .add_message(ChatMessage::assistant(&ack))
9417 .await?;
9418
9419 let response = AgentResponse::new(&ack);
9420 self.finish_turn_if_root(&response).await?;
9421 Ok(DisambiguationDispatch::Terminal(response))
9422 }
9423 }
9424 }
9425 }
9426 }
9427
9428 async fn prepare_turn_context(&self) -> Result<()> {
9432 if !self.context_initialized.load(Ordering::SeqCst) {
9433 self.context_manager.initialize().await?;
9434 self.context_initialized.store(true, Ordering::SeqCst);
9435 debug!("Context manager initialized (defaults, env, builtins)");
9436 }
9437
9438 self.check_turn_timeout().await?;
9439 self.context_manager.refresh_per_turn().await?;
9440 self.context_manager.validate()
9441 }
9442
9443 async fn run_loop(&self, input: &str) -> Result<AgentResponse> {
9446 self.init_storage().await?;
9450 self.begin_root_turn();
9451 let _root_cleanup = RootTurnCleanup::new(self);
9452 info!(input_len = input.len(), "Starting chat");
9453
9454 self.hooks.on_message_received(input).await;
9455
9456 self.prepare_turn_context().await?;
9457
9458 self.clear_disambiguation_context();
9461
9462 let input_to_run = match self.resolve_disambiguation(input).await? {
9465 DisambiguationDispatch::Terminal(response) => return Ok(response),
9466 DisambiguationDispatch::RecheckSkill {
9467 skill_id,
9468 enriched_input,
9469 disambiguation_epoch,
9470 state_generation,
9471 } => {
9472 return self
9473 .recheck_skill_disambiguation(
9474 &skill_id,
9475 &enriched_input,
9476 disambiguation_epoch,
9477 state_generation,
9478 )
9479 .await;
9480 }
9481 DisambiguationDispatch::Proceed(input) => input,
9482 };
9483
9484 self.run_loop_internal(&input_to_run).await
9485 }
9486
9487 async fn generate_localized_apology(&self, instruction: &str, reason: &str) -> Result<String> {
9489 let llm = self.llm_registry.router().map_err(|e| {
9490 AgentError::LLM(format!(
9491 "Router LLM not available for localized response: {}",
9492 e
9493 ))
9494 })?;
9495
9496 let recent: Vec<String> = self
9497 .memory
9498 .get_messages(Some(3))
9499 .await?
9500 .iter()
9501 .map(|m| m.content.clone())
9502 .collect();
9503
9504 let context_hint = if recent.is_empty() {
9505 String::new()
9506 } else {
9507 format!(
9508 "\nRecent conversation (detect the user's language from this):\n{}\n",
9509 recent.join("\n")
9510 )
9511 };
9512
9513 let prompt = format!(
9514 "{}\nReason: {}\n{}Respond in the same language as the user. Output ONLY the message, nothing else.",
9515 instruction, reason, context_hint
9516 );
9517
9518 let messages = vec![ChatMessage::user(&prompt)];
9519 let response = self
9520 .observe_purpose(
9521 ObservationPurpose::DisambiguationClarification,
9522 llm.complete(&messages, None),
9523 )
9524 .await
9525 .map_err(|e| AgentError::LLM(format!("Localized response generation failed: {}", e)))?;
9526
9527 Ok(response.content.trim().to_string())
9528 }
9529
9530 fn render_action_args(&self, args: &Value) -> Value {
9534 let context = self.build_context_with_overlays();
9535 match args {
9536 Value::Object(map) => {
9537 let mut rendered = serde_json::Map::new();
9538 for (k, v) in map {
9539 match v {
9540 Value::String(s) if s.contains("{{") => {
9541 match self.template_renderer.render(s, &context) {
9542 Ok(rendered_str) => {
9543 rendered.insert(k.clone(), Value::String(rendered_str));
9544 }
9545 Err(_) => {
9546 rendered.insert(k.clone(), v.clone());
9547 }
9548 }
9549 }
9550 _ => {
9551 rendered.insert(k.clone(), v.clone());
9552 }
9553 }
9554 }
9555 Value::Object(rendered)
9556 }
9557 _ => args.clone(),
9558 }
9559 }
9560
9561 fn clear_disambiguation_context(&self) {
9563 let _ = self
9564 .context_manager
9565 .set("resolved_intent", serde_json::Value::Null);
9566
9567 let all = self.context_manager.get_all();
9568 for key in all.keys() {
9569 if key.starts_with("disambiguation.") {
9570 let _ = self.context_manager.set(key, serde_json::Value::Null);
9571 }
9572 }
9573 }
9574
9575 async fn recheck_skill_disambiguation(
9581 &self,
9582 skill_id: &str,
9583 enriched_input: &str,
9584 expected_disambiguation_epoch: u64,
9585 expected_state_generation: Option<u64>,
9586 ) -> Result<AgentResponse> {
9587 let skill = self
9588 .skill_router
9589 .as_ref()
9590 .and_then(|r| r.get_skill(skill_id).cloned());
9591
9592 if let Some(ref skill) = skill
9594 && let Some(ref skill_disambig) = skill.disambiguation
9595 && skill_disambig.enabled.unwrap_or(false)
9596 && let Some(ref disambiguator) = self.disambiguation_manager
9597 {
9598 let context = self.build_disambiguation_context().await?;
9599 let state_override = self
9600 .state_machine
9601 .as_ref()
9602 .and_then(|sm| sm.current_definition())
9603 .and_then(|def| def.disambiguation.clone());
9604
9605 let disambiguation_result = self
9606 .observe_purpose(
9607 ObservationPurpose::DisambiguationDetection,
9608 disambiguator.process_input_with_override(
9609 enriched_input,
9610 &context,
9611 state_override.as_ref(),
9612 Some(skill_disambig),
9613 ),
9614 )
9615 .await?;
9616 let current_state_generation = self
9617 .state_machine
9618 .as_ref()
9619 .map(|state_machine| state_machine.generation());
9620 if current_state_generation != expected_state_generation
9621 || self.disambiguation_epoch.load(Ordering::SeqCst) != expected_disambiguation_epoch
9622 {
9623 disambiguator.clear_pending().await;
9624 *self.pending_skill_id.write() = None;
9625 return Err(AgentError::Other(
9626 "State or reset ownership changed during skill disambiguation recheck"
9627 .to_string(),
9628 ));
9629 }
9630 match disambiguation_result {
9631 DisambiguationResult::Clear => {
9632 debug!(skill_id = %skill_id, "Skill re-check: all fields present");
9633 }
9634 DisambiguationResult::NeedsClarification {
9635 question,
9636 detection,
9637 } => {
9638 let admission = self
9639 .admit_disambiguation_redispatch(
9640 expected_disambiguation_epoch,
9641 expected_state_generation,
9642 )
9643 .await?;
9644 let awaiting_confirmation = disambiguator.has_pending_confirmation().await;
9645 info!(
9646 skill_id = %skill_id,
9647 ambiguity_type = ?detection.ambiguity_type,
9648 what_is_unclear = ?detection.what_is_unclear,
9649 "Skill re-check: still missing fields, asking again"
9650 );
9651 self.memory
9655 .add_message(ChatMessage::user(enriched_input))
9656 .await?;
9657 self.memory
9658 .add_message(ChatMessage::assistant(&question.question))
9659 .await?;
9660
9661 let response = AgentResponse::new(&question.question).with_metadata(
9662 "disambiguation",
9663 serde_json::json!({
9664 "status": if awaiting_confirmation { "awaiting_confirmation" } else { "awaiting_clarification" },
9665 "skill_id": skill_id,
9666 "options": question.options,
9667 "clarifying": question.clarifying,
9668 "detection": {
9669 "type": detection.ambiguity_type,
9670 "confidence": detection.confidence,
9671 "what_is_unclear": detection.what_is_unclear,
9672 }
9673 }),
9674 );
9675 drop(admission);
9676 self.finish_turn_if_root(&response).await?;
9677 return Ok(response);
9678 }
9679 DisambiguationResult::Clarified {
9680 enriched_input: re_enriched,
9681 ..
9682 } => {
9683 debug!(skill_id = %skill_id, "Skill re-check: clarified immediately, executing");
9684 let admission = self
9685 .admit_disambiguation_redispatch(
9686 expected_disambiguation_epoch,
9687 expected_state_generation,
9688 )
9689 .await?;
9690 *self.pending_skill_id.write() = None;
9691 drop(admission);
9692 let skill_response = self.execute_skill_by_id(skill_id, &re_enriched).await?;
9693 self.memory
9694 .add_message(ChatMessage::user(&re_enriched))
9695 .await?;
9696 return self
9697 .handle_skill_response(
9698 &re_enriched,
9699 skill_id,
9700 skill_response,
9701 &HashMap::new(),
9702 )
9703 .await;
9704 }
9705 DisambiguationResult::ProceedWithBestGuess {
9706 enriched_input: re_enriched,
9707 } => {
9708 debug!(skill_id = %skill_id, "Skill re-check: proceeding with best guess");
9709 let admission = self
9710 .admit_disambiguation_redispatch(
9711 expected_disambiguation_epoch,
9712 expected_state_generation,
9713 )
9714 .await?;
9715 *self.pending_skill_id.write() = None;
9716 drop(admission);
9717 let skill_response = self.execute_skill_by_id(skill_id, &re_enriched).await?;
9718 self.memory
9719 .add_message(ChatMessage::user(&re_enriched))
9720 .await?;
9721 return self
9722 .handle_skill_response(
9723 &re_enriched,
9724 skill_id,
9725 skill_response,
9726 &HashMap::new(),
9727 )
9728 .await;
9729 }
9730 DisambiguationResult::GiveUp { reason } => {
9731 *self.pending_skill_id.write() = None;
9732 let apology = self
9733 .generate_localized_apology(
9734 "Generate a brief, polite apology saying you couldn't understand the request. Be concise.",
9735 &reason,
9736 )
9737 .await
9738 .unwrap_or_else(|_| {
9739 format!("I'm sorry, I couldn't understand your request: {}", reason)
9740 });
9741 let response = AgentResponse::new(&apology);
9742 self.finish_turn_if_root(&response).await?;
9743 return Ok(response);
9744 }
9745 DisambiguationResult::Escalate { reason } => {
9746 *self.pending_skill_id.write() = None;
9747 let apology = self
9748 .generate_localized_apology(
9749 "Explain briefly that you're transferring the user to a human agent for help.",
9750 &reason,
9751 )
9752 .await
9753 .unwrap_or_else(|_| {
9754 format!("I need human assistance to help with your request: {}", reason)
9755 });
9756 let response = AgentResponse::new(&apology);
9757 self.finish_turn_if_root(&response).await?;
9758 return Ok(response);
9759 }
9760 DisambiguationResult::Abandoned { new_input } => {
9761 *self.pending_skill_id.write() = None;
9764 debug!(skill_id = %skill_id, "Skill re-check: abandoned by user");
9765 if let Some(fresh) = new_input {
9766 return self.run_loop_internal(&fresh).await;
9767 }
9768 let ack = self
9769 .generate_localized_apology(
9770 "The user changed their mind about their previous request. \
9771 Generate a brief, friendly acknowledgment (e.g. 'OK, no problem. What else can I help with?'). \
9772 Do NOT apologize excessively. Be concise.",
9773 "User abandoned clarification",
9774 )
9775 .await
9776 .unwrap_or_else(|_| {
9777 "OK, no problem. What else can I help with?".to_string()
9778 });
9779 self.memory
9780 .add_message(ChatMessage::assistant(&ack))
9781 .await?;
9782 let response = AgentResponse::new(&ack);
9783 self.finish_turn_if_root(&response).await?;
9784 return Ok(response);
9785 }
9786 }
9787 }
9788
9789 let admission = self
9791 .admit_disambiguation_redispatch(
9792 expected_disambiguation_epoch,
9793 expected_state_generation,
9794 )
9795 .await?;
9796 *self.pending_skill_id.write() = None;
9797 drop(admission);
9798 let skill_response = self.execute_skill_by_id(skill_id, enriched_input).await?;
9799 self.memory
9800 .add_message(ChatMessage::user(enriched_input))
9801 .await?;
9802 self.handle_skill_response(enriched_input, skill_id, skill_response, &HashMap::new())
9803 .await
9804 }
9805
9806 async fn handle_skill_response(
9809 &self,
9810 processed_input: &str,
9811 skill_id: &str,
9812 skill_response: String,
9813 input_context: &HashMap<String, Value>,
9814 ) -> Result<AgentResponse> {
9815 let output_data = self.process_output(&skill_response, input_context).await?;
9816 let final_response = output_data.content;
9817
9818 self.memory
9819 .add_message(ChatMessage::assistant(&final_response))
9820 .await?;
9821
9822 self.check_memory_compression().await?;
9823
9824 self.increment_turn();
9825 self.evaluate_transitions(processed_input, &final_response)
9826 .await?;
9827
9828 let response = AgentResponse::new(final_response)
9829 .with_metadata("skill_id", serde_json::json!(skill_id));
9830 self.finish_turn_if_root(&response).await?;
9831 Ok(response)
9832 }
9833
9834 async fn handle_plan_and_execute(
9837 &self,
9838 processed_input: &str,
9839 input_context: &HashMap<String, Value>,
9840 auto_detected: bool,
9841 ) -> Result<AgentResponse> {
9842 let effective = self.get_effective_reasoning_config();
9843 let plan_reflection = effective
9844 .get_planning()
9845 .map(|c| c.reflection.clone())
9846 .unwrap_or_default();
9847
9848 let max_attempts = if plan_reflection.enabled {
9849 1 + plan_reflection.max_replans
9850 } else {
9851 1
9852 };
9853
9854 let mut plan = self.generate_plan(processed_input).await?;
9855 info!(
9856 plan_id = %plan.id,
9857 steps = plan.steps.len(),
9858 "Plan generated"
9859 );
9860
9861 let mut plan_result = String::new();
9862
9863 for attempt in 0..max_attempts {
9864 *self.current_plan.write() = Some(plan.clone());
9865 plan_result = self.execute_plan(&mut plan).await?;
9866
9867 info!(
9868 plan_status = ?plan.status,
9869 completed_steps = plan.completed_steps().count(),
9870 attempt = attempt + 1,
9871 "Plan execution completed"
9872 );
9873
9874 if !plan_reflection.enabled {
9875 break;
9876 }
9877
9878 let has_failures = plan.steps.iter().any(|s| s.status.is_failed());
9879 if !has_failures {
9880 break;
9881 }
9882
9883 if attempt + 1 >= max_attempts {
9884 break;
9885 }
9886
9887 match plan_reflection.on_step_failure {
9888 StepFailureAction::Replan => {
9889 info!(attempt = attempt + 1, "Plan had failures, replanning");
9890 plan = self.generate_plan(processed_input).await?;
9891 }
9892 StepFailureAction::Abort => {
9893 warn!("Plan step failed, aborting");
9894 break;
9895 }
9896 StepFailureAction::Skip | StepFailureAction::Continue => {
9897 break;
9898 }
9899 }
9900 }
9901
9902 *self.current_plan.write() = Some(plan);
9903
9904 let output_data = self.process_output(&plan_result, input_context).await?;
9905 let final_content = output_data.content;
9906
9907 self.memory
9908 .add_message(ChatMessage::assistant(&final_content))
9909 .await?;
9910
9911 self.check_memory_compression().await?;
9912 self.increment_turn();
9913 self.evaluate_transitions(processed_input, &final_content)
9914 .await?;
9915
9916 let reasoning_metadata =
9917 ReasoningMetadata::new(ReasoningMode::PlanAndExecute).with_auto_detected(auto_detected);
9918
9919 let response = AgentResponse::new(&final_content).with_metadata(
9920 "reasoning",
9921 serde_json::to_value(&reasoning_metadata).unwrap_or_default(),
9922 );
9923
9924 self.finish_turn_if_root(&response).await?;
9925 Ok(response)
9926 }
9927
9928 fn inject_reasoning_prompt(
9930 &self,
9931 messages: &mut [ChatMessage],
9932 reasoning_mode: &ReasoningMode,
9933 is_first_iteration: bool,
9934 ) {
9935 if !is_first_iteration {
9936 return;
9937 }
9938 match reasoning_mode {
9939 ReasoningMode::CoT => {
9940 if let Some(msg) = messages.first_mut()
9941 && matches!(msg.role, ai_agents_core::Role::System)
9942 {
9943 msg.content = self.build_cot_system_prompt(&msg.content);
9944 debug!("Applied Chain-of-Thought system prompt");
9945 }
9946 }
9947 ReasoningMode::React => {
9948 if let Some(msg) = messages.first_mut()
9949 && matches!(msg.role, ai_agents_core::Role::System)
9950 {
9951 msg.content = self.build_react_system_prompt(&msg.content);
9952 debug!("Applied ReAct system prompt");
9953 }
9954 }
9955 _ => {}
9956 }
9957 }
9958
9959 async fn generate_main_response_draft(
9964 &self,
9965 processed_input: &str,
9966 reasoning_mode: &ReasoningMode,
9967 ) -> Result<MainResponseDraft> {
9968 let llm = self.get_state_llm()?;
9969 let protocol = self.main_tool_protocol(llm.as_ref(), true).await?;
9970 let mut messages = self
9971 .build_messages_internal(false, Some(processed_input), protocol.choice.is_none())
9972 .await?;
9973 self.inject_reasoning_prompt(&mut messages, reasoning_mode, true);
9974 let response = self
9975 .complete_main_llm_with_recovery(llm, &messages, &protocol)
9976 .await?;
9977 let content = response.content.trim().to_string();
9978 let (thinking, answer) = self.extract_thinking(&content);
9979 if let Some(calls) = self.parse_main_tool_calls(&content, &protocol)? {
9980 return Ok(MainResponseDraft::ToolCalls {
9981 raw_content: content,
9982 calls,
9983 thinking,
9984 });
9985 }
9986 Ok(MainResponseDraft::Text {
9987 raw_content: answer,
9988 thinking,
9989 })
9990 }
9991
9992 async fn commit_main_response_draft(
9997 &self,
9998 processed_input: &str,
9999 input_context: &HashMap<String, Value>,
10000 draft: MainResponseDraft,
10001 reasoning_mode: ReasoningMode,
10002 auto_detected: bool,
10003 ) -> Result<AgentResponse> {
10004 self.commit_root_user_message(processed_input).await?;
10005 match draft {
10006 MainResponseDraft::Text {
10007 raw_content,
10008 thinking,
10009 } => {
10010 self.finish_text_response_from_model(CommittedTextResponse {
10011 processed_input,
10012 input_context,
10013 answer: raw_content,
10014 reasoning_mode,
10015 auto_detected,
10016 iterations: 1,
10017 thinking_content: thinking,
10018 all_tool_calls: Vec::new(),
10019 })
10020 .await
10021 }
10022 MainResponseDraft::ToolCalls {
10023 raw_content,
10024 calls,
10025 thinking: _,
10026 } => {
10027 let mut all_tool_calls = Vec::new();
10028 match self
10029 .handle_tool_calls(
10030 processed_input,
10031 &raw_content,
10032 calls,
10033 &mut all_tool_calls,
10034 None,
10035 )
10036 .await?
10037 {
10038 ToolCallOutcome::Rejected(response) => {
10039 self.finish_turn_if_root(&response).await?;
10040 Ok(response)
10041 }
10042 ToolCallOutcome::Continue | ToolCallOutcome::TransitionFired => {
10043 self.continue_after_committed_tool_draft(processed_input)
10044 .await
10045 }
10046 }
10047 }
10048 }
10049 }
10050
10051 async fn continue_after_committed_tool_draft(
10056 &self,
10057 processed_input: &str,
10058 ) -> Result<AgentResponse> {
10059 *self.redispatch_depth.write() += 1;
10060 if let Some(context) = self.active_turn_context.write().as_mut() {
10061 context.enter_redispatch();
10062 }
10063 let result = Box::pin(self.run_loop_internal(processed_input)).await;
10064 *self.redispatch_depth.write() -= 1;
10065 if let Some(context) = self.active_turn_context.write().as_mut() {
10066 context.exit_redispatch();
10067 }
10068 let response = result?;
10069 self.finish_turn_if_root(&response).await?;
10070 Ok(response)
10071 }
10072
10073 async fn finish_text_response_from_model(
10078 &self,
10079 response: CommittedTextResponse<'_>,
10080 ) -> Result<AgentResponse> {
10081 let CommittedTextResponse {
10082 processed_input,
10083 input_context,
10084 answer,
10085 reasoning_mode,
10086 auto_detected,
10087 iterations,
10088 thinking_content,
10089 all_tool_calls,
10090 } = response;
10091 let output_data = self.process_output(&answer, input_context).await?;
10092 let mut final_content = if output_data.metadata.rejected {
10093 output_data
10094 .metadata
10095 .rejection_reason
10096 .unwrap_or_else(|| answer.to_string())
10097 } else {
10098 output_data.content
10099 };
10100 let llm = self.get_state_llm()?;
10101 let reflection_metadata;
10102 (final_content, reflection_metadata) = self
10103 .run_reflection(&*llm, processed_input, final_content)
10104 .await?;
10105 final_content =
10106 self.format_response_with_thinking(thinking_content.as_deref(), &final_content);
10107 let final_content = {
10108 let result = self
10109 .post_loop_processing(processed_input, final_content)
10110 .await?;
10111 self.apply_post_loop_result(processed_input, result)
10112 .await?
10113 .content
10114 };
10115 let response = self.build_agent_response(AgentResponseParts {
10116 content: final_content,
10117 all_tool_calls,
10118 reasoning_mode,
10119 auto_detected,
10120 iterations,
10121 thinking: thinking_content,
10122 reflection_metadata,
10123 });
10124 self.finish_turn_if_root(&response).await?;
10125 Ok(response)
10126 }
10127
10128 async fn run_committed_response_loop_with_reasoning(
10133 &self,
10134 processed_input: &str,
10135 input_context: &HashMap<String, Value>,
10136 reasoning_mode: ReasoningMode,
10137 auto_detected: bool,
10138 ) -> Result<AgentResponse> {
10139 self.commit_root_user_message(processed_input).await?;
10140 let llm = self.get_state_llm()?;
10141 let mut iterations = 0u32;
10142 let mut all_tool_calls = Vec::new();
10143 let mut thinking_content = None;
10144 loop {
10145 let effective_max = if reasoning_mode != ReasoningMode::None {
10146 let rc = self.get_effective_reasoning_config();
10147 self.max_iterations.min(rc.max_iterations)
10148 } else {
10149 self.max_iterations
10150 };
10151 if iterations >= effective_max {
10152 return Err(AgentError::Other(format!(
10153 "Max iterations ({}) exceeded",
10154 effective_max
10155 )));
10156 }
10157 iterations += 1;
10158 *self.iteration_count.write() = iterations;
10159 let protocol = self.main_tool_protocol(llm.as_ref(), false).await?;
10160 let mut messages = self
10161 .build_messages_internal(true, None, protocol.choice.is_none())
10162 .await?;
10163 self.inject_reasoning_prompt(&mut messages, &reasoning_mode, iterations == 1);
10164 self.hooks.on_llm_start(&messages).await;
10165 let llm_start = Instant::now();
10166 let response = self
10167 .complete_main_llm_with_recovery(Arc::clone(&llm), &messages, &protocol)
10168 .await?;
10169 let llm_duration_ms = llm_start.elapsed().as_millis() as u64;
10170 self.hooks.on_llm_complete(&response, llm_duration_ms).await;
10171 let content = response.content.trim();
10172 if let Some(tool_calls) = self.parse_main_tool_calls(content, &protocol)? {
10173 match self
10174 .handle_tool_calls(
10175 processed_input,
10176 content,
10177 tool_calls,
10178 &mut all_tool_calls,
10179 None,
10180 )
10181 .await?
10182 {
10183 ToolCallOutcome::Continue | ToolCallOutcome::TransitionFired => continue,
10184 ToolCallOutcome::Rejected(resp) => {
10185 self.finish_turn_if_root(&resp).await?;
10186 return Ok(resp);
10187 }
10188 }
10189 }
10190 let (extracted_thinking, answer) = self.extract_thinking(content);
10191 if extracted_thinking.is_some() {
10192 thinking_content = extracted_thinking;
10193 }
10194 return self
10195 .finish_text_response_from_model(CommittedTextResponse {
10196 processed_input,
10197 input_context,
10198 answer,
10199 reasoning_mode,
10200 auto_detected,
10201 iterations,
10202 thinking_content,
10203 all_tool_calls,
10204 })
10205 .await;
10206 }
10207 }
10208
10209 async fn handle_tool_calls(
10215 &self,
10216 processed_input: &str,
10217 content: &str,
10218 tool_calls: Vec<ToolCall>,
10219 all_tool_calls: &mut Vec<ToolCall>,
10220 mut events: Option<&mut Vec<StreamChunk>>,
10221 ) -> Result<ToolCallOutcome> {
10222 let include_tool_events = self.streaming.include_tool_events;
10223 let transition_content = native_readable_projection(content)
10227 .map_err(|error| AgentError::LLM(error.to_string()))?;
10228 let transition_fired = self
10229 .evaluate_transitions(processed_input, &transition_content)
10230 .await?;
10231 if transition_fired {
10232 self.memory
10233 .add_message(ChatMessage::assistant(
10234 "(Transitioned to new state — tool call handled by workflow)",
10235 ))
10236 .await?;
10237 if let Some(events) = events.as_deref_mut()
10238 && self.streaming.include_state_events
10239 && let Some(state) = self.current_state()
10240 {
10241 events.push(StreamChunk::state_transition(None, state));
10242 }
10243 return Ok(ToolCallOutcome::TransitionFired);
10244 }
10245
10246 self.memory
10248 .add_message(ChatMessage::assistant(content))
10249 .await?;
10250 self.remember_committed_native_exchange(content).await?;
10251 let native_tool_call = Self::is_native_tool_call_content(content)?;
10252
10253 if let Some(events) = events.as_deref_mut()
10254 && include_tool_events
10255 {
10256 for tool_call in &tool_calls {
10257 events.push(StreamChunk::tool_start(&tool_call.id, &tool_call.name));
10258 }
10259 }
10260 let results = self.execute_tools_parallel(&tool_calls).await;
10261 let mut rejection = None;
10262
10263 for ((_id, result), tool_call) in results.into_iter().zip(tool_calls.iter()) {
10264 match result {
10265 Ok(output) => {
10266 if let Some(events) = events.as_deref_mut()
10267 && include_tool_events
10268 {
10269 events.push(StreamChunk::tool_result(
10270 &tool_call.id,
10271 &tool_call.name,
10272 &output,
10273 true,
10274 ));
10275 }
10276 self.memory
10277 .add_message(Self::tool_result_message(
10278 tool_call,
10279 &output,
10280 native_tool_call,
10281 )?)
10282 .await?;
10283 }
10284 Err(e) => {
10285 if matches!(e, AgentError::HITLRejected(_)) {
10286 if !native_tool_call {
10287 self.memory
10288 .add_message(ChatMessage::assistant(format!(
10289 "The operation was rejected by the approver: {e}"
10290 )))
10291 .await?;
10292 return Ok(ToolCallOutcome::Rejected(AgentResponse {
10293 content: format!("Operation cancelled: {e}"),
10294 metadata: None,
10295 tool_calls: Some(all_tool_calls.clone()),
10296 }));
10297 }
10298 if rejection.is_none() {
10299 rejection = Some(e.to_string());
10300 }
10301 }
10302 if let Some(events) = events.as_deref_mut()
10303 && include_tool_events
10304 {
10305 events.push(StreamChunk::tool_result(
10306 &tool_call.id,
10307 &tool_call.name,
10308 e.to_string(),
10309 false,
10310 ));
10311 }
10312 self.memory
10313 .add_message(Self::tool_result_message(
10314 tool_call,
10315 &format!("Error: {}", e),
10316 native_tool_call,
10317 )?)
10318 .await?;
10319 }
10320 }
10321 all_tool_calls.push(tool_call.clone());
10322 if let Some(events) = events.as_deref_mut()
10323 && include_tool_events
10324 {
10325 events.push(StreamChunk::tool_end(&tool_call.id));
10326 }
10327 }
10328 if let Some(rejection) = rejection {
10329 self.memory
10330 .add_message(ChatMessage::assistant(format!(
10331 "The operation was rejected by the approver: {rejection}"
10332 )))
10333 .await?;
10334 return Ok(ToolCallOutcome::Rejected(AgentResponse {
10335 content: format!("Operation cancelled: {rejection}"),
10336 metadata: None,
10337 tool_calls: Some(all_tool_calls.clone()),
10338 }));
10339 }
10340 Ok(ToolCallOutcome::Continue)
10341 }
10342
10343 async fn run_reflection(
10345 &self,
10346 llm: &dyn LLMProvider,
10347 processed_input: &str,
10348 mut content: String,
10349 ) -> Result<(String, Option<ReflectionMetadata>)> {
10350 let should_reflect = self.should_reflect(processed_input, &content).await?;
10351 if !should_reflect {
10352 return Ok((content, None));
10353 }
10354
10355 info!("Starting response reflection evaluation");
10356 let mut attempts = 0u32;
10357 let max_retries = self.reflection_config.max_retries;
10358 let mut history: Vec<ReflectionAttempt> = Vec::new();
10359
10360 loop {
10361 let evaluation = self.evaluate_response(processed_input, &content).await?;
10362
10363 if evaluation.passed || attempts >= max_retries {
10364 info!(
10365 passed = evaluation.passed,
10366 confidence = evaluation.confidence,
10367 attempts = attempts + 1,
10368 "Reflection evaluation complete"
10369 );
10370 let reflection_metadata = Some(
10371 ReflectionMetadata::new(evaluation)
10372 .with_attempts(attempts + 1)
10373 .with_history(history),
10374 );
10375 return Ok((content, reflection_metadata));
10376 }
10377
10378 debug!(
10379 attempt = attempts + 1,
10380 failed_criteria = evaluation.failed_criteria().count(),
10381 "Response did not meet criteria, retrying"
10382 );
10383
10384 history.push(
10385 ReflectionAttempt::new(&content, evaluation.clone())
10386 .with_feedback("Response did not meet quality criteria"),
10387 );
10388
10389 let feedback: Vec<String> = evaluation
10390 .failed_criteria()
10391 .map(|c| format!("- {}", c.criterion))
10392 .collect();
10393
10394 let retry_prompt = format!(
10395 "Your previous response did not meet these criteria:\n{}\n\nPlease provide an improved response.",
10396 feedback.join("\n")
10397 );
10398
10399 self.memory
10400 .add_message(ChatMessage::user(&retry_prompt))
10401 .await?;
10402
10403 let retry_messages = self.build_messages().await?;
10404 let retry_response = self
10405 .observe_purpose(
10406 ObservationPurpose::ReflectionEvaluation,
10407 llm.complete(&retry_messages, None),
10408 )
10409 .await
10410 .map_err(|e| AgentError::LLM(e.to_string()))?;
10411
10412 content = retry_response.content.trim().to_string();
10413 attempts += 1;
10414 }
10415 }
10416
10417 async fn post_loop_processing(
10420 &self,
10421 processed_input: &str,
10422 content: String,
10423 ) -> Result<PostLoopResult> {
10424 self.increment_turn();
10429
10430 self.run_context_extractors(processed_input).await;
10432
10433 let transitioned = self.evaluate_transitions(processed_input, &content).await?;
10434
10435 if !transitioned {
10436 self.memory
10437 .add_message(ChatMessage::assistant(&content))
10438 .await?;
10439 self.check_memory_compression().await?;
10440 return Ok(PostLoopResult::NoTransition(content));
10441 }
10442
10443 if !self.should_regenerate_after_transition() {
10445 self.memory
10446 .add_message(ChatMessage::assistant(&content))
10447 .await?;
10448 self.check_memory_compression().await?;
10449 return Ok(PostLoopResult::Transitioned {
10450 content,
10451 regenerated: false,
10452 });
10453 }
10454
10455 if self.needs_redispatch_for_new_state() {
10459 info!("Post-transition NeedsRedispatch: new state requires full dispatch");
10460 return Ok(PostLoopResult::NeedsRedispatch);
10463 }
10464
10465 self.memory
10468 .add_message(ChatMessage::assistant(&content))
10469 .await?;
10470 self.check_memory_compression().await?;
10471
10472 let new_llm = self.get_state_llm()?;
10478 let mut final_content;
10479
10480 for post_iter in 0..self.max_iterations {
10481 let protocol = self.main_tool_protocol(new_llm.as_ref(), false).await?;
10482 let new_messages = self
10483 .build_messages_internal(true, None, protocol.choice.is_none())
10484 .await?;
10485 if post_iter == 0
10486 && let Some(system_msg) = new_messages.first()
10487 && system_msg.role == ai_agents_core::Role::System
10488 {
10489 debug!(
10490 prompt_preview =
10491 &system_msg.content[system_msg.content.len().saturating_sub(200)..],
10492 "Post-transition system prompt (last 200 chars)"
10493 );
10494 }
10495
10496 let new_response = self
10497 .complete_main_llm_with_recovery(Arc::clone(&new_llm), &new_messages, &protocol)
10498 .await?;
10499 final_content = new_response.content.trim().to_string();
10500
10501 if let Some(tool_calls) = self.parse_main_tool_calls(&final_content, &protocol)? {
10504 let native_tool_call = Self::is_native_tool_call_content(&final_content)?;
10505 debug!(
10506 post_iter = post_iter,
10507 tools = tool_calls.len(),
10508 "Post-transition tool call detected, executing"
10509 );
10510
10511 self.memory
10512 .add_message(ChatMessage::assistant(&final_content))
10513 .await?;
10514 self.remember_committed_native_exchange(&final_content)
10515 .await?;
10516
10517 let results = self.execute_tools_parallel(&tool_calls).await;
10518 let mut rejection = None;
10519 for ((_id, result), tool_call) in results.into_iter().zip(tool_calls.iter()) {
10520 match result {
10521 Ok(output) => {
10522 self.memory
10523 .add_message(Self::tool_result_message(
10524 tool_call,
10525 &output,
10526 native_tool_call,
10527 )?)
10528 .await?;
10529 }
10530 Err(e) => {
10531 if native_tool_call
10532 && rejection.is_none()
10533 && matches!(e, AgentError::HITLRejected(_))
10534 {
10535 rejection = Some(e.to_string());
10536 }
10537 self.memory
10538 .add_message(Self::tool_result_message(
10539 tool_call,
10540 &format!("Error: {}", e),
10541 native_tool_call,
10542 )?)
10543 .await?;
10544 }
10545 }
10546 }
10547 if let Some(rejection) = rejection {
10548 self.memory
10549 .add_message(ChatMessage::assistant(format!(
10550 "The operation was rejected by the approver: {rejection}"
10551 )))
10552 .await?;
10553 return Err(AgentError::HITLRejected(rejection));
10554 }
10555 continue;
10557 }
10558
10559 self.memory
10561 .add_message(ChatMessage::assistant(&final_content))
10562 .await?;
10563 return Ok(PostLoopResult::Transitioned {
10564 content: final_content,
10565 regenerated: true,
10566 });
10567 }
10568
10569 final_content = "Post-transition processing completed.".to_string();
10571 self.memory
10572 .add_message(ChatMessage::assistant(&final_content))
10573 .await?;
10574
10575 Ok(PostLoopResult::Transitioned {
10576 content: final_content,
10577 regenerated: true,
10578 })
10579 }
10580
10581 fn should_regenerate_after_transition(&self) -> bool {
10584 if let Some(ref sm) = self.state_machine {
10585 if !sm.config().regenerate_on_transition {
10587 return false;
10588 }
10589 if let Some(def) = sm.current_definition()
10591 && let Some(regen) = def.regenerate_on_enter
10592 {
10593 return regen;
10594 }
10595 }
10596 true
10597 }
10598
10599 fn needs_redispatch_for_new_state(&self) -> bool {
10602 if let Some(ref sm) = self.state_machine
10603 && let Some(def) = sm.current_definition()
10604 {
10605 if def.concurrent.is_some()
10606 || def.group_chat.is_some()
10607 || def.pipeline.is_some()
10608 || def.handoff.is_some()
10609 || def.delegate.is_some()
10610 {
10611 return true;
10612 }
10613 let effective = self.get_effective_reasoning_config();
10615 if !matches!(effective.mode, ReasoningMode::None) {
10616 return true;
10617 }
10618 }
10619 false
10620 }
10621
10622 async fn apply_post_loop_result(
10628 &self,
10629 processed_input: &str,
10630 result: PostLoopResult,
10631 ) -> Result<AppliedPostLoop> {
10632 match result {
10633 PostLoopResult::NoTransition(content) => Ok(AppliedPostLoop {
10634 content,
10635 transitioned: false,
10636 regenerated: false,
10637 }),
10638 PostLoopResult::Transitioned {
10639 content,
10640 regenerated,
10641 } => Ok(AppliedPostLoop {
10642 content,
10643 transitioned: true,
10644 regenerated,
10645 }),
10646 PostLoopResult::NeedsRedispatch => {
10647 const MAX_REDISPATCH_DEPTH: u32 = 3;
10648 let current_depth = *self.redispatch_depth.read();
10649 if current_depth >= MAX_REDISPATCH_DEPTH {
10650 warn!(
10651 depth = current_depth,
10652 "Post-transition re-dispatch depth limit reached, returning empty response"
10653 );
10654 let content = String::new();
10655 self.memory
10656 .add_message(ChatMessage::assistant(&content))
10657 .await?;
10658 return Ok(AppliedPostLoop {
10660 content,
10661 transitioned: true,
10662 regenerated: false,
10663 });
10664 }
10665 *self.redispatch_depth.write() += 1;
10666 if let Some(context) = self.active_turn_context.write().as_mut() {
10667 context.enter_redispatch();
10668 }
10669 info!(
10670 depth = current_depth + 1,
10671 "Re-dispatching for new state after transition"
10672 );
10673 let resp = Box::pin(self.run_loop_internal(processed_input)).await;
10674 *self.redispatch_depth.write() -= 1;
10675 if let Some(context) = self.active_turn_context.write().as_mut() {
10676 context.exit_redispatch();
10677 }
10678 resp.map(|r| AppliedPostLoop {
10679 content: r.content,
10680 transitioned: true,
10681 regenerated: true,
10682 })
10683 }
10684 }
10685 }
10686
10687 fn build_agent_response(&self, parts: AgentResponseParts) -> AgentResponse {
10689 let AgentResponseParts {
10690 content,
10691 all_tool_calls,
10692 reasoning_mode,
10693 auto_detected,
10694 iterations,
10695 thinking,
10696 reflection_metadata,
10697 } = parts;
10698 let reasoning_metadata = ReasoningMetadata::new(reasoning_mode.clone())
10699 .with_thinking(thinking.clone().unwrap_or_default())
10700 .with_iterations(iterations)
10701 .with_auto_detected(auto_detected);
10702
10703 let mut response = AgentResponse::new(&content);
10704 if !all_tool_calls.is_empty() {
10705 response = response.with_tool_calls(all_tool_calls);
10706 }
10707
10708 if let Some(state) = self.current_state() {
10709 response = response.with_metadata("current_state", serde_json::json!(state));
10710 }
10711
10712 response = response.with_metadata(
10713 "reasoning",
10714 serde_json::to_value(&reasoning_metadata).unwrap_or_default(),
10715 );
10716
10717 if let Some(ref refl_meta) = reflection_metadata {
10718 response = response.with_metadata(
10719 "reflection",
10720 serde_json::to_value(refl_meta).unwrap_or_default(),
10721 );
10722 }
10723
10724 response
10725 }
10726
10727 async fn handle_delegated_state(
10729 &self,
10730 input: &str,
10731 delegate_id: &str,
10732 state_def: &ai_agents_state::StateDefinition,
10733 ) -> Result<AgentResponse> {
10734 use std::time::Instant;
10735
10736 let registry = self.spawner_registry.as_ref().ok_or_else(|| {
10737 AgentError::Config(format!(
10738 "State delegates to '{}' but no agent registry is configured. \
10739 Add a spawner section with auto_spawn to your YAML.",
10740 delegate_id
10741 ))
10742 })?;
10743
10744 let state_name = self
10745 .state_machine
10746 .as_ref()
10747 .map(|sm| sm.current())
10748 .unwrap_or_else(|| "unknown".to_string());
10749
10750 self.hooks.on_delegate_start(delegate_id, &state_name).await;
10751 let start = Instant::now();
10752
10753 let delegate = registry.get(delegate_id).ok_or_else(|| {
10754 AgentError::Other(format!(
10755 "State '{}' delegates to '{}' but no agent with that ID exists in the registry.",
10756 state_name, delegate_id
10757 ))
10758 })?;
10759
10760 let context_mode = state_def.delegate_context.clone().unwrap_or_default();
10762 let effective_input = self
10763 .observe_purpose(
10764 ObservationPurpose::OrchestrationRouting,
10765 crate::orchestration::context::prepare_delegate_input(
10766 input,
10767 &context_mode,
10768 &*self.memory,
10769 self.llm_registry.get("router").ok().as_deref(),
10770 ),
10771 )
10772 .await?;
10773
10774 let response = delegate
10775 .chat_with_actor_context(&effective_input, self.outbound_actor_context())
10776 .await?;
10777
10778 let duration_ms = start.elapsed().as_millis() as u64;
10779 self.hooks
10780 .on_delegate_complete(delegate_id, &state_name, duration_ms)
10781 .await;
10782
10783 let ctx_key = format!("delegation.{}.last_response", delegate_id);
10785 let _ = self.context_manager.set(
10786 &ctx_key,
10787 serde_json::Value::String(response.content.clone()),
10788 );
10789
10790 let _ = self.context_manager.set(
10792 "orchestration",
10793 serde_json::json!({
10794 "type": "delegate",
10795 "agent": delegate_id,
10796 "state": state_name,
10797 "response": response.content,
10798 "duration_ms": duration_ms,
10799 }),
10800 );
10801
10802 self.commit_root_user_message(input).await?;
10803
10804 let post_result = self
10807 .post_loop_processing(
10808 input,
10809 format!("[Delegated to {}]: {}", delegate_id, response.content),
10810 )
10811 .await?;
10812 let final_content = self
10813 .apply_post_loop_result(input, post_result)
10814 .await?
10815 .content;
10816
10817 let mut result = AgentResponse::new(final_content);
10818
10819 let metadata = serde_json::json!({
10820 "orchestration": {
10821 "type": "delegate",
10822 "agent": delegate_id,
10823 "state": state_name,
10824 "response": response.content,
10825 "duration_ms": duration_ms,
10826 }
10827 });
10828 result.metadata = Some(
10829 serde_json::from_value::<std::collections::HashMap<String, serde_json::Value>>(
10830 metadata,
10831 )
10832 .unwrap_or_default(),
10833 );
10834
10835 self.finish_turn_if_root(&result).await?;
10836 Ok(result)
10837 }
10838
10839 async fn handle_concurrent_state(
10841 &self,
10842 input: &str,
10843 config: &ai_agents_state::ConcurrentStateConfig,
10844 ) -> Result<AgentResponse> {
10845 use std::time::Instant;
10846
10847 let registry = self.spawner_registry.as_ref().ok_or_else(|| {
10848 AgentError::Config(
10849 "Concurrent state requires an agent registry. Add a spawner section.".into(),
10850 )
10851 })?;
10852
10853 let context_mode = config.context_mode.clone().unwrap_or_default();
10858 let context_input = self
10859 .observe_purpose(
10860 ObservationPurpose::OrchestrationRouting,
10861 crate::orchestration::context::prepare_delegate_input(
10862 input,
10863 &context_mode,
10864 &*self.memory,
10865 self.llm_registry.get("router").ok().as_deref(),
10866 ),
10867 )
10868 .await?;
10869
10870 let effective_input = if let Some(ref tmpl) = config.input {
10871 render_concurrent_template(tmpl, &context_input, &self.build_context_with_overlays())
10872 .unwrap_or_else(|_| context_input.clone())
10873 } else {
10874 context_input
10875 };
10876
10877 let start = Instant::now();
10878
10879 let llm_name = config
10880 .aggregation
10881 .synthesizer_llm
10882 .as_deref()
10883 .unwrap_or("router");
10884 let llm_provider = self.llm_registry.get(llm_name).ok();
10885
10886 let vote_parallelism = if self.runtime_config.optimization.enabled
10887 && self
10888 .runtime_config
10889 .optimization
10890 .parallel_orchestration_vote_extraction
10891 {
10892 Some(self.runtime_config.optimization.max_parallel_runtime_tasks)
10893 } else {
10894 None
10895 };
10896
10897 let result = self
10898 .observe_purpose(
10899 ObservationPurpose::OrchestrationAggregation,
10900 scope_actor_context(
10901 self.outbound_actor_context(),
10902 crate::orchestration::concurrent(
10903 registry,
10904 &effective_input,
10905 &config.agents,
10906 &config.aggregation,
10907 llm_provider.as_deref(),
10908 config.min_required,
10909 config.timeout_ms,
10910 config.on_partial_failure.clone(),
10911 vote_parallelism,
10912 ),
10913 ),
10914 )
10915 .await?;
10916
10917 let duration_ms = start.elapsed().as_millis() as u64;
10918 let agent_ids: Vec<String> = config.agents.iter().map(|a| a.id().to_string()).collect();
10919 let strategy = format!("{:?}", config.aggregation.strategy);
10920 self.hooks
10921 .on_concurrent_complete(&agent_ids, &strategy, duration_ms)
10922 .await;
10923
10924 let _ = self.context_manager.set(
10926 "concurrent.result",
10927 serde_json::Value::String(result.response.content.clone()),
10928 );
10929
10930 let agents_json: Vec<serde_json::Value> = result
10932 .agent_results
10933 .iter()
10934 .map(|ar| {
10935 serde_json::json!({
10936 "id": ar.agent_id,
10937 "response": ar.response.as_ref().map(|r| r.content.as_str()),
10938 "success": ar.success,
10939 "error": ar.error,
10940 "duration_ms": ar.duration_ms,
10941 })
10942 })
10943 .collect();
10944
10945 let _ = self.context_manager.set(
10947 "orchestration",
10948 serde_json::json!({
10949 "type": "concurrent",
10950 "result": result.response.content,
10951 "strategy": strategy,
10952 "agents": agents_json,
10953 "duration_ms": duration_ms,
10954 }),
10955 );
10956
10957 self.commit_root_user_message(input).await?;
10958
10959 let post_result = self
10960 .post_loop_processing(input, result.response.content.clone())
10961 .await?;
10962 let final_content = self
10963 .apply_post_loop_result(input, post_result)
10964 .await?
10965 .content;
10966
10967 let mut response = AgentResponse::new(final_content);
10968 let metadata = serde_json::json!({
10969 "orchestration": {
10970 "type": "concurrent",
10971 "result": result.response.content,
10972 "strategy": strategy,
10973 "agents": agents_json,
10974 "duration_ms": duration_ms,
10975 }
10976 });
10977 response.metadata = Some(
10978 serde_json::from_value::<std::collections::HashMap<String, serde_json::Value>>(
10979 metadata,
10980 )
10981 .unwrap_or_default(),
10982 );
10983
10984 self.finish_turn_if_root(&response).await?;
10985 Ok(response)
10986 }
10987
10988 async fn handle_group_chat_state(
10990 &self,
10991 input: &str,
10992 config: &ai_agents_state::GroupChatStateConfig,
10993 ) -> Result<AgentResponse> {
10994 use std::time::Instant;
10995
10996 let registry = self.spawner_registry.as_ref().ok_or_else(|| {
10997 AgentError::Config(
10998 "Group chat state requires an agent registry. Add a spawner section.".into(),
10999 )
11000 })?;
11001
11002 let start = Instant::now();
11003
11004 let llm_provider = self.llm_registry.get("router").ok();
11005
11006 let context_mode = config.context_mode.clone().unwrap_or_default();
11008 let context_input = self
11009 .observe_purpose(
11010 ObservationPurpose::OrchestrationRouting,
11011 crate::orchestration::context::prepare_delegate_input(
11012 input,
11013 &context_mode,
11014 &*self.memory,
11015 self.llm_registry.get("router").ok().as_deref(),
11016 ),
11017 )
11018 .await?;
11019
11020 let effective_topic = if let Some(ref tmpl) = config.input {
11022 render_concurrent_template(tmpl, &context_input, &self.build_context_with_overlays())
11023 .unwrap_or_else(|_| context_input.clone())
11024 } else {
11025 context_input
11026 };
11027
11028 let result = self
11029 .observe_purpose(
11030 ObservationPurpose::OrchestrationConversation,
11031 scope_actor_context(
11032 self.outbound_actor_context(),
11033 crate::orchestration::group_chat(
11034 registry,
11035 &effective_topic,
11036 config,
11037 llm_provider.as_deref(),
11038 Some(&*self.hooks),
11039 ),
11040 ),
11041 )
11042 .await?;
11043
11044 let duration_ms = start.elapsed().as_millis() as u64;
11045
11046 let _ = self.context_manager.set(
11048 "group_chat.conclusion",
11049 serde_json::Value::String(result.response.content.clone()),
11050 );
11051
11052 let transcript_json: Vec<serde_json::Value> = result
11054 .transcript
11055 .iter()
11056 .map(|t| {
11057 serde_json::json!({
11058 "speaker": t.speaker,
11059 "round": t.round,
11060 "content": t.content,
11061 })
11062 })
11063 .collect();
11064
11065 let _ = self.context_manager.set(
11067 "orchestration",
11068 serde_json::json!({
11069 "type": "group_chat",
11070 "conclusion": result.response.content,
11071 "transcript": transcript_json,
11072 "rounds": result.rounds_completed,
11073 "termination": result.termination_reason,
11074 "duration_ms": duration_ms,
11075 }),
11076 );
11077
11078 self.commit_root_user_message(input).await?;
11079
11080 let post_result = self
11081 .post_loop_processing(input, result.response.content.clone())
11082 .await?;
11083 let final_content = self
11084 .apply_post_loop_result(input, post_result)
11085 .await?
11086 .content;
11087
11088 let mut response = AgentResponse::new(final_content);
11089 let metadata = serde_json::json!({
11090 "orchestration": {
11091 "type": "group_chat",
11092 "conclusion": result.response.content,
11093 "transcript": transcript_json,
11094 "rounds": result.rounds_completed,
11095 "termination": result.termination_reason,
11096 "duration_ms": duration_ms,
11097 }
11098 });
11099 response.metadata = Some(
11100 serde_json::from_value::<std::collections::HashMap<String, serde_json::Value>>(
11101 metadata,
11102 )
11103 .unwrap_or_default(),
11104 );
11105
11106 self.finish_turn_if_root(&response).await?;
11107 Ok(response)
11108 }
11109
11110 async fn handle_pipeline_state(
11112 &self,
11113 input: &str,
11114 config: &ai_agents_state::PipelineStateConfig,
11115 ) -> Result<AgentResponse> {
11116 use std::time::Instant;
11117
11118 let registry = self.spawner_registry.as_ref().ok_or_else(|| {
11119 AgentError::Config(
11120 "Pipeline state requires an agent registry. Add a spawner section.".into(),
11121 )
11122 })?;
11123
11124 let start = Instant::now();
11125
11126 let stages: Vec<crate::orchestration::PipelineStage> = config
11127 .stages
11128 .iter()
11129 .map(|entry| {
11130 let mut stage = crate::orchestration::PipelineStage::id(entry.id());
11131 if let Some(tmpl) = entry.input() {
11132 stage = stage.with_input(tmpl);
11133 }
11134 stage
11135 })
11136 .collect();
11137
11138 let context_mode = config.context_mode.clone().unwrap_or_default();
11140 let context_input = self
11141 .observe_purpose(
11142 ObservationPurpose::OrchestrationRouting,
11143 crate::orchestration::context::prepare_delegate_input(
11144 input,
11145 &context_mode,
11146 &*self.memory,
11147 self.llm_registry.get("router").ok().as_deref(),
11148 ),
11149 )
11150 .await?;
11151
11152 let context_values = self.build_context_with_overlays();
11153 let result = self
11154 .observe_purpose(
11155 ObservationPurpose::OrchestrationRouting,
11156 scope_actor_context(
11157 self.outbound_actor_context(),
11158 crate::orchestration::pipeline(
11159 registry,
11160 &context_input,
11161 &stages,
11162 config.timeout_ms,
11163 Some(&*self.hooks),
11164 Some(&context_values),
11165 ),
11166 ),
11167 )
11168 .await?;
11169
11170 let duration_ms = start.elapsed().as_millis() as u64;
11171
11172 let _ = self.context_manager.set(
11174 "pipeline.result",
11175 serde_json::Value::String(result.response.content.clone()),
11176 );
11177
11178 let stages_json: Vec<serde_json::Value> = result
11180 .stage_outputs
11181 .iter()
11182 .map(|s| {
11183 serde_json::json!({
11184 "agent_id": s.agent_id,
11185 "output": s.output,
11186 "duration_ms": s.duration_ms,
11187 "skipped": s.skipped,
11188 })
11189 })
11190 .collect();
11191
11192 let _ = self.context_manager.set(
11194 "orchestration",
11195 serde_json::json!({
11196 "type": "pipeline",
11197 "result": result.response.content,
11198 "stages": stages_json,
11199 "duration_ms": duration_ms,
11200 }),
11201 );
11202
11203 self.commit_root_user_message(input).await?;
11204
11205 let post_result = self
11206 .post_loop_processing(input, result.response.content.clone())
11207 .await?;
11208 let final_content = self
11209 .apply_post_loop_result(input, post_result)
11210 .await?
11211 .content;
11212
11213 let mut response = AgentResponse::new(final_content);
11214 let metadata = serde_json::json!({
11215 "orchestration": {
11216 "type": "pipeline",
11217 "result": result.response.content,
11218 "stages": stages_json,
11219 "duration_ms": duration_ms,
11220 }
11221 });
11222 response.metadata = Some(
11223 serde_json::from_value::<std::collections::HashMap<String, serde_json::Value>>(
11224 metadata,
11225 )
11226 .unwrap_or_default(),
11227 );
11228
11229 self.finish_turn_if_root(&response).await?;
11230 Ok(response)
11231 }
11232
11233 async fn handle_handoff_state(
11235 &self,
11236 input: &str,
11237 config: &ai_agents_state::HandoffStateConfig,
11238 ) -> Result<AgentResponse> {
11239 use std::time::Instant;
11240
11241 let registry = self.spawner_registry.as_ref().ok_or_else(|| {
11242 AgentError::Config(
11243 "Handoff state requires an agent registry. Add a spawner section.".into(),
11244 )
11245 })?;
11246
11247 let llm = self
11248 .llm_registry
11249 .get("router")
11250 .map_err(|_| AgentError::Config("Handoff state requires a router LLM.".into()))?;
11251
11252 let start = Instant::now();
11253
11254 let context_mode = config.context_mode.clone().unwrap_or_default();
11256 let context_input = self
11257 .observe_purpose(
11258 ObservationPurpose::OrchestrationRouting,
11259 crate::orchestration::context::prepare_delegate_input(
11260 input,
11261 &context_mode,
11262 &*self.memory,
11263 self.llm_registry.get("router").ok().as_deref(),
11264 ),
11265 )
11266 .await?;
11267
11268 let effective_input = if let Some(ref tmpl) = config.input {
11270 render_concurrent_template(tmpl, &context_input, &self.build_context_with_overlays())
11271 .unwrap_or_else(|_| context_input.clone())
11272 } else {
11273 context_input
11274 };
11275
11276 let result = self
11277 .observe_purpose(
11278 ObservationPurpose::OrchestrationRouting,
11279 scope_actor_context(
11280 self.outbound_actor_context(),
11281 crate::orchestration::handoff(
11282 registry,
11283 &effective_input,
11284 &config.initial_agent,
11285 &config.available_agents,
11286 config.max_handoffs,
11287 llm.as_ref(),
11288 Some(&*self.hooks),
11289 ),
11290 ),
11291 )
11292 .await?;
11293
11294 let duration_ms = start.elapsed().as_millis() as u64;
11295
11296 let _ = self.context_manager.set(
11298 "handoff.result",
11299 serde_json::Value::String(result.response.content.clone()),
11300 );
11301
11302 let chain_json: Vec<serde_json::Value> = result
11304 .handoff_chain
11305 .iter()
11306 .map(|h| {
11307 serde_json::json!({
11308 "from": h.from_agent,
11309 "to": h.to_agent,
11310 "reason": h.reason,
11311 })
11312 })
11313 .collect();
11314
11315 let _ = self.context_manager.set(
11317 "orchestration",
11318 serde_json::json!({
11319 "type": "handoff",
11320 "result": result.response.content,
11321 "final_agent": result.final_agent,
11322 "handoff_chain": chain_json,
11323 "duration_ms": duration_ms,
11324 }),
11325 );
11326
11327 self.commit_root_user_message(input).await?;
11328
11329 let post_result = self
11330 .post_loop_processing(input, result.response.content.clone())
11331 .await?;
11332 let final_content = self
11333 .apply_post_loop_result(input, post_result)
11334 .await?
11335 .content;
11336
11337 let mut response = AgentResponse::new(final_content);
11338 let metadata = serde_json::json!({
11339 "orchestration": {
11340 "type": "handoff",
11341 "result": result.response.content,
11342 "final_agent": result.final_agent,
11343 "handoff_chain": chain_json,
11344 "duration_ms": duration_ms,
11345 }
11346 });
11347 response.metadata = Some(
11348 serde_json::from_value::<std::collections::HashMap<String, serde_json::Value>>(
11349 metadata,
11350 )
11351 .unwrap_or_default(),
11352 );
11353
11354 self.finish_turn_if_root(&response).await?;
11355 Ok(response)
11356 }
11357
11358 async fn run_loop_internal(&self, input: &str) -> Result<AgentResponse> {
11360 self.begin_root_turn();
11361 self.pre_turn_session_lifecycle().await;
11363
11364 let input_data = self.process_input(input).await?;
11365 self.update_active_turn_context(&input_data.content, input_data.context.clone());
11366
11367 for (key, value) in &input_data.context {
11370 let _ = self.context_manager.set(key, value.clone());
11371 }
11372
11373 if input_data.metadata.rejected {
11374 let reason = input_data
11375 .metadata
11376 .rejection_reason
11377 .unwrap_or_else(|| "Input rejected".to_string());
11378 warn!(reason = %reason, "Input rejected");
11379 let response = AgentResponse::new(reason);
11380 self.finish_turn_if_root(&response).await?;
11381 return Ok(response);
11382 }
11383
11384 let processed_input = &input_data.content;
11385
11386 if let Some(response) = self.try_pre_response_transition(processed_input).await? {
11387 return Ok(response);
11388 }
11389
11390 if let Some(ref sm) = self.state_machine
11392 && let Some(def) = sm.current_definition()
11393 {
11394 if let Some(ref delegate_id) = def.delegate {
11395 return self
11396 .handle_delegated_state(processed_input, delegate_id, &def)
11397 .await;
11398 }
11399 if let Some(ref concurrent_config) = def.concurrent {
11400 return self
11401 .handle_concurrent_state(processed_input, concurrent_config)
11402 .await;
11403 }
11404 if let Some(ref group_chat_config) = def.group_chat {
11405 return self
11406 .handle_group_chat_state(processed_input, group_chat_config)
11407 .await;
11408 }
11409 if let Some(ref pipeline_config) = def.pipeline {
11410 return self
11411 .handle_pipeline_state(processed_input, pipeline_config)
11412 .await;
11413 }
11414 if let Some(ref handoff_config) = def.handoff {
11415 return self
11416 .handle_handoff_state(processed_input, handoff_config)
11417 .await;
11418 }
11419 }
11420
11421 if let Some(response) =
11426 Box::pin(self.try_speculative_branches(processed_input, &input_data.context)).await?
11427 {
11428 return Ok(response);
11429 }
11430
11431 match self.try_skill_route(processed_input).await? {
11432 SkillRouteResult::Response { skill_id, content } => {
11433 self.commit_root_user_message(processed_input).await?;
11434 return self
11435 .handle_skill_response(processed_input, &skill_id, content, &input_data.context)
11436 .await;
11437 }
11438 SkillRouteResult::NeedsClarification {
11439 response,
11440 ownership,
11441 } => {
11442 let admission = self
11443 .admit_optional_disambiguation_ownership(ownership)
11444 .await?;
11445 self.commit_root_user_message(processed_input).await?;
11446 if Self::skill_clarification_needs_memory_record(&response) {
11447 self.memory
11450 .add_message(ChatMessage::assistant(&response.content))
11451 .await?;
11452 }
11453 drop(admission);
11454 self.finish_turn_if_root(&response).await?;
11455 return Ok(response);
11456 }
11457 SkillRouteResult::NoMatch => {} }
11459
11460 let effective_reasoning = self.get_effective_reasoning_config();
11461 let reasoning_mode = self.determine_reasoning_mode(processed_input).await?;
11462 let auto_detected = matches!(effective_reasoning.mode, ReasoningMode::Auto);
11463
11464 info!(
11465 reasoning_mode = ?reasoning_mode,
11466 auto_detected = auto_detected,
11467 reflection_enabled = ?self.reflection_config.enabled,
11468 "Reasoning mode determined"
11469 );
11470
11471 if matches!(reasoning_mode, ReasoningMode::PlanAndExecute) {
11472 self.commit_root_user_message(processed_input).await?;
11473 return self
11474 .handle_plan_and_execute(processed_input, &input_data.context, auto_detected)
11475 .await;
11476 }
11477
11478 self.commit_root_user_message(processed_input).await?;
11479
11480 let mut iterations = 0u32;
11481 let mut all_tool_calls: Vec<ToolCall> = Vec::new();
11482 let mut thinking_content: Option<String> = None;
11483
11484 let llm = self.get_state_llm()?;
11485
11486 loop {
11487 let effective_max = if reasoning_mode != ReasoningMode::None {
11489 let rc = self.get_effective_reasoning_config();
11490 self.max_iterations.min(rc.max_iterations)
11491 } else {
11492 self.max_iterations
11493 };
11494
11495 if iterations >= effective_max {
11496 let err = AgentError::Other(format!("Max iterations ({}) exceeded", effective_max));
11497 self.hooks.on_error(&err).await;
11498 error!(iterations = iterations, "Max iterations exceeded");
11499 return Err(err);
11500 }
11501 iterations += 1;
11502 *self.iteration_count.write() = iterations;
11503
11504 debug!(iteration = iterations, max = effective_max, "LLM call");
11505
11506 let protocol = self.main_tool_protocol(llm.as_ref(), false).await?;
11507 let mut messages = self
11508 .build_messages_internal(true, None, protocol.choice.is_none())
11509 .await?;
11510 self.inject_reasoning_prompt(&mut messages, &reasoning_mode, iterations == 1);
11511
11512 self.hooks.on_llm_start(&messages).await;
11513 let llm_start = Instant::now();
11514 let response = self
11515 .complete_main_llm_with_recovery(Arc::clone(&llm), &messages, &protocol)
11516 .await?;
11517
11518 let llm_duration_ms = llm_start.elapsed().as_millis() as u64;
11519 self.hooks.on_llm_complete(&response, llm_duration_ms).await;
11520
11521 let content = response.content.trim();
11522
11523 if let Some(tool_calls) = self.parse_main_tool_calls(content, &protocol)? {
11524 match self
11525 .handle_tool_calls(
11526 processed_input,
11527 content,
11528 tool_calls,
11529 &mut all_tool_calls,
11530 None,
11531 )
11532 .await?
11533 {
11534 ToolCallOutcome::Continue | ToolCallOutcome::TransitionFired => continue,
11535 ToolCallOutcome::Rejected(resp) => {
11536 self.finish_turn_if_root(&resp).await?;
11537 return Ok(resp);
11538 }
11539 }
11540 }
11541
11542 let (extracted_thinking, answer) = self.extract_thinking(content);
11543 if extracted_thinking.is_some() {
11544 thinking_content = extracted_thinking;
11545 }
11546
11547 let output_data = self.process_output(&answer, &input_data.context).await?;
11548
11549 let mut final_content = if output_data.metadata.rejected {
11550 output_data
11551 .metadata
11552 .rejection_reason
11553 .unwrap_or_else(|| answer.to_string())
11554 } else {
11555 output_data.content
11556 };
11557
11558 let reflection_metadata;
11560 (final_content, reflection_metadata) = self
11561 .run_reflection(&*llm, processed_input, final_content)
11562 .await?;
11563
11564 final_content =
11565 self.format_response_with_thinking(thinking_content.as_deref(), &final_content);
11566
11567 let final_content = {
11571 let result = self
11572 .post_loop_processing(processed_input, final_content)
11573 .await?;
11574 self.apply_post_loop_result(processed_input, result)
11575 .await?
11576 .content
11577 };
11578
11579 let reflected = reflection_metadata.is_some();
11580 let reasoning_mode_debug = format!("{:?}", reasoning_mode);
11581
11582 let response = self.build_agent_response(AgentResponseParts {
11583 content: final_content,
11584 all_tool_calls,
11585 reasoning_mode,
11586 auto_detected,
11587 iterations,
11588 thinking: thinking_content,
11589 reflection_metadata,
11590 });
11591
11592 self.finish_turn_if_root(&response).await?;
11593
11594 let tool_call_count = response.tool_calls.as_ref().map(|tc| tc.len()).unwrap_or(0);
11595 info!(
11596 tool_calls = tool_call_count,
11597 response_len = response.content.len(),
11598 reasoning_mode = %reasoning_mode_debug,
11599 reflected = reflected,
11600 "Chat completed"
11601 );
11602 return Ok(response);
11603 }
11604 }
11605
11606 async fn generate_buffered_streaming_draft(
11607 &self,
11608 processed_input: &str,
11609 routing_resolved: Arc<AtomicBool>,
11610 ) -> Result<StreamingDraftResult> {
11611 let llm = self.get_state_llm()?;
11612 if llm.configured_tool_choice().is_some() {
11613 let draft = self
11614 .generate_main_response_draft(processed_input, &ReasoningMode::None)
11615 .await?;
11616 return Ok(StreamingDraftResult::new(draft, Vec::new()));
11617 }
11618 let protocol = self.main_tool_protocol(llm.as_ref(), true).await?;
11620 let messages = self.build_messages_for_draft(processed_input).await?;
11621 let source = self
11622 .open_main_stream_with_recovery(Arc::clone(&llm), &messages, &protocol)
11623 .await?;
11624 let mut buffer = crate::optimization::StreamBranchBuffer::new(self.streaming.buffer_size)?;
11625 let mut chunks = Vec::new();
11626 let mut accumulated = String::new();
11627 match source {
11628 MainStreamSource::StaticResponse(text) => {
11629 accumulated.push_str(&text);
11630 let stream_chunk = StreamChunk::content(text);
11631 if routing_resolved.load(Ordering::SeqCst) {
11632 chunks.push(stream_chunk);
11633 } else {
11634 buffer.push(stream_chunk)?;
11635 }
11636 }
11637 MainStreamSource::Stream(mut stream) => {
11638 while let Some(chunk_result) = stream.next().await {
11639 let chunk = chunk_result.map_err(|e| AgentError::LLM(e.to_string()))?;
11640 accumulated.push_str(&chunk.delta);
11641 let stream_chunk = StreamChunk::content(chunk.delta);
11642 if routing_resolved.load(Ordering::SeqCst) {
11643 chunks.push(stream_chunk);
11644 } else {
11645 buffer.push(stream_chunk)?;
11646 }
11647 }
11648 }
11649 }
11650 chunks.splice(0..0, buffer.drain());
11651 let content = accumulated.trim().to_string();
11652 let draft = if let Some(calls) = self.parse_tool_calls(&content)? {
11653 MainResponseDraft::ToolCalls {
11654 raw_content: content,
11655 calls,
11656 thinking: None,
11657 }
11658 } else {
11659 MainResponseDraft::Text {
11660 raw_content: content,
11661 thinking: None,
11662 }
11663 };
11664 Ok(StreamingDraftResult::new(draft, chunks))
11665 }
11666
11667 async fn try_buffered_streaming_branches(
11668 &self,
11669 processed_input: &str,
11670 input_context: &HashMap<String, Value>,
11671 ) -> Result<Option<(AgentResponse, Vec<StreamChunk>)>> {
11672 let optimization = &self.runtime_config.optimization;
11673 if !optimization.enabled {
11674 return Ok(None);
11675 }
11676 if !matches!(
11682 self.get_effective_reasoning_config().mode,
11683 ReasoningMode::None
11684 ) {
11685 return Ok(None);
11686 }
11687 let transition_enabled =
11688 optimization.speculative_state_transitions && self.has_parallel_transition_candidates();
11689 if !transition_enabled {
11690 return Ok(None);
11691 }
11692 let mut branch_scheduler =
11693 TurnBranchScheduler::new(optimization.max_parallel_runtime_tasks)?;
11694 if !branch_scheduler.reserve_task() {
11695 return Ok(None);
11696 }
11697 if !self
11698 .reserve_active_speculative_llm_call(RuntimeOptimizationKind::BufferedStreamingRouting)
11699 {
11700 branch_scheduler.release_task();
11701 return Ok(None);
11702 }
11703 if !branch_scheduler.reserve_task() {
11704 branch_scheduler.release_task();
11705 return Ok(None);
11706 }
11707 let mut main_branch = RuntimeBranch::new(
11708 RuntimeTaskPurpose::MainResponse,
11709 RuntimeOptimizationKind::BufferedStreamingRouting,
11710 RuntimeTaskPriority::Normal,
11711 RuntimeCommitBehavior::FinalResponse,
11712 );
11713 let mut transition_branch = RuntimeBranch::new(
11714 RuntimeTaskPurpose::StateTransition,
11715 RuntimeOptimizationKind::ParallelStateTransition,
11716 RuntimeTaskPriority::Critical,
11717 RuntimeCommitBehavior::TransitionDecision,
11718 );
11719 let main_id = main_branch.branch_id();
11720 let transition_id = transition_branch.branch_id();
11721 let routing_resolved = Arc::new(AtomicBool::new(false));
11722 let mut main_future =
11723 Box::pin(crate::optimization::observability::with_branch_observation(
11724 &main_id,
11725 RuntimeOptimizationKind::BufferedStreamingRouting,
11726 RuntimeCommitBehavior::FinalResponse,
11727 self.generate_buffered_streaming_draft(
11728 processed_input,
11729 Arc::clone(&routing_resolved),
11730 ),
11731 ));
11732 let mut transition_future =
11733 Box::pin(crate::optimization::observability::with_branch_observation(
11734 &transition_id,
11735 RuntimeOptimizationKind::ParallelStateTransition,
11736 RuntimeCommitBehavior::TransitionDecision,
11737 self.select_parallel_transition_candidate(processed_input),
11738 ));
11739 let mut main_pending = true;
11740 let mut transition_pending = true;
11741 let mut main_result: Option<Result<StreamingDraftResult>> = None;
11742 let mut transition_finalized = false;
11743 let mut transition_candidate: Option<TransitionCandidate> = None;
11744 loop {
11745 if let Some(candidate) = transition_candidate.take() {
11746 if self
11747 .approve_transition_target(&candidate.from_state, candidate.target())
11748 .await?
11749 {
11750 drop(main_future);
11752 drop(transition_future);
11753 self.finalize_branch_loss(
11754 &main_id,
11755 RuntimeOptimizationKind::BufferedStreamingRouting,
11756 RuntimeCommitBehavior::FinalResponse,
11757 main_pending,
11758 main_result.as_ref().map(|result| result.is_err()),
11759 );
11760 if !self
11761 .apply_pre_response_transition_candidate(
11762 &candidate,
11763 &HashMap::new(),
11764 processed_input,
11765 )
11766 .await?
11767 {
11768 self.finalize_optional_branch(
11769 &transition_id,
11770 RuntimeOptimizationKind::ParallelStateTransition,
11771 RuntimeCommitBehavior::TransitionDecision,
11772 "discarded",
11773 false,
11774 );
11775 return Ok(None);
11776 }
11777 self.finalize_optional_branch(
11778 &transition_id,
11779 RuntimeOptimizationKind::ParallelStateTransition,
11780 RuntimeCommitBehavior::TransitionDecision,
11781 "committed",
11782 true,
11783 );
11784 let response = self.redispatch_current_state(processed_input).await?;
11785 return Ok(Some((
11786 response.clone(),
11787 vec![StreamChunk::content(response.content)],
11788 )));
11789 }
11790 self.finalize_optional_branch(
11791 &transition_id,
11792 RuntimeOptimizationKind::ParallelStateTransition,
11793 RuntimeCommitBehavior::TransitionDecision,
11794 "discarded",
11795 false,
11796 );
11797 transition_finalized = true;
11798 }
11799 if transition_finalized && !routing_resolved.load(Ordering::SeqCst) {
11805 match self
11806 .resolve_buffered_skill_after_transition(processed_input, &routing_resolved)
11807 .await
11808 {
11809 Ok(Some(candidate)) => {
11810 drop(main_future);
11812 drop(transition_future);
11813 self.finalize_branch_loss(
11814 &main_id,
11815 RuntimeOptimizationKind::BufferedStreamingRouting,
11816 RuntimeCommitBehavior::FinalResponse,
11817 main_pending,
11818 main_result.as_ref().map(|result| result.is_err()),
11819 );
11820 return match self
11821 .commit_winning_skill_candidate(
11822 candidate,
11823 processed_input,
11824 input_context,
11825 )
11826 .await?
11827 {
11828 Some(response) => Ok(Some((
11829 response.clone(),
11830 vec![StreamChunk::content(response.content)],
11831 ))),
11832 None => Ok(None),
11833 };
11834 }
11835 Ok(None) => {}
11836 Err(error) => {
11837 drop(main_future);
11838 drop(transition_future);
11839 self.finalize_branch_loss(
11840 &main_id,
11841 RuntimeOptimizationKind::BufferedStreamingRouting,
11842 RuntimeCommitBehavior::FinalResponse,
11843 main_pending,
11844 main_result.as_ref().map(|result| result.is_err()),
11845 );
11846 return Err(error);
11847 }
11848 }
11849 }
11850 if transition_finalized
11851 && routing_resolved.load(Ordering::SeqCst)
11852 && let Some(result) = main_result.take()
11853 {
11854 let stream_draft = match result {
11855 Ok(stream_draft) => stream_draft,
11856 Err(error) => {
11857 self.finalize_optional_branch(
11858 &main_id,
11859 RuntimeOptimizationKind::BufferedStreamingRouting,
11860 RuntimeCommitBehavior::FinalResponse,
11861 "failed",
11862 false,
11863 );
11864 return Err(error);
11865 }
11866 };
11867 let raw_draft_content = stream_draft.draft.raw_content().to_string();
11868 let buffered_chunks = stream_draft.chunks;
11869 self.finalize_optional_branch(
11870 &main_id,
11871 RuntimeOptimizationKind::BufferedStreamingRouting,
11872 RuntimeCommitBehavior::FinalResponse,
11873 "committed",
11874 true,
11875 );
11876 let response = self
11877 .commit_main_response_draft(
11878 processed_input,
11879 input_context,
11880 stream_draft.draft,
11881 ReasoningMode::None,
11882 false,
11883 )
11884 .await?;
11885 let chunks = if response.content == raw_draft_content {
11886 buffered_chunks
11887 } else {
11888 vec![StreamChunk::content(response.content.clone())]
11889 };
11890 return Ok(Some((response, chunks)));
11891 }
11892 tokio::select! {
11893 result = &mut main_future, if main_pending => {
11894 main_pending = false;
11895 main_branch.transition_to(RuntimeBranchStatus::Completed)?;
11896 main_result = Some(result);
11897 }
11898 result = &mut transition_future, if transition_pending => {
11899 transition_pending = false;
11900 transition_branch.transition_to(RuntimeBranchStatus::Completed)?;
11901 match result {
11902 Ok(ParallelTransitionSelection::Candidate(candidate)) => {
11903 transition_candidate = Some(candidate)
11904 }
11905 Ok(ParallelTransitionSelection::NoMatch) => {
11906 self.finalize_optional_branch(
11907 &transition_id,
11908 RuntimeOptimizationKind::ParallelStateTransition,
11909 RuntimeCommitBehavior::TransitionDecision,
11910 "discarded",
11911 false,
11912 );
11913 transition_finalized = true;
11914 }
11915 Ok(ParallelTransitionSelection::ReservationExhausted) => {
11916 self.finalize_optional_branch(
11917 &transition_id,
11918 RuntimeOptimizationKind::ParallelStateTransition,
11919 RuntimeCommitBehavior::TransitionDecision,
11920 "cancelled",
11921 false,
11922 );
11923 routing_resolved.store(true, Ordering::SeqCst);
11924 self.finalize_branch_loss(
11925 &main_id,
11926 RuntimeOptimizationKind::BufferedStreamingRouting,
11927 RuntimeCommitBehavior::FinalResponse,
11928 main_pending,
11929 main_result.as_ref().map(|result| result.is_err()),
11930 );
11931 return Ok(None);
11932 }
11933 Err(_) => {
11934 self.finalize_optional_branch(
11935 &transition_id,
11936 RuntimeOptimizationKind::ParallelStateTransition,
11937 RuntimeCommitBehavior::TransitionDecision,
11938 "failed",
11939 false,
11940 );
11941 transition_finalized = true;
11942 }
11943 }
11944 }
11945 }
11946 }
11947 }
11948
11949 async fn resolve_buffered_skill_after_transition(
11955 &self,
11956 processed_input: &str,
11957 routing_resolved: &AtomicBool,
11958 ) -> Result<Option<SkillCandidate>> {
11959 let candidate = if self.skill_router.is_some() {
11960 self.select_skill_candidate(processed_input).await?
11961 } else {
11962 None
11963 };
11964 if candidate.is_none() {
11965 routing_resolved.store(true, Ordering::SeqCst);
11966 }
11967 Ok(candidate)
11968 }
11969
11970 fn run_loop_internal_stream<'a>(
11974 &'a self,
11975 input: &'a str,
11976 terminal: RuntimeStreamTerminalSlot,
11977 ) -> Pin<Box<dyn Stream<Item = StreamChunk> + Send + 'a>> {
11978 let include_state_events = self.streaming.include_state_events;
11979
11980 Box::pin(async_stream::stream! {
11981 self.begin_root_turn();
11982 self.pre_turn_session_lifecycle().await;
11984
11985 let input_data = match self.process_input(input).await {
11986 Ok(data) => data,
11987 Err(e) => {
11988 yield StreamChunk::error(e.to_string());
11989 return;
11990 }
11991 };
11992 self.update_active_turn_context(&input_data.content, input_data.context.clone());
11993
11994 for (key, value) in &input_data.context {
11996 let _ = self.context_manager.set(key, value.clone());
11997 }
11998
11999 if input_data.metadata.rejected {
12000 let reason = input_data
12001 .metadata
12002 .rejection_reason
12003 .unwrap_or_else(|| "Input rejected".to_string());
12004 warn!(reason = %reason, "Input rejected (stream)");
12005 let response = AgentResponse::new(&reason);
12008 if let Err(e) = self.finish_turn_if_root(&response).await {
12009 yield StreamChunk::error(e.to_string());
12010 return;
12011 }
12012 yield StreamChunk::content(&reason);
12013 record_runtime_stream_final(&terminal, response);
12014 yield StreamChunk::Done {};
12015 return;
12016 }
12017
12018 let processed_input = &input_data.content;
12019
12020 let streaming_policy = self.runtime_config.optimization.streaming_policy;
12021
12022 if self.runtime_config.optimization.enabled
12029 && !matches!(
12030 streaming_policy,
12031 crate::optimization::StreamingOptimizationPolicy::Disabled
12032 )
12033 {
12034 match self.try_pre_response_transition(processed_input).await {
12035 Ok(Some(response)) => {
12036 yield StreamChunk::content(&response.content);
12037 record_runtime_stream_final(&terminal, response);
12038 yield StreamChunk::Done {};
12039 return;
12040 }
12041 Ok(None) => {}
12042 Err(e) => {
12043 yield StreamChunk::error(e.to_string());
12044 return;
12045 }
12046 }
12047 }
12048
12049 if self.runtime_config.optimization.enabled
12050 && matches!(
12051 streaming_policy,
12052 crate::optimization::StreamingOptimizationPolicy::BufferUntilRoutingDone
12053 )
12054 {
12055 match Box::pin(self.try_buffered_streaming_branches(processed_input, &input_data.context)).await {
12060 Ok(Some((response, chunks))) => {
12061 for chunk in chunks {
12062 yield chunk;
12063 }
12064 record_runtime_stream_final(&terminal, response);
12065 yield StreamChunk::Done {};
12066 return;
12067 }
12068 Ok(None) => {}
12069 Err(e) => {
12070 yield StreamChunk::error(e.to_string());
12071 return;
12072 }
12073 }
12074 }
12075
12076 if let Some(ref sm) = self.state_machine
12078 && let Some(def) = sm.current_definition()
12079 {
12080 let orchestration_result = if let Some(ref delegate_id) = def.delegate {
12081 Some(self.handle_delegated_state(processed_input, delegate_id, &def).await)
12082 } else if let Some(ref concurrent_config) = def.concurrent {
12083 Some(self.handle_concurrent_state(processed_input, concurrent_config).await)
12084 } else if let Some(ref group_chat_config) = def.group_chat {
12085 Some(self.handle_group_chat_state(processed_input, group_chat_config).await)
12086 } else if let Some(ref pipeline_config) = def.pipeline {
12087 Some(self.handle_pipeline_state(processed_input, pipeline_config).await)
12088 } else if let Some(ref handoff_config) = def.handoff {
12089 Some(self.handle_handoff_state(processed_input, handoff_config).await)
12090 } else {
12091 None
12092 };
12093
12094 if let Some(result) = orchestration_result {
12095 match result {
12096 Ok(response) => {
12097 yield StreamChunk::content(&response.content);
12098 record_runtime_stream_final(&terminal, response);
12099 yield StreamChunk::Done {};
12100 }
12101 Err(e) => {
12102 yield StreamChunk::error(e.to_string());
12103 }
12104 }
12105 return;
12106 }
12107 }
12108
12109 match self.try_skill_route(processed_input).await {
12111 Ok(SkillRouteResult::Response { skill_id, content }) => {
12112 if let Err(e) = self.commit_root_user_message(processed_input).await {
12113 yield StreamChunk::error(e.to_string());
12114 return;
12115 }
12116 match self.handle_skill_response(processed_input, &skill_id, content, &input_data.context).await {
12117 Ok(resp) => {
12118 yield StreamChunk::content(&resp.content);
12119 record_runtime_stream_final(&terminal, resp);
12120 yield StreamChunk::Done {};
12121 return;
12122 }
12123 Err(e) => {
12124 yield StreamChunk::error(e.to_string());
12125 return;
12126 }
12127 }
12128 }
12129 Ok(SkillRouteResult::NeedsClarification {
12130 response,
12131 ownership,
12132 }) => {
12133 let admission = match self
12134 .admit_optional_disambiguation_ownership(ownership)
12135 .await
12136 {
12137 Ok(admission) => admission,
12138 Err(e) => {
12139 yield StreamChunk::error(e.to_string());
12140 return;
12141 }
12142 };
12143 if let Err(e) = self.commit_root_user_message(processed_input).await {
12144 yield StreamChunk::error(e.to_string());
12145 return;
12146 }
12147 if Self::skill_clarification_needs_memory_record(&response)
12149 && let Err(e) = self.memory.add_message(ChatMessage::assistant(&response.content)).await
12150 {
12151 yield StreamChunk::error(e.to_string());
12152 return;
12153 }
12154 drop(admission);
12155 if let Err(e) = self.finish_turn_if_root(&response).await {
12156 yield StreamChunk::error(e.to_string());
12157 return;
12158 }
12159 yield StreamChunk::content(&response.content);
12160 record_runtime_stream_final(&terminal, response);
12161 yield StreamChunk::Done {};
12162 return;
12163 }
12164 Ok(SkillRouteResult::NoMatch) => {} Err(e) => {
12166 yield StreamChunk::error(e.to_string());
12167 return;
12168 }
12169 }
12170
12171 let effective_reasoning = self.get_effective_reasoning_config();
12173 let reasoning_mode = match self.determine_reasoning_mode(processed_input).await {
12174 Ok(mode) => mode,
12175 Err(e) => {
12176 yield StreamChunk::error(e.to_string());
12177 return;
12178 }
12179 };
12180 let auto_detected = matches!(effective_reasoning.mode, ReasoningMode::Auto);
12181
12182 info!(
12183 reasoning_mode = ?reasoning_mode,
12184 auto_detected = auto_detected,
12185 "Reasoning mode determined (stream)"
12186 );
12187
12188 if matches!(reasoning_mode, ReasoningMode::PlanAndExecute) {
12190 if let Err(e) = self.commit_root_user_message(processed_input).await {
12191 yield StreamChunk::error(e.to_string());
12192 return;
12193 }
12194 match self.handle_plan_and_execute(processed_input, &input_data.context, auto_detected).await {
12195 Ok(resp) => {
12196 yield StreamChunk::content(&resp.content);
12197 record_runtime_stream_final(&terminal, resp);
12198 yield StreamChunk::Done {};
12199 return;
12200 }
12201 Err(e) => {
12202 yield StreamChunk::error(e.to_string());
12203 return;
12204 }
12205 }
12206 }
12207
12208 if let Err(e) = self.commit_root_user_message(processed_input).await {
12209 yield StreamChunk::error(e.to_string());
12210 return;
12211 }
12212
12213 let llm = match self.get_state_llm() {
12214 Ok(llm) => llm,
12215 Err(e) => {
12216 yield StreamChunk::error(e.to_string());
12217 return;
12218 }
12219 };
12220
12221 let mut iterations = 0u32;
12222 let mut all_tool_calls: Vec<ToolCall> = Vec::new();
12223 let mut thinking_content: Option<String> = None;
12224
12225 loop {
12226 let effective_max = if reasoning_mode != ReasoningMode::None {
12228 let rc = self.get_effective_reasoning_config();
12229 self.max_iterations.min(rc.max_iterations)
12230 } else {
12231 self.max_iterations
12232 };
12233
12234 if iterations >= effective_max {
12235 let err_msg = format!("Max iterations ({}) exceeded", effective_max);
12236 let err = AgentError::Other(err_msg.clone());
12237 self.hooks.on_error(&err).await;
12238 error!(iterations = iterations, "Max iterations exceeded (stream)");
12239 yield StreamChunk::error(err_msg);
12240 return;
12241 }
12242 iterations += 1;
12243 *self.iteration_count.write() = iterations;
12244
12245 debug!(iteration = iterations, max = effective_max, "LLM call (stream)");
12246
12247 let protocol = match self.main_tool_protocol(llm.as_ref(), false).await {
12248 Ok(protocol) => protocol,
12249 Err(e) => {
12250 yield StreamChunk::error(e.to_string());
12251 return;
12252 }
12253 };
12254 let mut messages = match self
12255 .build_messages_internal(true, None, protocol.choice.is_none())
12256 .await
12257 {
12258 Ok(m) => m,
12259 Err(e) => {
12260 yield StreamChunk::error(e.to_string());
12261 return;
12262 }
12263 };
12264 self.inject_reasoning_prompt(&mut messages, &reasoning_mode, iterations == 1);
12265
12266 self.hooks.on_llm_start(&messages).await;
12267 let llm_start = Instant::now();
12268
12269 let buffered_decision = self.main_stream_must_buffer(&reasoning_mode, &protocol);
12270 let content = if buffered_decision {
12271 let response = match self
12275 .complete_main_llm_with_recovery(
12276 Arc::clone(&llm),
12277 &messages,
12278 &protocol,
12279 )
12280 .await
12281 {
12282 Ok(r) => r,
12283 Err(e) => {
12284 yield StreamChunk::error(e.to_string());
12285 return;
12286 }
12287 };
12288 let llm_duration_ms = llm_start.elapsed().as_millis() as u64;
12289 self.hooks.on_llm_complete(&response, llm_duration_ms).await;
12290 response.content.trim().to_string()
12291 } else {
12292 let source = match self
12294 .open_main_stream_with_recovery(Arc::clone(&llm), &messages, &protocol)
12295 .await
12296 {
12297 Ok(source) => source,
12298 Err(e) => {
12299 yield StreamChunk::error(e.to_string());
12300 return;
12301 }
12302 };
12303 let mut accumulated = String::new();
12304 match source {
12305 MainStreamSource::StaticResponse(text) => {
12306 accumulated.push_str(&text);
12307 yield StreamChunk::content(text);
12308 }
12309 MainStreamSource::Stream(mut stream_inner) => {
12310 while let Some(chunk_result) = stream_inner.next().await {
12311 match chunk_result {
12312 Ok(chunk) => {
12313 accumulated.push_str(&chunk.delta);
12314 yield StreamChunk::content(chunk.delta);
12315 }
12316 Err(e) => {
12317 yield StreamChunk::error(e.to_string());
12319 return;
12320 }
12321 }
12322 }
12323 }
12324 }
12325 let llm_duration_ms = llm_start.elapsed().as_millis() as u64;
12326 let llm_response = ai_agents_core::LLMResponse::new(
12328 accumulated.trim(),
12329 ai_agents_core::FinishReason::Stop,
12330 );
12331 self.hooks.on_llm_complete(&llm_response, llm_duration_ms).await;
12332 accumulated.trim().to_string()
12333 };
12334
12335 let parsed_tool_calls = match self.parse_main_tool_calls(&content, &protocol) {
12337 Ok(calls) => calls,
12338 Err(error) => {
12339 yield StreamChunk::error(error.to_string());
12340 return;
12341 }
12342 };
12343 if let Some(tool_calls) = parsed_tool_calls {
12344 let mut events = Vec::new();
12347 let outcome = self
12348 .handle_tool_calls(
12349 processed_input,
12350 &content,
12351 tool_calls,
12352 &mut all_tool_calls,
12353 Some(&mut events),
12354 )
12355 .await;
12356 for chunk in events.drain(..) {
12357 yield chunk;
12358 }
12359 match outcome {
12360 Ok(ToolCallOutcome::Continue) | Ok(ToolCallOutcome::TransitionFired) => continue,
12361 Ok(ToolCallOutcome::Rejected(response)) => {
12362 if let Err(finalize_error) = self.finish_turn_if_root(&response).await {
12363 yield StreamChunk::error(finalize_error.to_string());
12364 return;
12365 }
12366 let legacy_error = response.content.clone();
12367 record_runtime_stream_final(&terminal, response);
12368 yield StreamChunk::error(legacy_error);
12369 yield StreamChunk::Done {};
12370 return;
12371 }
12372 Err(e) => {
12373 yield StreamChunk::error(e.to_string());
12374 return;
12375 }
12376 }
12377 }
12378
12379 let (extracted_thinking, answer) = self.extract_thinking(&content);
12381 if extracted_thinking.is_some() {
12382 thinking_content = extracted_thinking;
12383 }
12384
12385 let output_data = match self.process_output(&answer, &input_data.context).await {
12386 Ok(d) => d,
12387 Err(e) => {
12388 yield StreamChunk::error(e.to_string());
12389 return;
12390 }
12391 };
12392
12393 let final_content = if output_data.metadata.rejected {
12394 output_data
12395 .metadata
12396 .rejection_reason
12397 .unwrap_or_else(|| answer.to_string())
12398 } else {
12399 output_data.content
12400 };
12401
12402 let (final_content, reflection_metadata) = match self
12404 .run_reflection(&*llm, processed_input, final_content)
12405 .await
12406 {
12407 Ok(r) => r,
12408 Err(e) => {
12409 yield StreamChunk::error(e.to_string());
12410 return;
12411 }
12412 };
12413
12414 let final_content = self.format_response_with_thinking(
12415 thinking_content.as_deref(),
12416 &final_content,
12417 );
12418
12419 if buffered_decision {
12421 yield StreamChunk::content(&final_content);
12422 }
12423
12424 let post_result = match self
12428 .post_loop_processing(processed_input, final_content)
12429 .await
12430 {
12431 Ok(r) => r,
12432 Err(e) => {
12433 yield StreamChunk::error(e.to_string());
12434 return;
12435 }
12436 };
12437
12438 let applied = match self.apply_post_loop_result(processed_input, post_result).await {
12439 Ok(applied) => applied,
12440 Err(e) => {
12441 yield StreamChunk::error(e.to_string());
12442 return;
12443 }
12444 };
12445
12446 if applied.transitioned {
12447 if include_state_events
12448 && let Some(state) = self.current_state()
12449 {
12450 yield StreamChunk::state_transition(None, state);
12451 }
12452 if applied.regenerated {
12458 yield StreamChunk::content(&applied.content);
12459 }
12460 }
12461 let final_content = applied.content;
12462
12463 let final_response = self.build_agent_response(AgentResponseParts {
12465 content: final_content,
12466 all_tool_calls,
12467 reasoning_mode,
12468 auto_detected,
12469 iterations,
12470 thinking: thinking_content,
12471 reflection_metadata,
12472 });
12473 if let Err(e) = self.finish_turn_if_root(&final_response).await {
12474 yield StreamChunk::error(e.to_string());
12475 return;
12476 }
12477
12478 record_runtime_stream_final(&terminal, final_response);
12479 yield StreamChunk::Done {};
12480 return;
12481 }
12482 })
12483 }
12484
12485 fn run_loop_stream<'a>(
12488 &'a self,
12489 input: &'a str,
12490 terminal: RuntimeStreamTerminalSlot,
12491 ) -> Pin<Box<dyn Stream<Item = StreamChunk> + Send + 'a>> {
12492 Box::pin(async_stream::stream! {
12493 self.begin_root_turn();
12494 let _root_cleanup = RootTurnCleanup::new(self);
12495 self.hooks.on_message_received(input).await;
12496
12497 if let Err(e) = self.prepare_turn_context().await {
12499 yield StreamChunk::error(e.to_string());
12500 return;
12501 }
12502
12503 self.clear_disambiguation_context();
12505
12506 let input_to_run = match self.resolve_disambiguation(input).await {
12510 Err(e) => {
12511 yield StreamChunk::error(e.to_string());
12512 return;
12513 }
12514 Ok(DisambiguationDispatch::Terminal(response)) => {
12515 yield StreamChunk::content(&response.content);
12516 record_runtime_stream_final(&terminal, response);
12517 yield StreamChunk::Done {};
12518 return;
12519 }
12520 Ok(DisambiguationDispatch::RecheckSkill {
12521 skill_id,
12522 enriched_input,
12523 disambiguation_epoch,
12524 state_generation,
12525 }) => {
12526 match self
12527 .recheck_skill_disambiguation(
12528 &skill_id,
12529 &enriched_input,
12530 disambiguation_epoch,
12531 state_generation,
12532 )
12533 .await
12534 {
12535 Ok(resp) => {
12536 yield StreamChunk::content(&resp.content);
12537 record_runtime_stream_final(&terminal, resp);
12538 yield StreamChunk::Done {};
12539 return;
12540 }
12541 Err(e) => {
12542 yield StreamChunk::error(e.to_string());
12543 return;
12544 }
12545 }
12546 }
12547 Ok(DisambiguationDispatch::Proceed(input)) => input,
12548 };
12549
12550 let mut inner = self.run_loop_internal_stream(&input_to_run, Arc::clone(&terminal));
12551 while let Some(chunk) = inner.next().await {
12552 yield chunk;
12553 }
12554 })
12555 }
12556
12557 pub fn info(&self) -> AgentInfo {
12558 self.info.clone()
12559 }
12560
12561 pub fn skills(&self) -> &[SkillDefinition] {
12562 &self.skills
12563 }
12564
12565 async fn reset_runtime_state(&self) -> Result<()> {
12567 let _admission = self.disambiguation_admission.write().await;
12568 if self.state_transition_reserved.load(Ordering::SeqCst) {
12569 return Err(AgentError::Other(
12570 "Cannot reset while a state transition is in progress".to_string(),
12571 ));
12572 }
12573 self.disambiguation_epoch.fetch_add(1, Ordering::SeqCst);
12574 *self.pending_skill_id.write() = None;
12575 if let Some(disambiguator) = self.disambiguation_manager.as_ref() {
12576 disambiguator.clear_pending().await;
12577 }
12578 self.memory.clear().await?;
12579 self.active_native_exchanges.write().clear();
12580 *self.iteration_count.write() = 0;
12581 self.tool_call_history.write().clear();
12582 if let Some(ref sm) = self.state_machine {
12583 sm.reset();
12584 }
12585 Ok(())
12586 }
12587
12588 pub async fn reset(&self) -> Result<()> {
12590 self.reset_runtime_state().await
12591 }
12592
12593 pub fn max_context_tokens(&self) -> u32 {
12594 self.max_context_tokens
12595 }
12596
12597 pub fn llm_registry(&self) -> &Arc<LLMRegistry> {
12598 &self.llm_registry
12599 }
12600
12601 pub fn state_machine(&self) -> Option<&Arc<StateMachine>> {
12602 self.state_machine.as_ref()
12603 }
12604
12605 pub fn context_manager(&self) -> &Arc<ContextManager> {
12606 &self.context_manager
12607 }
12608
12609 pub fn tool_call_history(&self) -> Vec<ToolCallRecord> {
12610 self.tool_call_history.read().clone()
12611 }
12612
12613 pub fn memory_token_budget(&self) -> Option<&MemoryTokenBudget> {
12614 self.memory_token_budget.as_ref()
12615 }
12616
12617 pub fn parallel_tools_config(&self) -> &ParallelToolsConfig {
12618 &self.parallel_tools
12619 }
12620
12621 pub fn streaming_config(&self) -> &StreamingConfig {
12622 &self.streaming
12623 }
12624
12625 pub fn hooks(&self) -> &Arc<dyn AgentHooks> {
12626 &self.hooks
12627 }
12628
12629 pub fn hitl_engine(&self) -> Option<&HITLEngine> {
12630 self.hitl_engine.as_ref()
12631 }
12632
12633 pub fn approval_handler(&self) -> &Arc<dyn ApprovalHandler> {
12634 &self.approval_handler
12635 }
12636
12637 fn build_hitl_language_context(&self) -> HashMap<String, Value> {
12639 let mut ctx = HashMap::new();
12640 for key in &["user.language", "input.detected.language", "language"] {
12641 if let Some(val) = self.context_manager.get(key) {
12642 ctx.insert(key.to_string(), val);
12643 }
12644 }
12645 ctx
12646 }
12647
12648 async fn request_hitl_approval(&self, check_result: HITLCheckResult) -> Result<ApprovalResult> {
12650 let Some(request) = check_result.into_request() else {
12651 return Ok(ApprovalResult::Approved);
12652 };
12653
12654 self.hooks.on_approval_requested(&request).await;
12655
12656 let timeout = request.timeout;
12657
12658 let raw_result = if let Some(duration) = timeout {
12659 match tokio::time::timeout(
12660 duration,
12661 self.approval_handler.request_approval(request.clone()),
12662 )
12663 .await
12664 {
12665 Ok(result) => result,
12666 Err(_) => ApprovalResult::timeout(),
12667 }
12668 } else {
12669 self.approval_handler
12670 .request_approval(request.clone())
12671 .await
12672 };
12673
12674 self.hooks
12675 .on_approval_result(&request.id, &raw_result)
12676 .await;
12677
12678 let (outcome, effective_result): (ApprovalResolvedOutcome, Result<ApprovalResult>) =
12679 match &raw_result {
12680 ApprovalResult::Approved => (
12681 ApprovalResolvedOutcome::Approved,
12682 Ok(ApprovalResult::Approved),
12683 ),
12684 ApprovalResult::Rejected { reason } => (
12685 ApprovalResolvedOutcome::Rejected {
12686 reason: reason.clone(),
12687 },
12688 Ok(ApprovalResult::Rejected {
12689 reason: reason.clone(),
12690 }),
12691 ),
12692 ApprovalResult::Modified { changes } => (
12693 ApprovalResolvedOutcome::Modified {
12694 changes: changes.clone(),
12695 },
12696 Ok(ApprovalResult::Modified {
12697 changes: changes.clone(),
12698 }),
12699 ),
12700 ApprovalResult::Timeout => {
12701 if let Some(ref engine) = self.hitl_engine {
12702 match engine.config().on_timeout {
12703 TimeoutAction::Approve => (
12704 ApprovalResolvedOutcome::Approved,
12705 Ok(ApprovalResult::Approved),
12706 ),
12707 TimeoutAction::Reject => {
12708 let reason = Some("Timeout".to_string());
12709 (
12710 ApprovalResolvedOutcome::Rejected {
12711 reason: reason.clone(),
12712 },
12713 Ok(ApprovalResult::Rejected { reason }),
12714 )
12715 }
12716 TimeoutAction::Error => {
12717 let message = "HITL approval timeout".to_string();
12718 (
12719 ApprovalResolvedOutcome::Error {
12720 message: message.clone(),
12721 },
12722 Err(AgentError::Other(message)),
12723 )
12724 }
12725 }
12726 } else {
12727 let reason = Some("Timeout (no engine)".to_string());
12728 (
12729 ApprovalResolvedOutcome::Rejected {
12730 reason: reason.clone(),
12731 },
12732 Ok(ApprovalResult::Rejected { reason }),
12733 )
12734 }
12735 }
12736 };
12737
12738 self.hooks
12739 .on_approval_resolved(&request, &raw_result, &outcome)
12740 .await;
12741
12742 effective_result
12743 }
12744
12745 pub async fn check_state_hitl(&self, from: Option<&str>, to: &str) -> Result<bool> {
12746 if let Some(ref hitl_engine) = self.hitl_engine {
12747 let hitl_lang_ctx = self.build_hitl_language_context();
12748 let check_result = self
12749 .observe_purpose(
12750 ObservationPurpose::HitlLocalization,
12751 hitl_engine.check_state_transition_with_localization(
12752 from,
12753 to,
12754 &hitl_lang_ctx,
12755 self.approval_handler.as_ref(),
12756 Some(&self.llm_registry),
12757 ),
12758 )
12759 .await?;
12760 if check_result.is_required() {
12761 let result = self.request_hitl_approval(check_result).await?;
12762 return Ok(matches!(
12763 result,
12764 ApprovalResult::Approved | ApprovalResult::Modified { .. }
12765 ));
12766 }
12767 }
12768 Ok(true)
12769 }
12770
12771 async fn execute_tools_parallel(
12773 &self,
12774 tool_calls: &[ToolCall],
12775 ) -> Vec<(String, Result<String>)> {
12776 let can_run_parallel = tool_calls.iter().all(|tc| {
12777 self.tools
12778 .resolve(&tc.name)
12779 .map(|resolved| resolved.tool.classify_call(&tc.arguments).concurrency_safe)
12780 .unwrap_or(false)
12781 });
12782
12783 if !self.parallel_tools.enabled || tool_calls.len() <= 1 || !can_run_parallel {
12784 let mut results = Vec::new();
12785 for tc in tool_calls {
12786 let result = self
12787 .observe_purpose(
12788 current_observation_context()
12789 .map(|context| context.purpose)
12790 .unwrap_or_default(),
12791 self.execute_tool_smart(tc),
12792 )
12793 .await;
12794 results.push((tc.id.clone(), result));
12795 }
12796 return results;
12797 }
12798
12799 let chunks: Vec<_> = tool_calls
12800 .chunks(self.parallel_tools.max_parallel)
12801 .collect();
12802
12803 let mut all_results = Vec::new();
12804
12805 for chunk in chunks {
12806 let futures: Vec<_> = chunk
12807 .iter()
12808 .map(|tc| {
12809 let tc = tc.clone();
12810 async move {
12811 let result = self.execute_tool_smart(&tc).await;
12812 (tc.id.clone(), result)
12813 }
12814 })
12815 .collect();
12816
12817 let results = futures::future::join_all(futures).await;
12818 all_results.extend(results);
12819 }
12820
12821 all_results
12822 }
12823
12824 pub async fn chat_stream<'a>(
12828 &'a self,
12829 input: &'a str,
12830 ) -> Result<Pin<Box<dyn Stream<Item = StreamChunk> + Send + 'a>>> {
12831 let RootTurnAdmission {
12832 guard: root_turn_guard,
12833 identity_stack,
12834 } = self.acquire_root_turn().await?;
12835 scope_runtime_gate_identity_stack(&identity_stack, self.init_storage()).await?;
12839 info!(input_len = input.len(), "Starting streaming chat");
12840 let terminal = new_runtime_stream_terminal_slot();
12841 let inner = self.run_loop_stream(input, terminal);
12842 let observation_context = self.build_observation_context(None);
12843 let stream: Pin<Box<dyn Stream<Item = StreamChunk> + Send + 'a>> =
12844 Box::pin(async_stream::stream! {
12845 let mut root_turn_guard = Some(root_turn_guard);
12846 let mut inner = inner;
12847 loop {
12848 let next = scope_runtime_gate_identity_stack(&identity_stack, async {
12849 if let Some(context) = observation_context.as_ref() {
12850 with_observation_context(context.clone(), inner.next()).await
12851 } else {
12852 inner.next().await
12853 }
12854 })
12855 .await;
12856 match next {
12857 Some(StreamChunk::Done {}) => {
12858 while scope_runtime_gate_identity_stack(&identity_stack, async {
12859 if let Some(context) = observation_context.as_ref() {
12860 with_observation_context(context.clone(), inner.next())
12861 .await
12862 .is_some()
12863 } else {
12864 inner.next().await.is_some()
12865 }
12866 })
12867 .await
12868 {}
12869 if observation_context.is_some() {
12870 scope_runtime_gate_identity_stack(
12871 &identity_stack,
12872 self.export_observability_if_configured(),
12873 )
12874 .await;
12875 }
12876 drop(root_turn_guard.take());
12877 yield StreamChunk::Done {};
12878 return;
12879 }
12880 Some(chunk) => yield chunk,
12881 None => {
12882 if observation_context.is_some() {
12883 scope_runtime_gate_identity_stack(
12884 &identity_stack,
12885 self.export_observability_if_configured(),
12886 )
12887 .await;
12888 }
12889 drop(root_turn_guard.take());
12890 return;
12891 }
12892 }
12893 }
12894 });
12895 Ok(stream)
12896 }
12897
12898 pub async fn chat_stream_events<'a>(
12902 &'a self,
12903 input: &'a str,
12904 ) -> Result<Pin<Box<dyn Stream<Item = AgentStreamEvent> + Send + 'a>>> {
12905 let RootTurnAdmission {
12906 guard,
12907 identity_stack,
12908 } = self.acquire_root_turn().await?;
12909 scope_runtime_gate_identity_stack(&identity_stack, self.init_storage()).await?;
12913 info!(input_len = input.len(), "Starting streaming chat events");
12914 let terminal = new_runtime_stream_terminal_slot();
12915 let inner = self.run_loop_stream(input, Arc::clone(&terminal));
12916 let observation_context = self.build_observation_context(None);
12917 Ok(self.drive_event_stream(
12918 inner,
12919 terminal,
12920 guard,
12921 identity_stack,
12922 observation_context,
12923 None,
12924 ))
12925 }
12926
12927 pub async fn chat_stream_events_with_actor_context<'a>(
12933 &'a self,
12934 input: &'a str,
12935 actor_context: crate::TurnActorContext,
12936 ) -> Result<Pin<Box<dyn Stream<Item = AgentStreamEvent> + Send + 'a>>> {
12937 let RootTurnAdmission {
12938 guard,
12939 identity_stack,
12940 } = self.acquire_root_turn().await?;
12941 scope_runtime_gate_identity_stack(&identity_stack, self.init_storage()).await?;
12942 info!(
12943 input_len = input.len(),
12944 "Starting streaming chat events with actor context"
12945 );
12946 let actor_id = actor_context.effective_actor_id().map(str::to_string);
12947 let terminal = new_runtime_stream_terminal_slot();
12948 let inner = self.run_loop_stream(input, Arc::clone(&terminal));
12949 let observation_context = self.build_observation_context(actor_id);
12950 Ok(self.drive_event_stream(
12951 inner,
12952 terminal,
12953 guard,
12954 identity_stack,
12955 observation_context,
12956 Some(actor_context),
12957 ))
12958 }
12959
12960 fn drive_event_stream<'a>(
12969 &'a self,
12970 mut inner: Pin<Box<dyn Stream<Item = StreamChunk> + Send + 'a>>,
12971 terminal: RuntimeStreamTerminalSlot,
12972 root_turn_guard: tokio::sync::OwnedMutexGuard<()>,
12973 identity_stack: RootTurnGateIdentityStack,
12974 observation_context: Option<SpanContext>,
12975 actor_context: Option<crate::TurnActorContext>,
12976 ) -> Pin<Box<dyn Stream<Item = AgentStreamEvent> + Send + 'a>> {
12977 Box::pin(async_stream::stream! {
12978 let mut root_turn_guard = Some(root_turn_guard);
12979 loop {
12980 let next = poll_scoped_chunk(
12981 &mut inner,
12982 &identity_stack,
12983 observation_context.as_ref(),
12984 actor_context.as_ref(),
12985 )
12986 .await;
12987 match next {
12988 Some(StreamChunk::Done {}) => {
12989 let terminal_event = { terminal.write().take() };
12990 if let Some(response) = terminal_event {
12991 while poll_scoped_chunk(
12992 &mut inner,
12993 &identity_stack,
12994 observation_context.as_ref(),
12995 actor_context.as_ref(),
12996 )
12997 .await
12998 .is_some()
12999 {}
13000 if observation_context.is_some() {
13001 scope_runtime_gate_identity_stack(
13002 &identity_stack,
13003 self.export_observability_if_configured(),
13004 )
13005 .await;
13006 }
13007 drop(root_turn_guard.take());
13008 yield AgentStreamEvent::Final(response);
13009 return;
13010 }
13011 }
13012 Some(StreamChunk::Error { message }) => {
13013 let finalized = { terminal.read().is_some() };
13014 if finalized {
13015 continue;
13016 }
13017 while poll_scoped_chunk(
13018 &mut inner,
13019 &identity_stack,
13020 observation_context.as_ref(),
13021 actor_context.as_ref(),
13022 )
13023 .await
13024 .is_some()
13025 {}
13026 if observation_context.is_some() {
13027 scope_runtime_gate_identity_stack(
13028 &identity_stack,
13029 self.export_observability_if_configured(),
13030 )
13031 .await;
13032 }
13033 drop(root_turn_guard.take());
13034 yield AgentStreamEvent::Chunk(StreamChunk::Error { message });
13035 return;
13036 }
13037 Some(chunk) => yield AgentStreamEvent::Chunk(chunk),
13038 None => {
13039 if observation_context.is_some() {
13040 scope_runtime_gate_identity_stack(
13041 &identity_stack,
13042 self.export_observability_if_configured(),
13043 )
13044 .await;
13045 }
13046 drop(root_turn_guard.take());
13047 return;
13048 }
13049 }
13050 }
13051 })
13052 }
13053}
13054
13055async fn poll_scoped_chunk<'a>(
13061 inner: &mut Pin<Box<dyn Stream<Item = StreamChunk> + Send + 'a>>,
13062 identity_stack: &RootTurnGateIdentityStack,
13063 observation_context: Option<&SpanContext>,
13064 actor_context: Option<&crate::TurnActorContext>,
13065) -> Option<StreamChunk> {
13066 scope_runtime_gate_identity_stack(identity_stack, async {
13067 let next = inner.next();
13068 match (observation_context, actor_context) {
13069 (Some(observation), Some(actor)) => {
13070 with_observation_context(
13071 observation.clone(),
13072 scope_actor_context(actor.clone(), next),
13073 )
13074 .await
13075 }
13076 (Some(observation), None) => with_observation_context(observation.clone(), next).await,
13077 (None, Some(actor)) => scope_actor_context(actor.clone(), next).await,
13078 (None, None) => next.await,
13079 }
13080 })
13081 .await
13082}
13083
13084#[async_trait]
13085impl ToolInvoker for RuntimeAgent {
13086 async fn invoke_tool(&self, request: ToolExecutionRequest) -> Result<ToolExecutionRecord> {
13087 self.execute_tool_record(request).await
13088 }
13089}
13090
13091#[async_trait]
13092impl Agent for RuntimeAgent {
13093 async fn chat(&self, input: &str) -> Result<AgentResponse> {
13095 let RootTurnAdmission {
13096 guard,
13097 identity_stack,
13098 } = self.acquire_root_turn().await?;
13099 let result = scope_runtime_gate_identity_stack(&identity_stack, async {
13100 let result = if let Some(context) = self.build_observation_context(None) {
13101 with_observation_context(context, self.run_loop(input)).await
13102 } else {
13103 self.run_loop(input).await
13104 };
13105 self.export_observability_if_configured().await;
13106 result
13107 })
13108 .await;
13109 drop(guard);
13110 result
13111 }
13112
13113 fn info(&self) -> AgentInfo {
13114 self.info.clone()
13115 }
13116
13117 async fn reset(&self) -> Result<()> {
13119 self.reset_runtime_state().await
13120 }
13121}
13122
13123fn background_maintenance_tags(
13133 label: &str,
13134 stage: &str,
13135 reason: Option<&str>,
13136 policy: Option<&crate::optimization::config::MaintenanceTaskPolicy>,
13137) -> HashMap<String, String> {
13138 let mut tags = HashMap::new();
13139 tags.insert("runtime.background".to_string(), "true".to_string());
13140 tags.insert("runtime.maintenance".to_string(), label.to_string());
13141 tags.insert("runtime.maintenance_stage".to_string(), stage.to_string());
13142 if let Some(policy) = policy {
13143 tags.insert(
13144 "runtime.await_before_next_turn".to_string(),
13145 await_before_next_turn_label(policy.await_before_next_turn).to_string(),
13146 );
13147 tags.insert(
13148 "runtime.maintenance_mode".to_string(),
13149 maintenance_mode_label(policy.mode).to_string(),
13150 );
13151 }
13152 if let Some(reason) = reason {
13153 tags.insert("runtime.reason".to_string(), reason.to_string());
13154 }
13155 tags
13156}
13157
13158fn await_before_next_turn_label(policy: AwaitBeforeNextTurn) -> &'static str {
13159 match policy {
13160 AwaitBeforeNextTurn::Never => "never",
13161 AwaitBeforeNextTurn::SameActor => "same_actor",
13162 AwaitBeforeNextTurn::Always => "always",
13163 }
13164}
13165
13166fn maintenance_mode_label(mode: MaintenanceMode) -> &'static str {
13167 match mode {
13168 MaintenanceMode::InlineSerial => "inline_serial",
13169 MaintenanceMode::InlineParallel => "inline_parallel",
13170 MaintenanceMode::Background => "background",
13171 }
13172}
13173
13174fn record_background_maintenance_event(
13176 manager: Option<&Arc<ObservabilityManager>>,
13177 label: &str,
13178 status: EventStatus,
13179 duration_ms: u64,
13180 stage: &str,
13181 reason: Option<String>,
13182 policy: Option<&crate::optimization::config::MaintenanceTaskPolicy>,
13183) {
13184 if let Some(manager) = manager {
13185 manager.record_lifecycle_event(
13186 EventType::MemoryOperation {
13187 operation: format!("{}_background_{}", label, stage),
13188 },
13189 ObservationPurpose::Other(format!("{}_maintenance", label)),
13190 status,
13191 duration_ms,
13192 background_maintenance_tags(label, stage, reason.as_deref(), policy),
13193 None,
13194 );
13195 }
13196}
13197
13198fn effective_maintenance_mode(mode: MaintenanceMode, force_parallel: bool) -> MaintenanceMode {
13199 if force_parallel && matches!(mode, MaintenanceMode::InlineSerial) {
13200 MaintenanceMode::InlineParallel
13201 } else {
13202 mode
13203 }
13204}
13205
13206fn observation_purpose_for_process(hint: ProcessPurposeHint) -> ObservationPurpose {
13207 match hint {
13208 ProcessPurposeHint::Detect => ObservationPurpose::ProcessDetect,
13209 ProcessPurposeHint::Extract => ObservationPurpose::ProcessExtract,
13210 ProcessPurposeHint::Validate => ObservationPurpose::ProcessValidate,
13211 ProcessPurposeHint::Transform | ProcessPurposeHint::Other => {
13212 ObservationPurpose::ProcessTransform
13213 }
13214 }
13215}
13216
13217fn new_tool_resource_locks() -> ToolResourceLocks {
13218 Arc::new(RwLock::new(HashMap::new()))
13219}
13220
13221fn tool_resource_lock_keys(
13226 _canonical_id: &str,
13227 args: &Value,
13228 bindings: &ai_agents_core::ToolPolicyBindings,
13229 classification: &ai_agents_core::ToolCallClassification,
13230) -> Vec<String> {
13231 if classification.concurrency_safe {
13232 return Vec::new();
13233 }
13234
13235 let mut keys = Vec::new();
13236 let mut has_path_resource = false;
13237 for binding in &bindings.path_fields {
13238 let value = value_at_argument_path(args, &binding.field)
13239 .cloned()
13240 .or_else(|| {
13241 binding
13242 .default_path
13243 .as_ref()
13244 .map(|path| Value::String(path.clone()))
13245 });
13246 if let Some(value) = value {
13247 collect_resource_strings(&value, |_| {
13248 has_path_resource = true;
13249 });
13250 }
13251 }
13252 for binding in &bindings.domain_fields {
13253 if let Some(value) = value_at_argument_path(args, &binding.field) {
13254 collect_resource_strings(value, |domain| {
13255 let normalized = if binding.is_url {
13256 normalized_url_resource_key(domain)
13257 } else {
13258 domain.trim().trim_end_matches('.').to_ascii_lowercase()
13259 };
13260 keys.push(format!("domain:{}", normalized));
13261 });
13262 }
13263 }
13264 for binding in &bindings.command_fields {
13265 if !matches!(binding.kind, ai_agents_core::CommandBindingKind::Cwd) {
13266 continue;
13267 }
13268 if let Some(value) = value_at_argument_path(args, &binding.field) {
13269 collect_resource_strings(value, |_| {
13270 has_path_resource = true;
13271 });
13272 }
13273 }
13274 if has_path_resource {
13275 keys.push("path-mutation:global".to_string());
13276 }
13277 if keys.is_empty() {
13278 keys.push("side-effect:unbound".to_string());
13279 }
13280 keys.sort();
13281 keys.dedup();
13282 keys
13283}
13284
13285fn value_at_argument_path<'a>(value: &'a Value, field: &str) -> Option<&'a Value> {
13286 let mut current = value;
13287 for segment in field.split('.') {
13288 if segment.is_empty() {
13289 return None;
13290 }
13291 current = current.get(segment)?;
13292 }
13293 Some(current)
13294}
13295
13296fn collect_resource_strings(value: &Value, mut collect: impl FnMut(&str)) {
13297 match value {
13298 Value::String(value) => collect(value),
13299 Value::Array(values) => {
13300 for value in values {
13301 if let Some(value) = value.as_str() {
13302 collect(value);
13303 }
13304 }
13305 }
13306 _ => {}
13307 }
13308}
13309
13310fn normalized_url_resource_key(value: &str) -> String {
13311 let value = value.trim();
13312 let Some((scheme, remainder)) = value.split_once("://") else {
13313 return value.to_ascii_lowercase();
13314 };
13315 let authority_end = remainder.find(['/', '?', '#']).unwrap_or(remainder.len());
13316 let (authority, suffix) = remainder.split_at(authority_end);
13317 format!(
13318 "{}://{}{}",
13319 scheme.to_ascii_lowercase(),
13320 authority.to_ascii_lowercase(),
13321 suffix
13322 )
13323}
13324
13325fn render_concurrent_template(
13326 template: &str,
13327 user_input: &str,
13328 context_values: &std::collections::HashMap<String, serde_json::Value>,
13329) -> Result<String> {
13330 let mut env = minijinja::Environment::new();
13331 env.add_template("concurrent", template)
13332 .map_err(|e| AgentError::Other(format!("Concurrent template parse error: {}", e)))?;
13333
13334 let mut ctx = std::collections::BTreeMap::new();
13335 ctx.insert("user_input".to_string(), minijinja::Value::from(user_input));
13336
13337 let context_obj = minijinja::Value::from_serialize(context_values);
13339 ctx.insert("context".to_string(), context_obj);
13340
13341 let tmpl = env
13342 .get_template("concurrent")
13343 .map_err(|e| AgentError::Other(format!("Concurrent template error: {}", e)))?;
13344
13345 tmpl.render(minijinja::Value::from_serialize(&ctx))
13346 .map_err(|e| AgentError::Other(format!("Concurrent template render error: {}", e)))
13347}
13348
13349#[cfg(test)]
13350mod tests {
13351 use super::*;
13352 use crate::AgentBuilder;
13353 use ai_agents_core::{LLMChunk, LLMConfig, LLMError, LLMFeature, Tool};
13354 use ai_agents_llm::mock::MockLLMProvider;
13355 use ai_agents_skills::{SkillDefinition, SkillStep};
13356 use ai_agents_tools::{
13357 CalculatorTool, CopyPathTool, DeletePathTool, FileWriteTool, MovePathTool, ToolAliases,
13358 ToolDescriptor, ToolProvider, ToolProviderError, ToolProviderType, WebFetchResolver,
13359 WebFetchTool, WebFetchTransport, WebFetchTransportRequest, WebFetchTransportResponse,
13360 };
13361
13362 fn mock_with_response(response: &str) -> MockLLMProvider {
13363 let mut mock = MockLLMProvider::new("test");
13364 mock.set_response(response);
13365 mock
13366 }
13367
13368 fn mock_with_responses(responses: Vec<&str>) -> MockLLMProvider {
13369 let mut mock = MockLLMProvider::new("test");
13370 mock.set_responses(responses.into_iter().map(String::from).collect(), true);
13371 mock
13372 }
13373
13374 async fn collect_stream_events(
13376 agent: &RuntimeAgent,
13377 input: &str,
13378 ) -> (String, Vec<StreamChunk>, Option<AgentResponse>) {
13379 use futures::StreamExt;
13380 let mut events = agent.chat_stream_events(input).await.expect("stream opens");
13381 let mut content = String::new();
13382 let mut chunks = Vec::new();
13383 let mut final_response = None;
13384 while let Some(event) = events.next().await {
13385 match event {
13386 AgentStreamEvent::Chunk(chunk) => {
13387 if let StreamChunk::Content { text } = &chunk {
13388 content.push_str(text);
13389 }
13390 chunks.push(chunk);
13391 }
13392 AgentStreamEvent::Final(response) => final_response = Some(response),
13393 }
13394 }
13395 (content, chunks, final_response)
13396 }
13397
13398 fn metadata_keys(response: &AgentResponse) -> std::collections::BTreeSet<String> {
13399 response
13400 .metadata
13401 .as_ref()
13402 .map(|m| m.keys().cloned().collect())
13403 .unwrap_or_default()
13404 }
13405
13406 async fn assert_blocking_streaming_parity<F>(
13409 build: F,
13410 input: &str,
13411 ) -> (AgentResponse, AgentResponse, Vec<StreamChunk>)
13412 where
13413 F: Fn() -> RuntimeAgent,
13414 {
13415 let blocking_agent = build();
13416 let streaming_agent = build();
13417
13418 let blocking = blocking_agent
13419 .chat(input)
13420 .await
13421 .expect("blocking chat succeeds");
13422 let (_, chunks, final_response) = collect_stream_events(&streaming_agent, input).await;
13423 let streamed = final_response.expect("streaming must emit Final when blocking succeeds");
13424
13425 assert_eq!(
13426 blocking.content, streamed.content,
13427 "committed content differs"
13428 );
13429 assert_eq!(
13430 metadata_keys(&blocking),
13431 metadata_keys(&streamed),
13432 "metadata key sets differ"
13433 );
13434 assert_eq!(
13435 blocking.tool_calls.as_ref().map(Vec::len),
13436 streamed.tool_calls.as_ref().map(Vec::len),
13437 "tool call counts differ"
13438 );
13439 assert_eq!(
13440 blocking_agent.current_state(),
13441 streaming_agent.current_state(),
13442 "final states differ"
13443 );
13444 (blocking, streamed, chunks)
13445 }
13446
13447 fn signed_calculator_response(
13448 exchange_id: &str,
13449 call_id: &str,
13450 expression: &str,
13451 ) -> LLMResponse {
13452 let call = ToolCall {
13453 id: call_id.to_string(),
13454 name: "calculator".to_string(),
13455 arguments: serde_json::json!({"expression": expression}),
13456 };
13457 let state = ai_agents_core::NativeProviderState::new(
13458 exchange_id,
13459 "fixture",
13460 "native-tools",
13461 ai_agents_core::NativeProviderTarget::new("https://fixture.invalid/", "fixture-model")
13462 .unwrap(),
13463 serde_json::json!({
13464 "role": "model",
13465 "parts": [{
13466 "functionCall": {"name": "calculator", "args": {"expression": expression}},
13467 "thoughtSignature": format!("signature-{exchange_id}")
13468 }]
13469 }),
13470 vec![ai_agents_core::NativeCallBinding::new(call_id, 0).unwrap()],
13471 )
13472 .unwrap();
13473 LLMResponse::new("", FinishReason::ToolCall)
13474 .with_provider_state(state)
13475 .unwrap()
13476 .with_tool_calls(vec![call])
13477 .unwrap()
13478 }
13479
13480 struct TerminalHistoryProvider {
13481 calls: Arc<std::sync::atomic::AtomicU32>,
13482 }
13483
13484 struct DroppingSignedAssistantMemory {
13485 messages: RwLock<Vec<ChatMessage>>,
13486 }
13487
13488 struct DroppingEarlierSequentialMemory {
13489 messages: RwLock<Vec<ChatMessage>>,
13490 signed_seen: std::sync::atomic::AtomicUsize,
13491 }
13492
13493 #[async_trait]
13494 impl ai_agents_core::Memory for DroppingSignedAssistantMemory {
13495 async fn add_message(&self, message: ChatMessage) -> Result<()> {
13496 let signed = message.role == ai_agents_core::Role::Assistant
13497 && ai_agents_core::decode_native_tool_call_markers(&message.content)
13498 .map_err(|error| AgentError::LLM(error.to_string()))?
13499 .is_some_and(|batch| batch.provider_state().is_some());
13500 if !signed {
13501 self.messages.write().push(message);
13502 }
13503 Ok(())
13504 }
13505
13506 async fn get_messages(&self, limit: Option<usize>) -> Result<Vec<ChatMessage>> {
13507 let messages = self.messages.read();
13508 let start = limit
13509 .map(|limit| messages.len().saturating_sub(limit))
13510 .unwrap_or(0);
13511 Ok(messages[start..].to_vec())
13512 }
13513
13514 async fn clear(&self) -> Result<()> {
13515 self.messages.write().clear();
13516 Ok(())
13517 }
13518
13519 fn len(&self) -> usize {
13520 self.messages.read().len()
13521 }
13522
13523 async fn restore(&self, snapshot: ai_agents_core::MemorySnapshot) -> Result<()> {
13524 *self.messages.write() = snapshot.messages;
13525 Ok(())
13526 }
13527 }
13528
13529 #[async_trait]
13530 impl ai_agents_memory::Memory for DroppingSignedAssistantMemory {}
13531
13532 #[async_trait]
13533 impl ai_agents_core::Memory for DroppingEarlierSequentialMemory {
13534 async fn add_message(&self, message: ChatMessage) -> Result<()> {
13535 let signed = message.role == ai_agents_core::Role::Assistant
13536 && ai_agents_core::decode_native_tool_call_markers(&message.content)
13537 .map_err(|error| AgentError::LLM(error.to_string()))?
13538 .is_some_and(|batch| batch.provider_state().is_some());
13539 let mut messages = self.messages.write();
13540 if signed && self.signed_seen.fetch_add(1, Ordering::SeqCst) == 1 {
13541 messages.retain(|stored| !stored.content.contains("seq-call-1"));
13542 }
13543 messages.push(message);
13544 Ok(())
13545 }
13546
13547 async fn get_messages(&self, limit: Option<usize>) -> Result<Vec<ChatMessage>> {
13548 let messages = self.messages.read();
13549 let start = limit
13550 .map(|limit| messages.len().saturating_sub(limit))
13551 .unwrap_or(0);
13552 Ok(messages[start..].to_vec())
13553 }
13554
13555 async fn clear(&self) -> Result<()> {
13556 self.messages.write().clear();
13557 self.signed_seen.store(0, Ordering::SeqCst);
13558 Ok(())
13559 }
13560
13561 fn len(&self) -> usize {
13562 self.messages.read().len()
13563 }
13564
13565 async fn restore(&self, snapshot: ai_agents_core::MemorySnapshot) -> Result<()> {
13566 *self.messages.write() = snapshot.messages;
13567 self.signed_seen.store(0, Ordering::SeqCst);
13568 Ok(())
13569 }
13570 }
13571
13572 #[async_trait]
13573 impl ai_agents_memory::Memory for DroppingEarlierSequentialMemory {}
13574
13575 #[async_trait]
13576 impl LLMProvider for TerminalHistoryProvider {
13577 async fn complete(
13578 &self,
13579 _messages: &[ChatMessage],
13580 _config: Option<&LLMConfig>,
13581 ) -> std::result::Result<LLMResponse, LLMError> {
13582 self.calls.fetch_add(1, Ordering::SeqCst);
13583 Err(LLMError::Serialization(
13584 "native history integrity failure".to_string(),
13585 ))
13586 }
13587
13588 async fn complete_stream(
13589 &self,
13590 _messages: &[ChatMessage],
13591 _config: Option<&LLMConfig>,
13592 ) -> std::result::Result<
13593 Box<dyn Stream<Item = std::result::Result<LLMChunk, LLMError>> + Unpin + Send>,
13594 LLMError,
13595 > {
13596 Err(LLMError::Serialization(
13597 "native history integrity failure".to_string(),
13598 ))
13599 }
13600
13601 fn provider_name(&self) -> &str {
13602 "terminal-history"
13603 }
13604
13605 fn supports(&self, _feature: LLMFeature) -> bool {
13606 false
13607 }
13608
13609 fn is_terminal_error(&self, error: &LLMError) -> bool {
13610 matches!(error, LLMError::Serialization(_))
13611 }
13612 }
13613
13614 fn disambiguation_state_machine(
13616 state_enabled: Option<bool>,
13617 require_confirmation: bool,
13618 ) -> Arc<StateMachine> {
13619 let definition = ai_agents_state::StateDefinition {
13620 prompt: Some("Handle the resolved request.".to_string()),
13621 disambiguation: Some(ai_agents_disambiguation::StateDisambiguationOverride {
13622 enabled: state_enabled,
13623 require_confirmation,
13624 ..Default::default()
13625 }),
13626 ..Default::default()
13627 };
13628 let review = ai_agents_state::StateDefinition {
13629 prompt: Some("Review a fresh request.".to_string()),
13630 ..Default::default()
13631 };
13632 Arc::new(
13633 StateMachine::new(ai_agents_state::StateConfig {
13634 initial: "active".to_string(),
13635 states: std::collections::HashMap::from([
13636 ("active".to_string(), definition),
13637 ("review".to_string(), review),
13638 ]),
13639 global_transitions: Vec::new(),
13640 fallback: None,
13641 max_no_transition: None,
13642 regenerate_on_transition: true,
13643 })
13644 .unwrap(),
13645 )
13646 }
13647
13648 fn state_disambiguation_agent(
13650 responses: Vec<&str>,
13651 manager_enabled: bool,
13652 state_enabled: Option<bool>,
13653 require_confirmation: bool,
13654 ) -> (RuntimeAgent, MockLLMProvider) {
13655 state_disambiguation_agent_with_skills(
13656 responses,
13657 manager_enabled,
13658 state_enabled,
13659 require_confirmation,
13660 Vec::new(),
13661 )
13662 }
13663
13664 fn state_disambiguation_agent_with_skills(
13666 responses: Vec<&str>,
13667 manager_enabled: bool,
13668 state_enabled: Option<bool>,
13669 require_confirmation: bool,
13670 skills: Vec<SkillDefinition>,
13671 ) -> (RuntimeAgent, MockLLMProvider) {
13672 let mut mock = MockLLMProvider::new("state-confirmation");
13673 mock.set_responses(responses.into_iter().map(String::from).collect(), false);
13674 let observed = mock.clone();
13675 let agent = AgentBuilder::new()
13676 .system_prompt("Handle requests.")
13677 .llm(Arc::new(mock.clone()))
13678 .llm_alias("router", Arc::new(mock))
13679 .state_machine(disambiguation_state_machine(
13680 state_enabled,
13681 require_confirmation,
13682 ))
13683 .skills(skills)
13684 .build()
13685 .unwrap()
13686 .with_disambiguation(DisambiguationConfig {
13687 enabled: manager_enabled,
13688 ..Default::default()
13689 });
13690 (agent, observed)
13691 }
13692
13693 fn confirmation_skill() -> SkillDefinition {
13695 SkillDefinition {
13696 id: "send_report".to_string(),
13697 description: "Send a report after clarification".to_string(),
13698 trigger: "When the user asks to send a report".to_string(),
13699 steps: vec![SkillStep::Prompt {
13700 prompt: "Execute confirmed report skill for: {{ input }}".to_string(),
13701 llm: None,
13702 }],
13703 reasoning: None,
13704 reflection: None,
13705 disambiguation: Some(ai_agents_disambiguation::SkillDisambiguationOverride {
13706 enabled: Some(true),
13707 ..Default::default()
13708 }),
13709 }
13710 }
13711
13712 fn confirmation_skill_call_count(observed: &MockLLMProvider) -> usize {
13714 observed
13715 .call_history()
13716 .iter()
13717 .filter(|call| {
13718 call.messages
13719 .iter()
13720 .any(|message| message.content.contains("Execute confirmed report skill"))
13721 })
13722 .count()
13723 }
13724
13725 struct BlockingRuntimeConfirmationObserver {
13726 entered: tokio::sync::Barrier,
13727 release: tokio::sync::Notify,
13728 }
13729
13730 impl BlockingRuntimeConfirmationObserver {
13731 fn new() -> Self {
13732 Self {
13733 entered: tokio::sync::Barrier::new(2),
13734 release: tokio::sync::Notify::new(),
13735 }
13736 }
13737 }
13738
13739 struct ResetOnTransitionHooks {
13740 agent: parking_lot::Mutex<Option<Weak<RuntimeAgent>>>,
13741 invoked: AtomicBool,
13742 }
13743
13744 #[async_trait]
13745 impl AgentHooks for ResetOnTransitionHooks {
13746 async fn on_state_transition(&self, _from: Option<&str>, _to: &str, _reason: &str) {
13747 if self.invoked.swap(true, Ordering::SeqCst) {
13748 return;
13749 }
13750 let agent = self.agent.lock().as_ref().and_then(Weak::upgrade);
13751 if let Some(agent) = agent {
13752 agent.reset().await.unwrap();
13753 }
13754 }
13755 }
13756
13757 impl ClarificationObserver for BlockingRuntimeConfirmationObserver {
13758 fn observe_question<'a>(
13759 &'a self,
13760 future: ClarificationQuestionFuture<'a>,
13761 ) -> ClarificationQuestionFuture<'a> {
13762 future
13763 }
13764
13765 fn observe_parse<'a>(
13766 &'a self,
13767 future: ClarificationParseFuture<'a>,
13768 ) -> ClarificationParseFuture<'a> {
13769 future
13770 }
13771
13772 fn observe_confirmation_parse<'a>(
13773 &'a self,
13774 future: ConfirmationParseFuture<'a>,
13775 ) -> ConfirmationParseFuture<'a> {
13776 Box::pin(async move {
13777 self.entered.wait().await;
13778 self.release.notified().await;
13779 future.await
13780 })
13781 }
13782 }
13783
13784 #[tokio::test]
13785 async fn state_confirmation_blocks_redispatch_until_explicit_agreement() {
13786 let (agent, observed) = state_disambiguation_agent(
13787 vec![
13788 r#"{"is_ambiguous":true,"confidence":0.2,"ambiguity_type":"missing_target","reasoning":"target missing","what_is_unclear":["target"],"detected_language":"en"}"#,
13789 r#"{"question":"What should I send?","options":null}"#,
13790 r#"{"status":"answered","selected_option":null,"enriched_input":"Send the report to Ada","resolved":{"intent":"send_report"}}"#,
13791 r#"{"question":"Should I send the report to Ada?"}"#,
13792 r#"{"status":"confirmed"}"#,
13793 "Request executed.",
13794 ],
13795 true,
13796 None,
13797 true,
13798 );
13799
13800 let clarification = agent.chat("Send it").await.unwrap();
13801 assert_eq!(clarification.content, "What should I send?");
13802 assert_eq!(observed.call_count(), 2);
13803
13804 let confirmation = agent.chat("The report to Ada").await.unwrap();
13805 assert_eq!(confirmation.content, "Should I send the report to Ada?");
13806 assert_eq!(
13807 confirmation
13808 .metadata
13809 .as_ref()
13810 .and_then(|metadata| metadata.get("disambiguation"))
13811 .and_then(|metadata| metadata.get("status"))
13812 .and_then(Value::as_str),
13813 Some("awaiting_confirmation")
13814 );
13815 assert_eq!(observed.call_count(), 4);
13816
13817 let completed = agent.chat("Yes").await.unwrap();
13818 assert_eq!(completed.content, "Request executed.");
13819 assert_eq!(observed.call_count(), 6);
13820 }
13821
13822 #[tokio::test]
13823 async fn streaming_state_confirmation_ends_the_turn_before_redispatch() {
13824 let (agent, observed) = state_disambiguation_agent(
13825 vec![
13826 r#"{"is_ambiguous":true,"confidence":0.2,"ambiguity_type":"missing_target","reasoning":"target missing","what_is_unclear":["target"],"detected_language":"en"}"#,
13827 r#"{"question":"What should I send?","options":null}"#,
13828 r#"{"status":"answered","selected_option":null,"enriched_input":"Send the report to Ada","resolved":{"intent":"send_report"}}"#,
13829 r#"{"question":"Should I send the report to Ada?"}"#,
13830 r#"{"status":"confirmed"}"#,
13831 "Request executed.",
13832 ],
13833 true,
13834 None,
13835 true,
13836 );
13837
13838 let mut clarification_stream = agent.chat_stream("Send it").await.unwrap();
13839 let mut clarification = String::new();
13840 while let Some(chunk) = clarification_stream.next().await {
13841 match chunk {
13842 StreamChunk::Content { text } => clarification.push_str(&text),
13843 StreamChunk::Done {} => break,
13844 StreamChunk::Error { message } => panic!("unexpected stream error: {message}"),
13845 _ => {}
13846 }
13847 }
13848 assert_eq!(clarification, "What should I send?");
13849 assert_eq!(observed.call_count(), 2);
13850
13851 let mut confirmation_stream = agent.chat_stream_events("The report to Ada").await.unwrap();
13852 let mut confirmation = None;
13853 while let Some(event) = confirmation_stream.next().await {
13854 match event {
13855 AgentStreamEvent::Final(response) => confirmation = Some(response),
13856 AgentStreamEvent::Chunk(StreamChunk::Error { message }) => {
13857 panic!("unexpected stream error: {message}")
13858 }
13859 AgentStreamEvent::Chunk(_) => {}
13860 }
13861 }
13862 let confirmation = confirmation.expect("confirmation must finalize");
13863 assert_eq!(confirmation.content, "Should I send the report to Ada?");
13864 assert_eq!(
13865 confirmation
13866 .metadata
13867 .as_ref()
13868 .and_then(|metadata| metadata.get("disambiguation"))
13869 .and_then(|metadata| metadata.get("status"))
13870 .and_then(Value::as_str),
13871 Some("awaiting_confirmation")
13872 );
13873 assert_eq!(observed.call_count(), 4);
13874
13875 let mut completed_stream = agent.chat_stream("Yes").await.unwrap();
13876 let mut completed = String::new();
13877 while let Some(chunk) = completed_stream.next().await {
13878 match chunk {
13879 StreamChunk::Content { text } => completed.push_str(&text),
13880 StreamChunk::Done {} => break,
13881 StreamChunk::Error { message } => panic!("unexpected stream error: {message}"),
13882 _ => {}
13883 }
13884 }
13885 assert_eq!(completed, "Request executed.");
13886 assert_eq!(observed.call_count(), 6);
13887 }
13888
13889 #[tokio::test]
13891 async fn root_turn_gate_serializes_blocking_and_streaming_entry_points() {
13892 let (complete_entered, mut complete_events) = tokio::sync::mpsc::unbounded_channel();
13893 let agent = Arc::new(
13894 AgentBuilder::new()
13895 .system_prompt("Serialize root turns.")
13896 .llm(Arc::new(RootTurnProbeProvider { complete_entered }))
13897 .build()
13898 .unwrap(),
13899 );
13900 let blocking_agent = Arc::clone(&agent);
13901
13902 let legacy_stream = agent.chat_stream("stream owner").await.unwrap();
13903 assert!(agent.root_turn_gate.try_lock().is_err());
13904 let blocking = tokio::spawn(async move { blocking_agent.chat("blocked").await.unwrap() });
13905 assert!(
13906 tokio::time::timeout(std::time::Duration::from_millis(50), complete_events.recv())
13907 .await
13908 .is_err(),
13909 "blocking turn reached the provider while the legacy stream owned the root gate"
13910 );
13911
13912 drop(legacy_stream);
13913 assert_eq!(
13914 tokio::time::timeout(std::time::Duration::from_secs(2), complete_events.recv())
13915 .await
13916 .expect("blocking turn did not enter after stream drop"),
13917 Some(())
13918 );
13919 let response = tokio::time::timeout(std::time::Duration::from_secs(2), blocking)
13920 .await
13921 .expect("blocking turn did not finish after stream drop")
13922 .unwrap();
13923 assert_eq!(response.content, "blocking complete");
13924
13925 let mut event_stream = agent.chat_stream_events("event terminal").await.unwrap();
13926 assert!(agent.root_turn_gate.try_lock().is_err());
13927 let mut saw_final = false;
13928 while let Some(event) = event_stream.next().await {
13929 if matches!(event, AgentStreamEvent::Final(_)) {
13930 saw_final = true;
13931 break;
13932 }
13933 }
13934 assert!(saw_final);
13935 assert!(
13936 agent.root_turn_gate.try_lock().is_ok(),
13937 "authoritative terminal event retained the root gate"
13938 );
13939 }
13940
13941 #[tokio::test]
13943 async fn response_hook_rejects_same_runtime_chat_reentry() {
13944 let hooks = Arc::new(ResponseChatHooks {
13945 target: parking_lot::Mutex::new(None),
13946 invoked: AtomicBool::new(false),
13947 nested_result: parking_lot::Mutex::new(None),
13948 });
13949 let agent = Arc::new(
13950 AgentBuilder::new()
13951 .system_prompt("Reject response hook reentry.")
13952 .llm(Arc::new(mock_with_response("outer response")))
13953 .hooks(hooks.clone())
13954 .build()
13955 .unwrap(),
13956 );
13957 *hooks.target.lock() = Some(Arc::downgrade(&agent));
13958
13959 let response = tokio::time::timeout(
13960 std::time::Duration::from_secs(2),
13961 agent.chat("outer request"),
13962 )
13963 .await
13964 .expect("same-runtime response hook reentry must fail without deadlocking")
13965 .unwrap();
13966
13967 assert_eq!(response.content, "outer response");
13968 let nested_result = hooks
13969 .nested_result
13970 .lock()
13971 .clone()
13972 .expect("response hook must record its nested call");
13973 let error = nested_result.expect_err("same-runtime nested chat must be rejected");
13974 assert!(error.contains("reentrant root turn ownership"));
13975 }
13976
13977 #[tokio::test]
13979 async fn root_turn_gate_allows_nested_runtime_and_rejects_cycles() {
13980 let agent_a = AgentBuilder::new()
13981 .system_prompt("Runtime A.")
13982 .llm(Arc::new(mock_with_response("response A")))
13983 .build()
13984 .unwrap();
13985 let agent_b = AgentBuilder::new()
13986 .system_prompt("Runtime B.")
13987 .llm(Arc::new(mock_with_response("response B")))
13988 .build()
13989 .unwrap();
13990 let RootTurnAdmission {
13991 guard: guard_a,
13992 identity_stack: stack_a,
13993 } = agent_a.acquire_root_turn().await.unwrap();
13994
13995 let cycle_error = scope_runtime_gate_identity_stack(&stack_a, async {
13996 let RootTurnAdmission {
13997 guard: guard_b,
13998 identity_stack: stack_b,
13999 } = agent_b
14000 .acquire_root_turn()
14001 .await
14002 .expect("runtime B must acquire a different gate");
14003 let result =
14004 scope_runtime_gate_identity_stack(&stack_b, agent_a.acquire_root_turn()).await;
14005 drop(guard_b);
14006 match result {
14007 Err(error) => error,
14008 Ok(_) => panic!("runtime A accepted a repeated gate identity"),
14009 }
14010 })
14011 .await;
14012 drop(guard_a);
14013
14014 assert!(
14015 cycle_error
14016 .to_string()
14017 .contains("reentrant root turn ownership")
14018 );
14019 }
14020
14021 #[tokio::test]
14023 async fn concurrent_orchestration_propagates_root_gate_ancestry() {
14024 let registry = Arc::new(crate::spawner::AgentRegistry::new());
14025 let hooks_a = Arc::new(ConcurrentResponseHooks {
14026 registry: Arc::downgrade(®istry),
14027 child_id: "runtime-b".to_string(),
14028 invoked: AtomicBool::new(false),
14029 nested_result: parking_lot::Mutex::new(None),
14030 });
14031 let hooks_b = Arc::new(ResponseChatHooks {
14032 target: parking_lot::Mutex::new(None),
14033 invoked: AtomicBool::new(false),
14034 nested_result: parking_lot::Mutex::new(None),
14035 });
14036 let agent_a = AgentBuilder::new()
14037 .system_prompt("Runtime A dispatches runtime B concurrently.")
14038 .llm(Arc::new(mock_with_response("response A")))
14039 .hooks(hooks_a.clone())
14040 .build()
14041 .unwrap();
14042 let agent_b = AgentBuilder::new()
14043 .system_prompt("Runtime B attempts to re-enter runtime A.")
14044 .llm(Arc::new(mock_with_response("response B")))
14045 .hooks(hooks_b.clone())
14046 .build()
14047 .unwrap();
14048 let spec_a = crate::spec::AgentSpec {
14049 name: "runtime-a".to_string(),
14050 system_prompt: "Runtime A dispatches runtime B concurrently.".to_string(),
14051 ..crate::spec::AgentSpec::default()
14052 };
14053 let spec_b = crate::spec::AgentSpec {
14054 name: "runtime-b".to_string(),
14055 system_prompt: "Runtime B attempts to re-enter runtime A.".to_string(),
14056 ..crate::spec::AgentSpec::default()
14057 };
14058 registry
14059 .register(crate::spawner::SpawnedAgent::from_runtime(
14060 "runtime-a".to_string(),
14061 agent_a,
14062 spec_a,
14063 ))
14064 .await
14065 .unwrap();
14066 registry
14067 .register(crate::spawner::SpawnedAgent::from_runtime(
14068 "runtime-b".to_string(),
14069 agent_b,
14070 spec_b,
14071 ))
14072 .await
14073 .unwrap();
14074 let runtime_a = registry.get("runtime-a").unwrap();
14075 *hooks_b.target.lock() = Some(Arc::downgrade(&runtime_a));
14076
14077 let response = tokio::time::timeout(
14078 std::time::Duration::from_secs(2),
14079 runtime_a.chat("outer concurrent request"),
14080 )
14081 .await
14082 .expect("concurrent orchestration cycle must fail without deadlocking")
14083 .unwrap();
14084
14085 assert_eq!(response.content, "response A");
14086 let child_result = hooks_a
14087 .nested_result
14088 .lock()
14089 .clone()
14090 .expect("runtime A hook must record runtime B completion");
14091 assert_eq!(child_result.unwrap(), "response B");
14092 let cycle_result = hooks_b
14093 .nested_result
14094 .lock()
14095 .clone()
14096 .expect("runtime B hook must record runtime A reentry");
14097 assert!(
14098 cycle_result
14099 .expect_err("runtime A accepted a repeated gate identity")
14100 .contains("reentrant root turn ownership")
14101 );
14102 }
14103
14104 fn skill_clarification_responses() -> Vec<&'static str> {
14107 vec![
14108 r#"{"is_ambiguous":false,"confidence":0.99,"ambiguity_type":null,"reasoning":"top-level clear","what_is_unclear":[],"detected_language":"en"}"#,
14109 "send_report",
14110 r#"{"is_ambiguous":true,"confidence":0.2,"ambiguity_type":"missing_target","reasoning":"target missing","what_is_unclear":["target"],"detected_language":"en"}"#,
14111 r#"{"question":"What should I send?","options":null}"#,
14112 ]
14113 }
14114
14115 #[tokio::test]
14121 async fn test_stream_skill_clarification_memory_matches_blocking() {
14122 let (blocking_agent, _) = state_disambiguation_agent_with_skills(
14123 skill_clarification_responses(),
14124 true,
14125 None,
14126 true,
14127 vec![confirmation_skill()],
14128 );
14129 let blocking = blocking_agent.chat("Send it").await.unwrap();
14130 let blocking_messages = blocking_agent.memory.get_messages(None).await.unwrap();
14131
14132 let (streaming_agent, _) = state_disambiguation_agent_with_skills(
14133 skill_clarification_responses(),
14134 true,
14135 None,
14136 true,
14137 vec![confirmation_skill()],
14138 );
14139 let (content, chunks, streamed) = collect_stream_events(&streaming_agent, "Send it").await;
14140 assert!(!chunks.iter().any(StreamChunk::is_error), "{chunks:?}");
14141 let streamed = streamed.expect("skill clarification must finalize as Final");
14142 let streaming_messages = streaming_agent.memory.get_messages(None).await.unwrap();
14143
14144 assert_eq!(blocking.content, "What should I send?");
14145 assert_eq!(streamed.content, blocking.content);
14146 assert_eq!(content, streamed.content);
14147 assert_eq!(
14148 blocking
14149 .metadata
14150 .as_ref()
14151 .and_then(|m| m.get("disambiguation")),
14152 streamed
14153 .metadata
14154 .as_ref()
14155 .and_then(|m| m.get("disambiguation")),
14156 );
14157 assert_eq!(
14158 streamed
14159 .metadata
14160 .as_ref()
14161 .and_then(|m| m.get("disambiguation"))
14162 .and_then(|d| d.get("status"))
14163 .and_then(Value::as_str),
14164 Some("awaiting_clarification"),
14165 );
14166 let shape = |messages: &[ChatMessage]| {
14167 messages
14168 .iter()
14169 .map(|m| (format!("{:?}", m.role), m.content.clone()))
14170 .collect::<Vec<_>>()
14171 };
14172 assert_eq!(shape(&blocking_messages), shape(&streaming_messages));
14173 assert_eq!(
14174 shape(&streaming_messages),
14175 vec![
14176 ("User".to_string(), "Send it".to_string()),
14177 ("Assistant".to_string(), "What should I send?".to_string()),
14178 ],
14179 );
14180 assert_eq!(
14181 *streaming_agent.pending_skill_id.read(),
14182 Some("send_report".to_string()),
14183 );
14184 }
14185
14186 #[tokio::test]
14188 async fn test_stream_skill_clarification_memory_failure_surfaces_as_error() {
14189 let build = || {
14191 let mut mock = MockLLMProvider::new("skill-clarification");
14192 mock.set_responses(
14193 skill_clarification_responses()
14194 .into_iter()
14195 .map(String::from)
14196 .collect(),
14197 false,
14198 );
14199 AgentBuilder::new()
14200 .system_prompt("Handle requests.")
14201 .llm(Arc::new(mock.clone()))
14202 .llm_alias("router", Arc::new(mock))
14203 .state_machine(disambiguation_state_machine(None, true))
14204 .skills(vec![confirmation_skill()])
14205 .memory(Arc::new(FailingMemory {
14206 messages: parking_lot::RwLock::new(Vec::new()),
14207 fail_on_add: 2,
14208 adds: std::sync::atomic::AtomicUsize::new(0),
14209 }))
14210 .build()
14211 .unwrap()
14212 .with_disambiguation(DisambiguationConfig {
14213 enabled: true,
14214 ..Default::default()
14215 })
14216 };
14217
14218 let blocking = build().chat("Send it").await;
14219 assert!(
14220 blocking.is_err(),
14221 "blocking must surface the failed clarification write: {blocking:?}"
14222 );
14223
14224 let (_, chunks, streamed) = collect_stream_events(&build(), "Send it").await;
14225 assert!(
14226 streamed.is_none(),
14227 "a failed write must not finalize the turn"
14228 );
14229 assert!(
14230 chunks.iter().any(|chunk| matches!(
14231 chunk,
14232 StreamChunk::Error { message } if message.contains("simulated memory failure")
14233 )),
14234 "streaming must surface the failed clarification write: {chunks:?}"
14235 );
14236 }
14237
14238 #[tokio::test]
14240 async fn confirmed_skill_route_executes_exactly_once() {
14241 let (agent, observed) = state_disambiguation_agent_with_skills(
14242 vec![
14243 r#"{"is_ambiguous":false,"confidence":0.99,"ambiguity_type":null,"reasoning":"top-level clear","what_is_unclear":[],"detected_language":"en"}"#,
14244 "send_report",
14245 r#"{"is_ambiguous":true,"confidence":0.2,"ambiguity_type":"missing_target","reasoning":"target missing","what_is_unclear":["target"],"detected_language":"en"}"#,
14246 r#"{"question":"What should I send?","options":null}"#,
14247 r#"{"status":"answered","selected_option":null,"enriched_input":"Send the report to Ada","resolved":{"intent":"send_report"}}"#,
14248 r#"{"question":"Should I send the report to Ada?"}"#,
14249 r#"{"status":"confirmed"}"#,
14250 r#"{"is_ambiguous":false,"confidence":0.99,"ambiguity_type":null,"reasoning":"resolved","what_is_unclear":[],"detected_language":"en"}"#,
14251 "Report skill executed.",
14252 ],
14253 true,
14254 None,
14255 true,
14256 vec![confirmation_skill()],
14257 );
14258
14259 let clarification = agent.chat("Send it").await.unwrap();
14260 assert_eq!(clarification.content, "What should I send?");
14261 assert_eq!(confirmation_skill_call_count(&observed), 0);
14262
14263 let confirmation = agent.chat("The report to Ada").await.unwrap();
14264 assert_eq!(confirmation.content, "Should I send the report to Ada?");
14265 assert_eq!(
14266 confirmation
14267 .metadata
14268 .as_ref()
14269 .and_then(|metadata| metadata.get("disambiguation"))
14270 .and_then(|metadata| metadata.get("status"))
14271 .and_then(Value::as_str),
14272 Some("awaiting_confirmation")
14273 );
14274 assert_eq!(confirmation_skill_call_count(&observed), 0);
14275
14276 let completed = agent.chat("Yes").await.unwrap();
14277 assert_eq!(completed.content, "Report skill executed.");
14278 assert_eq!(confirmation_skill_call_count(&observed), 1);
14279 assert!(agent.pending_skill_id.read().is_none());
14280 let messages = agent.memory.get_messages(None).await.unwrap();
14281 assert!(!messages.iter().any(|message| message.content == "Yes"));
14282 }
14283
14284 #[tokio::test]
14286 async fn confirmed_skill_recheck_preserves_new_clarification_metadata() {
14287 let (agent, observed) = state_disambiguation_agent_with_skills(
14288 vec![
14289 r#"{"is_ambiguous":false,"confidence":0.99,"ambiguity_type":null,"reasoning":"top-level clear","what_is_unclear":[],"detected_language":"en"}"#,
14290 "send_report",
14291 r#"{"is_ambiguous":true,"confidence":0.2,"ambiguity_type":"missing_target","reasoning":"target missing","what_is_unclear":["target"],"detected_language":"en"}"#,
14292 r#"{"question":"What should I send?","options":null}"#,
14293 r#"{"status":"answered","selected_option":null,"enriched_input":"Send the report to Ada","resolved":{"intent":"send_report"}}"#,
14294 r#"{"question":"Should I send the report to Ada?"}"#,
14295 r#"{"status":"confirmed"}"#,
14296 r#"{"is_ambiguous":true,"confidence":0.3,"ambiguity_type":"missing_parameters","reasoning":"timing missing","what_is_unclear":["timing"],"detected_language":"en"}"#,
14297 r#"{"question":"When should I send it?","options":null}"#,
14298 ],
14299 true,
14300 None,
14301 true,
14302 vec![confirmation_skill()],
14303 );
14304
14305 agent.chat("Send it").await.unwrap();
14306 agent.chat("The report to Ada").await.unwrap();
14307 let follow_up = agent.chat("Yes").await.unwrap();
14308
14309 assert_eq!(follow_up.content, "When should I send it?");
14310 let metadata = follow_up
14311 .metadata
14312 .as_ref()
14313 .and_then(|metadata| metadata.get("disambiguation"))
14314 .unwrap();
14315 assert_eq!(
14316 metadata.get("status").and_then(Value::as_str),
14317 Some("awaiting_clarification")
14318 );
14319 assert_eq!(
14320 metadata.get("skill_id").and_then(Value::as_str),
14321 Some("send_report")
14322 );
14323 assert!(metadata.get("detection").is_some());
14324 assert_eq!(confirmation_skill_call_count(&observed), 0);
14325 }
14326
14327 #[tokio::test]
14329 async fn rejected_skill_confirmation_never_executes() {
14330 let (agent, observed) = state_disambiguation_agent_with_skills(
14331 vec![
14332 r#"{"is_ambiguous":false,"confidence":0.99,"ambiguity_type":null,"reasoning":"top-level clear","what_is_unclear":[],"detected_language":"en"}"#,
14333 "send_report",
14334 r#"{"is_ambiguous":true,"confidence":0.2,"ambiguity_type":"missing_target","reasoning":"target missing","what_is_unclear":["target"],"detected_language":"en"}"#,
14335 r#"{"question":"What should I send?","options":null}"#,
14336 r#"{"status":"answered","selected_option":null,"enriched_input":"Send the report to Ada","resolved":{"intent":"send_report"}}"#,
14337 r#"{"question":"Should I send the report to Ada?"}"#,
14338 r#"{"status":"rejected"}"#,
14339 "Confirmation rejected.",
14340 ],
14341 true,
14342 None,
14343 true,
14344 vec![confirmation_skill()],
14345 );
14346
14347 agent.chat("Send it").await.unwrap();
14348 agent.chat("The report to Ada").await.unwrap();
14349 let rejected = agent.chat("No").await.unwrap();
14350
14351 assert_eq!(rejected.content, "Confirmation rejected.");
14352 assert_eq!(confirmation_skill_call_count(&observed), 0);
14353 assert!(agent.pending_skill_id.read().is_none());
14354 }
14355
14356 #[tokio::test]
14358 async fn reset_invalidates_pending_skill_confirmation_before_streaming_input() {
14359 let (agent, observed) = state_disambiguation_agent_with_skills(
14360 vec![
14361 r#"{"is_ambiguous":false,"confidence":0.99,"ambiguity_type":null,"reasoning":"top-level clear","what_is_unclear":[],"detected_language":"en"}"#,
14362 "send_report",
14363 r#"{"is_ambiguous":true,"confidence":0.2,"ambiguity_type":"missing_target","reasoning":"target missing","what_is_unclear":["target"],"detected_language":"en"}"#,
14364 r#"{"question":"What should I send?","options":null}"#,
14365 r#"{"status":"answered","selected_option":null,"enriched_input":"Send the report to Ada","resolved":{"intent":"send_report"}}"#,
14366 r#"{"question":"Should I send the report to Ada?"}"#,
14367 r#"{"is_ambiguous":false,"confidence":0.99,"ambiguity_type":null,"reasoning":"fresh input","what_is_unclear":[],"detected_language":"en"}"#,
14368 "none",
14369 "Fresh response.",
14370 ],
14371 true,
14372 None,
14373 true,
14374 vec![confirmation_skill()],
14375 );
14376
14377 agent.chat("Send it").await.unwrap();
14378 agent.chat("The report to Ada").await.unwrap();
14379 agent.reset().await.unwrap();
14380 assert!(agent.pending_skill_id.read().is_none());
14381 assert!(
14382 !agent
14383 .disambiguation_manager()
14384 .unwrap()
14385 .has_pending_clarification()
14386 .await
14387 );
14388
14389 let mut stream = agent.chat_stream("Yes").await.unwrap();
14390 let mut content = String::new();
14391 while let Some(chunk) = stream.next().await {
14392 match chunk {
14393 StreamChunk::Content { text } => content.push_str(&text),
14394 StreamChunk::Done {} => break,
14395 StreamChunk::Error { message } => panic!("unexpected stream error: {message}"),
14396 _ => {}
14397 }
14398 }
14399
14400 assert_eq!(content, "Fresh response.");
14401 assert_eq!(confirmation_skill_call_count(&observed), 0);
14402 }
14403
14404 #[tokio::test]
14406 async fn trait_reset_clears_pending_skill_confirmation() {
14407 let (agent, _) = state_disambiguation_agent_with_skills(
14408 vec![
14409 r#"{"is_ambiguous":false,"confidence":0.99,"ambiguity_type":null,"reasoning":"top-level clear","what_is_unclear":[],"detected_language":"en"}"#,
14410 "send_report",
14411 r#"{"is_ambiguous":true,"confidence":0.2,"ambiguity_type":"missing_target","reasoning":"target missing","what_is_unclear":["target"],"detected_language":"en"}"#,
14412 r#"{"question":"What should I send?","options":null}"#,
14413 r#"{"status":"answered","selected_option":null,"enriched_input":"Send the report to Ada","resolved":{"intent":"send_report"}}"#,
14414 r#"{"question":"Should I send the report to Ada?"}"#,
14415 ],
14416 true,
14417 None,
14418 true,
14419 vec![confirmation_skill()],
14420 );
14421
14422 agent.chat("Send it").await.unwrap();
14423 agent.chat("The report to Ada").await.unwrap();
14424 <RuntimeAgent as Agent>::reset(&agent).await.unwrap();
14425
14426 assert!(agent.pending_skill_id.read().is_none());
14427 assert!(
14428 !agent
14429 .disambiguation_manager()
14430 .unwrap()
14431 .has_pending_clarification()
14432 .await
14433 );
14434 }
14435
14436 #[tokio::test]
14438 async fn state_change_invalidates_pending_skill_confirmation() {
14439 let (agent, observed) = state_disambiguation_agent_with_skills(
14440 vec![
14441 r#"{"is_ambiguous":false,"confidence":0.99,"ambiguity_type":null,"reasoning":"top-level clear","what_is_unclear":[],"detected_language":"en"}"#,
14442 "send_report",
14443 r#"{"is_ambiguous":true,"confidence":0.2,"ambiguity_type":"missing_target","reasoning":"target missing","what_is_unclear":["target"],"detected_language":"en"}"#,
14444 r#"{"question":"What should I send?","options":null}"#,
14445 r#"{"status":"answered","selected_option":null,"enriched_input":"Send the report to Ada","resolved":{"intent":"send_report"}}"#,
14446 r#"{"question":"Should I send the report to Ada?"}"#,
14447 r#"{"is_ambiguous":false,"confidence":0.99,"ambiguity_type":null,"reasoning":"fresh input","what_is_unclear":[],"detected_language":"en"}"#,
14448 "none",
14449 "Fresh response.",
14450 ],
14451 true,
14452 None,
14453 true,
14454 vec![confirmation_skill()],
14455 );
14456
14457 agent.chat("Send it").await.unwrap();
14458 agent.chat("The report to Ada").await.unwrap();
14459 agent.transition_to("review").await.unwrap();
14460 let cancelled = agent.chat("Yes").await.unwrap();
14461
14462 assert_eq!(cancelled.content, "Fresh response.");
14463 assert_eq!(confirmation_skill_call_count(&observed), 0);
14464 assert!(agent.pending_skill_id.read().is_none());
14465 }
14466
14467 #[tokio::test]
14469 async fn in_flight_confirmation_cannot_redispatch_after_reset() {
14470 let (mut agent, observed) = state_disambiguation_agent_with_skills(
14471 vec![
14472 r#"{"is_ambiguous":false,"confidence":0.99,"ambiguity_type":null,"reasoning":"top-level clear","what_is_unclear":[],"detected_language":"en"}"#,
14473 "send_report",
14474 r#"{"is_ambiguous":true,"confidence":0.2,"ambiguity_type":"missing_target","reasoning":"target missing","what_is_unclear":["target"],"detected_language":"en"}"#,
14475 r#"{"question":"What should I send?","options":null}"#,
14476 r#"{"status":"answered","selected_option":null,"enriched_input":"Send the report to Ada","resolved":{"intent":"send_report"}}"#,
14477 r#"{"question":"Should I send the report to Ada?"}"#,
14478 r#"{"status":"confirmed"}"#,
14479 "Confirmation cancelled.",
14480 ],
14481 true,
14482 None,
14483 true,
14484 vec![confirmation_skill()],
14485 );
14486 let observer = Arc::new(BlockingRuntimeConfirmationObserver::new());
14487 let manager = agent
14488 .disambiguation_manager
14489 .take()
14490 .unwrap()
14491 .with_clarification_observer(observer.clone());
14492 agent.disambiguation_manager = Some(manager);
14493 let agent = Arc::new(agent);
14494
14495 agent.chat("Send it").await.unwrap();
14496 agent.chat("The report to Ada").await.unwrap();
14497
14498 let confirming_agent = Arc::clone(&agent);
14499 let confirmation = tokio::spawn(async move { confirming_agent.chat("Yes").await });
14500 observer.entered.wait().await;
14501 agent.reset().await.unwrap();
14502 observer.release.notify_one();
14503
14504 let response = confirmation.await.unwrap().unwrap();
14505 assert_eq!(response.content, "Confirmation cancelled.");
14506 assert_eq!(confirmation_skill_call_count(&observed), 0);
14507 assert!(agent.pending_skill_id.read().is_none());
14508 }
14509
14510 #[tokio::test]
14512 async fn queued_reset_prevents_stale_confirmation_question_publication() {
14513 let (agent, observed) = state_disambiguation_agent(
14514 vec![
14515 r#"{"is_ambiguous":true,"confidence":0.2,"ambiguity_type":"missing_target","reasoning":"target missing","what_is_unclear":["target"],"detected_language":"en"}"#,
14516 r#"{"question":"What should I send?","options":null}"#,
14517 r#"{"status":"answered","selected_option":null,"enriched_input":"Send the report to Ada","resolved":{"intent":"send_report"}}"#,
14518 r#"{"question":"Should I send the report to Ada?"}"#,
14519 ],
14520 true,
14521 None,
14522 true,
14523 );
14524 let agent = Arc::new(agent);
14525 agent.chat("Send it").await.unwrap();
14526
14527 let admission = agent.disambiguation_admission.write().await;
14528 let (started_tx, started_rx) = tokio::sync::oneshot::channel();
14529 let resetting_agent = Arc::clone(&agent);
14530 let reset = tokio::spawn(async move {
14531 let _ = started_tx.send(());
14532 resetting_agent.reset().await
14533 });
14534 started_rx.await.unwrap();
14535 tokio::task::yield_now().await;
14536
14537 let responding_agent = Arc::clone(&agent);
14538 let response =
14539 tokio::spawn(async move { responding_agent.chat("The report to Ada").await });
14540 tokio::time::timeout(std::time::Duration::from_secs(2), async {
14541 while observed.call_count() < 4 {
14542 tokio::task::yield_now().await;
14543 }
14544 })
14545 .await
14546 .expect("clarification processing must reach terminal publication");
14547 drop(admission);
14548
14549 reset.await.unwrap().unwrap();
14550 let error = response.await.unwrap().unwrap_err();
14551 assert!(error.to_string().contains("ownership changed"));
14552 assert!(
14553 !agent
14554 .disambiguation_manager()
14555 .unwrap()
14556 .has_pending_clarification()
14557 .await
14558 );
14559 assert!(agent.memory.get_messages(None).await.unwrap().is_empty());
14560 }
14561
14562 #[tokio::test]
14564 async fn queued_reset_prevents_stale_skill_clarification_publication() {
14565 let (agent, observed) = state_disambiguation_agent_with_skills(
14566 vec![
14567 r#"{"is_ambiguous":false,"confidence":0.99,"ambiguity_type":null,"reasoning":"top-level clear","what_is_unclear":[],"detected_language":"en"}"#,
14568 "send_report",
14569 r#"{"is_ambiguous":true,"confidence":0.2,"ambiguity_type":"missing_target","reasoning":"target missing","what_is_unclear":["target"],"detected_language":"en"}"#,
14570 r#"{"question":"What should I send?","options":null}"#,
14571 ],
14572 true,
14573 None,
14574 true,
14575 vec![confirmation_skill()],
14576 );
14577 let agent = Arc::new(agent);
14578 let admission = agent.disambiguation_admission.write().await;
14579 let (started_tx, started_rx) = tokio::sync::oneshot::channel();
14580 let resetting_agent = Arc::clone(&agent);
14581 let reset = tokio::spawn(async move {
14582 let _ = started_tx.send(());
14583 resetting_agent.reset().await
14584 });
14585 started_rx.await.unwrap();
14586 tokio::task::yield_now().await;
14587
14588 let responding_agent = Arc::clone(&agent);
14589 let response = tokio::spawn(async move { responding_agent.chat("Send it").await });
14590 tokio::time::timeout(std::time::Duration::from_secs(2), async {
14591 while observed.call_count() < 4 {
14592 tokio::task::yield_now().await;
14593 }
14594 })
14595 .await
14596 .expect("skill clarification must reach terminal publication");
14597 drop(admission);
14598
14599 reset.await.unwrap().unwrap();
14600 let error = response.await.unwrap().unwrap_err();
14601 assert!(error.to_string().contains("ownership changed"));
14602 assert_eq!(confirmation_skill_call_count(&observed), 0);
14603 assert!(agent.pending_skill_id.read().is_none());
14604 assert!(agent.memory.get_messages(None).await.unwrap().is_empty());
14605 }
14606
14607 #[tokio::test]
14609 async fn transition_hook_can_reset_without_admission_deadlock() {
14610 let hooks = Arc::new(ResetOnTransitionHooks {
14611 agent: parking_lot::Mutex::new(None),
14612 invoked: AtomicBool::new(false),
14613 });
14614 let agent = Arc::new(
14615 AgentBuilder::new()
14616 .system_prompt("Test transition hook reentrancy.")
14617 .llm(Arc::new(mock_with_response("done")))
14618 .state_machine(disambiguation_state_machine(None, false))
14619 .build()
14620 .unwrap()
14621 .with_hooks(hooks.clone()),
14622 );
14623 *hooks.agent.lock() = Some(Arc::downgrade(&agent));
14624
14625 let transitioned = tokio::time::timeout(
14626 std::time::Duration::from_secs(2),
14627 agent.apply_transition_target("active", "review", "test transition", None),
14628 )
14629 .await
14630 .expect("transition hook reset must not deadlock")
14631 .unwrap();
14632
14633 assert!(transitioned);
14634 assert!(hooks.invoked.load(Ordering::SeqCst));
14635 assert_eq!(agent.current_state().as_deref(), Some("active"));
14636 }
14637
14638 #[tokio::test]
14640 async fn concurrent_transition_cannot_duplicate_exit_actions() {
14641 let gate = PathMutationGate::new();
14642 let active = ai_agents_state::StateDefinition {
14643 on_exit: vec![StateAction::Tool {
14644 tool: "transition_exit".to_string(),
14645 args: Some(serde_json::json!({"path": "./transition-exit.txt"})),
14646 }],
14647 ..Default::default()
14648 };
14649 let state_machine = Arc::new(
14650 StateMachine::new(ai_agents_state::StateConfig {
14651 initial: "active".to_string(),
14652 states: HashMap::from([
14653 ("active".to_string(), active),
14654 (
14655 "review".to_string(),
14656 ai_agents_state::StateDefinition::default(),
14657 ),
14658 ]),
14659 global_transitions: Vec::new(),
14660 fallback: None,
14661 max_no_transition: None,
14662 regenerate_on_transition: true,
14663 })
14664 .unwrap(),
14665 );
14666 let agent = Arc::new(
14667 AgentBuilder::new()
14668 .system_prompt("Test transition reservation.")
14669 .llm(Arc::new(mock_with_response("done")))
14670 .tool(Arc::new(BlockingPathMutationTool {
14671 id: "transition_exit",
14672 path_fields: vec![ai_agents_core::PathPolicyBinding::write("path")],
14673 gate: gate.clone(),
14674 }))
14675 .state_machine(state_machine)
14676 .build()
14677 .unwrap(),
14678 );
14679
14680 let first_agent = Arc::clone(&agent);
14681 let first = tokio::spawn(async move { first_agent.transition_to("review").await });
14682 tokio::time::timeout(std::time::Duration::from_secs(2), gate.wait_until_entered())
14683 .await
14684 .expect("reserved transition must enter its exit action");
14685
14686 let second = tokio::time::timeout(
14687 std::time::Duration::from_secs(2),
14688 agent.transition_to("review"),
14689 )
14690 .await
14691 .expect("competing transition must fail without waiting for the exit action")
14692 .unwrap_err();
14693 assert!(second.to_string().contains("already in progress"));
14694
14695 gate.release();
14696 first.await.unwrap().unwrap();
14697 assert_eq!(agent.current_state().as_deref(), Some("review"));
14698 }
14699
14700 #[tokio::test]
14702 async fn concurrent_transition_cannot_overtake_enter_actions() {
14703 let gate = PathMutationGate::new();
14704 let review = ai_agents_state::StateDefinition {
14705 on_enter: vec![StateAction::Tool {
14706 tool: "transition_enter".to_string(),
14707 args: Some(serde_json::json!({"path": "./transition-enter.txt"})),
14708 }],
14709 ..Default::default()
14710 };
14711 let state_machine = Arc::new(
14712 StateMachine::new(ai_agents_state::StateConfig {
14713 initial: "active".to_string(),
14714 states: HashMap::from([
14715 (
14716 "active".to_string(),
14717 ai_agents_state::StateDefinition::default(),
14718 ),
14719 ("review".to_string(), review),
14720 ]),
14721 global_transitions: Vec::new(),
14722 fallback: None,
14723 max_no_transition: None,
14724 regenerate_on_transition: true,
14725 })
14726 .unwrap(),
14727 );
14728 let agent = Arc::new(
14729 AgentBuilder::new()
14730 .system_prompt("Test transition lifecycle reservation.")
14731 .llm(Arc::new(mock_with_response("done")))
14732 .tool(Arc::new(BlockingPathMutationTool {
14733 id: "transition_enter",
14734 path_fields: vec![ai_agents_core::PathPolicyBinding::write("path")],
14735 gate: gate.clone(),
14736 }))
14737 .state_machine(state_machine)
14738 .build()
14739 .unwrap(),
14740 );
14741
14742 let first_agent = Arc::clone(&agent);
14743 let first = tokio::spawn(async move { first_agent.transition_to("review").await });
14744 tokio::time::timeout(std::time::Duration::from_secs(2), gate.wait_until_entered())
14745 .await
14746 .expect("committed transition must enter its destination action");
14747
14748 let second = agent.transition_to("active").await.unwrap_err();
14749 assert!(second.to_string().contains("already in progress"));
14750 assert!(agent.reset().await.is_err());
14751
14752 gate.release();
14753 first.await.unwrap().unwrap();
14754 assert_eq!(agent.current_state().as_deref(), Some("review"));
14755 }
14756
14757 #[tokio::test]
14759 async fn same_state_restore_invalidates_pending_skill_confirmation() {
14760 let (agent, observed) = state_disambiguation_agent_with_skills(
14761 vec![
14762 r#"{"is_ambiguous":false,"confidence":0.99,"ambiguity_type":null,"reasoning":"top-level clear","what_is_unclear":[],"detected_language":"en"}"#,
14763 "send_report",
14764 r#"{"is_ambiguous":true,"confidence":0.2,"ambiguity_type":"missing_target","reasoning":"target missing","what_is_unclear":["target"],"detected_language":"en"}"#,
14765 r#"{"question":"What should I send?","options":null}"#,
14766 r#"{"status":"answered","selected_option":null,"enriched_input":"Send the report to Ada","resolved":{"intent":"send_report"}}"#,
14767 r#"{"question":"Should I send the report to Ada?"}"#,
14768 ],
14769 true,
14770 None,
14771 true,
14772 vec![confirmation_skill()],
14773 );
14774
14775 agent.chat("Send it").await.unwrap();
14776 agent.chat("The report to Ada").await.unwrap();
14777 let snapshot = agent.save_state().await.unwrap();
14778 assert_eq!(agent.current_state().as_deref(), Some("active"));
14779
14780 agent.restore_state(snapshot).await.unwrap();
14781
14782 assert_eq!(agent.current_state().as_deref(), Some("active"));
14783 assert!(agent.pending_skill_id.read().is_none());
14784 assert!(
14785 !agent
14786 .disambiguation_manager()
14787 .unwrap()
14788 .has_pending_clarification()
14789 .await
14790 );
14791 assert_eq!(confirmation_skill_call_count(&observed), 0);
14792 }
14793
14794 #[tokio::test]
14796 async fn direct_state_generation_change_invalidates_confirmation() {
14797 let (agent, observed) = state_disambiguation_agent_with_skills(
14798 vec![
14799 r#"{"is_ambiguous":false,"confidence":0.99,"ambiguity_type":null,"reasoning":"top-level clear","what_is_unclear":[],"detected_language":"en"}"#,
14800 "send_report",
14801 r#"{"is_ambiguous":true,"confidence":0.2,"ambiguity_type":"missing_target","reasoning":"target missing","what_is_unclear":["target"],"detected_language":"en"}"#,
14802 r#"{"question":"What should I send?","options":null}"#,
14803 r#"{"status":"answered","selected_option":null,"enriched_input":"Send the report to Ada","resolved":{"intent":"send_report"}}"#,
14804 r#"{"question":"Should I send the report to Ada?"}"#,
14805 "Confirmation cancelled.",
14806 ],
14807 true,
14808 None,
14809 true,
14810 vec![confirmation_skill()],
14811 );
14812
14813 agent.chat("Send it").await.unwrap();
14814 agent.chat("The report to Ada").await.unwrap();
14815 let state_machine = agent.state_machine().unwrap();
14816 state_machine
14817 .transition_to("review", "external test")
14818 .unwrap();
14819 state_machine
14820 .transition_to("active", "external test")
14821 .unwrap();
14822
14823 let response = agent.chat("Yes").await.unwrap();
14824
14825 assert_eq!(response.content, "Confirmation cancelled.");
14826 assert_eq!(confirmation_skill_call_count(&observed), 0);
14827 assert!(agent.pending_skill_id.read().is_none());
14828 }
14829
14830 #[tokio::test]
14831 async fn state_confirmation_does_not_add_a_question_for_clear_input() {
14832 let (agent, observed) = state_disambiguation_agent(
14833 vec![
14834 r#"{"is_ambiguous":false,"confidence":0.99,"ambiguity_type":null,"reasoning":"clear","what_is_unclear":[],"detected_language":"en"}"#,
14835 "Request executed.",
14836 ],
14837 true,
14838 None,
14839 true,
14840 );
14841
14842 let response = agent.chat("Send the report to Ada").await.unwrap();
14843
14844 assert_eq!(response.content, "Request executed.");
14845 assert_eq!(observed.call_count(), 2);
14846 }
14847
14848 #[tokio::test]
14849 async fn state_override_cannot_activate_a_disabled_top_level_manager() {
14850 let (agent, observed) =
14851 state_disambiguation_agent(vec!["Request executed."], false, Some(true), true);
14852
14853 assert!(!agent.has_disambiguation());
14854 let response = agent.chat("Send it").await.unwrap();
14855
14856 assert_eq!(response.content, "Request executed.");
14857 assert_eq!(observed.call_count(), 1);
14858 }
14859
14860 #[tokio::test]
14861 async fn native_required_choice_executes_through_the_shared_tool_path() {
14862 let mut mock = MockLLMProvider::new("native-required");
14863 mock.set_tool_choice(Some(ToolChoice::Required));
14864 let native_call = ToolCall {
14865 id: "provider-call-1".to_string(),
14866 name: "calculator".to_string(),
14867 arguments: serde_json::json!({"expression": "2 + 2"}),
14868 };
14869 let provider_state = ai_agents_core::NativeProviderState::new(
14870 "fixture-exchange-1",
14871 "fixture",
14872 "native-tools",
14873 ai_agents_core::NativeProviderTarget::new("https://fixture.invalid/", "fixture-model")
14874 .unwrap(),
14875 serde_json::json!({
14876 "role": "model",
14877 "parts": [{
14878 "functionCall": {"name": "calculator", "args": {"expression": "2 + 2"}},
14879 "thoughtSignature": "fixture-signature"
14880 }]
14881 }),
14882 vec![ai_agents_core::NativeCallBinding::new("provider-call-1", 0).unwrap()],
14883 )
14884 .unwrap();
14885 mock.add_response(
14886 LLMResponse::new("", FinishReason::ToolCall)
14887 .with_provider_state(provider_state)
14888 .unwrap()
14889 .with_tool_calls(vec![native_call])
14890 .unwrap(),
14891 );
14892 mock.add_response(LLMResponse::new("The answer is 4.", FinishReason::Stop));
14893 let observed = mock.clone();
14894 let agent = AgentBuilder::new()
14895 .system_prompt("Use the calculator when needed.")
14896 .llm(Arc::new(mock))
14897 .tool(Arc::new(CalculatorTool::new()))
14898 .build()
14899 .unwrap();
14900
14901 let response = agent.chat("What is 2 + 2?").await.unwrap();
14902
14903 assert_eq!(response.content, "The answer is 4.");
14904 assert_eq!(
14905 response.tool_calls.as_ref().unwrap()[0].id,
14906 "provider-call-1"
14907 );
14908 let calls = observed.call_history();
14909 assert_eq!(calls.len(), 2);
14910 assert!(matches!(
14911 calls[0].request.as_ref().map(|request| &request.choice),
14912 Some(ToolChoice::Required)
14913 ));
14914 assert!(matches!(
14915 calls[1].request.as_ref().map(|request| &request.choice),
14916 Some(ToolChoice::Auto)
14917 ));
14918 let replay_batch = calls[1]
14919 .messages
14920 .iter()
14921 .find_map(|message| {
14922 ai_agents_core::decode_native_tool_call_markers(&message.content).unwrap()
14923 })
14924 .expect("signed native call marker must be replayed");
14925 assert_eq!(
14926 replay_batch.provider_state().unwrap().exchange_id(),
14927 "fixture-exchange-1"
14928 );
14929 assert!(calls[1].messages.iter().any(|message| {
14930 ai_agents_core::decode_native_tool_result_markers(&message.content)
14931 .is_ok_and(|results| results.is_some())
14932 }));
14933 }
14934
14935 #[tokio::test]
14936 async fn custom_memory_loss_stops_before_signed_tool_execution() {
14937 let mut mock = MockLLMProvider::new("native-custom-memory");
14938 mock.set_tool_choice(Some(ToolChoice::Required));
14939 let call = ToolCall {
14940 id: "provider-call-drop".to_string(),
14941 name: "calculator".to_string(),
14942 arguments: serde_json::json!({"expression": "3 + 4"}),
14943 };
14944 let state = ai_agents_core::NativeProviderState::new(
14945 "fixture-exchange-drop",
14946 "fixture",
14947 "native-tools",
14948 ai_agents_core::NativeProviderTarget::new("https://fixture.invalid/", "fixture-model")
14949 .unwrap(),
14950 serde_json::json!({
14951 "role": "model",
14952 "parts": [{
14953 "functionCall": {"name": "calculator", "args": {"expression": "3 + 4"}},
14954 "thoughtSignature": "fixture-signature-drop"
14955 }]
14956 }),
14957 vec![ai_agents_core::NativeCallBinding::new("provider-call-drop", 0).unwrap()],
14958 )
14959 .unwrap();
14960 mock.add_response(
14961 LLMResponse::new("", FinishReason::ToolCall)
14962 .with_provider_state(state)
14963 .unwrap()
14964 .with_tool_calls(vec![call])
14965 .unwrap(),
14966 );
14967 let agent = AgentBuilder::new()
14968 .system_prompt("Use the calculator.")
14969 .llm(Arc::new(mock))
14970 .memory(Arc::new(DroppingSignedAssistantMemory {
14971 messages: RwLock::new(Vec::new()),
14972 }))
14973 .tool(Arc::new(CalculatorTool::new()))
14974 .build()
14975 .unwrap();
14976
14977 let error = agent.chat("What is 3 + 4?").await.unwrap_err();
14978
14979 assert!(
14980 error
14981 .to_string()
14982 .contains("removed before provider continuation")
14983 );
14984 assert!(agent.tool_call_history.read().is_empty());
14985 }
14986
14987 #[tokio::test]
14988 async fn sequential_signed_history_validates_every_prior_exchange() {
14989 let mut mock = MockLLMProvider::new("native-sequential-memory");
14990 mock.set_tool_choice(Some(ToolChoice::Required));
14991 mock.add_response(signed_calculator_response(
14992 "seq-exchange-1",
14993 "seq-call-1",
14994 "1 + 1",
14995 ));
14996 mock.add_response(signed_calculator_response(
14997 "seq-exchange-2",
14998 "seq-call-2",
14999 "2 + 2",
15000 ));
15001 let agent = AgentBuilder::new()
15002 .system_prompt("Use the calculator sequentially.")
15003 .llm(Arc::new(mock))
15004 .memory(Arc::new(DroppingEarlierSequentialMemory {
15005 messages: RwLock::new(Vec::new()),
15006 signed_seen: std::sync::atomic::AtomicUsize::new(0),
15007 }))
15008 .tool(Arc::new(CalculatorTool::new()))
15009 .build()
15010 .unwrap();
15011
15012 let error = agent.chat("Calculate twice.").await.unwrap_err();
15013
15014 assert!(error.to_string().contains("seq-exchange-1"));
15015 assert_eq!(agent.tool_call_history.read().len(), 1);
15016 }
15017
15018 #[tokio::test]
15019 async fn post_transition_signed_hitl_rejection_stops_before_continuation() {
15020 let mut native = MockLLMProvider::new("post-transition-native");
15021 native.set_tool_choice(Some(ToolChoice::Auto));
15022 let call = ToolCall {
15023 id: "post-transition-call".to_string(),
15024 name: "echo".to_string(),
15025 arguments: serde_json::json!({"message": "hello"}),
15026 };
15027 let state = ai_agents_core::NativeProviderState::new(
15028 "post-transition-exchange",
15029 "fixture",
15030 "native-tools",
15031 ai_agents_core::NativeProviderTarget::new("https://fixture.invalid/", "fixture-model")
15032 .unwrap(),
15033 serde_json::json!({
15034 "role": "model",
15035 "parts": [{
15036 "functionCall": {"name": "echo", "args": {"message": "hello"}},
15037 "thoughtSignature": "post-transition-signature"
15038 }]
15039 }),
15040 vec![ai_agents_core::NativeCallBinding::new("post-transition-call", 0).unwrap()],
15041 )
15042 .unwrap();
15043 native.add_response(
15044 LLMResponse::new("", FinishReason::ToolCall)
15045 .with_provider_state(state)
15046 .unwrap()
15047 .with_tool_calls(vec![call])
15048 .unwrap(),
15049 );
15050 let observed_native = native.clone();
15051 let yaml = r#"
15052name: PostTransitionNativeReject
15053system_prompt: test
15054tools: [echo]
15055hitl:
15056 tools:
15057 echo:
15058 require_approval: true
15059states:
15060 initial: intake
15061 states:
15062 intake:
15063 prompt: intake
15064 transitions:
15065 - to: active
15066 guard:
15067 context:
15068 route:
15069 eq: active
15070 active:
15071 prompt: active
15072 llm: native
15073"#;
15074 let agent = AgentBuilder::from_yaml(yaml)
15075 .unwrap()
15076 .llm(Arc::new(mock_with_response("stale intake response")))
15077 .llm_alias("native", Arc::new(native))
15078 .auto_configure_features()
15079 .unwrap()
15080 .build()
15081 .unwrap();
15082 agent
15083 .set_context("route", serde_json::json!("active"))
15084 .unwrap();
15085
15086 let error = agent.chat("move to active").await.unwrap_err();
15087
15088 assert!(matches!(error, AgentError::HITLRejected(_)));
15089 assert_eq!(observed_native.call_count(), 1);
15090 }
15091
15092 #[test]
15093 fn runtime_overflow_removes_a_past_signed_user_turn_as_one_prefix() {
15094 let call = ToolCall {
15095 id: "overflow-call".to_string(),
15096 name: "calculator".to_string(),
15097 arguments: serde_json::json!({"expression": "1 + 1"}),
15098 };
15099 let state = ai_agents_core::NativeProviderState::new(
15100 "overflow-exchange",
15101 "google",
15102 "generateContent",
15103 ai_agents_core::NativeProviderTarget::new("https://example.invalid/", "gemini-3")
15104 .unwrap(),
15105 serde_json::json!({
15106 "role": "model",
15107 "parts": [{
15108 "functionCall": {"name": "calculator", "args": {"expression": "1 + 1"}},
15109 "thoughtSignature": "overflow-signature"
15110 }]
15111 }),
15112 vec![ai_agents_core::NativeCallBinding::new("overflow-call", 0).unwrap()],
15113 )
15114 .unwrap();
15115 let call_marker = ai_agents_core::encode_native_tool_call_markers(
15116 std::slice::from_ref(&call),
15117 Some(&state),
15118 )
15119 .unwrap();
15120 let result_marker = ai_agents_core::encode_native_tool_result_marker(
15121 &call,
15122 serde_json::json!({"result": 2}),
15123 )
15124 .unwrap();
15125 let history = vec![
15126 ChatMessage::user("old question"),
15127 ChatMessage::assistant(call_marker),
15128 ChatMessage::function("calculator", result_marker),
15129 ChatMessage::assistant("old answer"),
15130 ChatMessage::user("new question"),
15131 ];
15132
15133 let removable = RuntimeAgent::native_safe_prefix_at_least(&history, 1).unwrap();
15134
15135 assert_eq!(removable, 4);
15136 }
15137
15138 #[test]
15139 fn auxiliary_projection_does_not_interpret_user_marker_text() {
15140 let user_text = serde_json::json!({
15141 "_ai_agents_native_tool_call": true,
15142 "id": "",
15143 "tool": "user-data",
15144 "arguments": {}
15145 })
15146 .to_string();
15147
15148 let projected =
15149 RuntimeAgent::readable_native_messages(vec![ChatMessage::user(&user_text)]).unwrap();
15150
15151 assert_eq!(projected[0].content, user_text);
15152 }
15153
15154 #[tokio::test]
15155 async fn terminal_provider_history_error_skips_retry_and_static_fallback() {
15156 let calls = Arc::new(std::sync::atomic::AtomicU32::new(0));
15157 let recovery = RecoveryManager::new(ai_agents_recovery::ErrorRecoveryConfig {
15158 default: ai_agents_recovery::RetryConfig {
15159 max_retries: 3,
15160 ..Default::default()
15161 },
15162 llm: ai_agents_recovery::LLMRecoveryConfig {
15163 on_failure: LLMFailureAction::FallbackResponse {
15164 message: "must not be returned".to_string(),
15165 },
15166 ..Default::default()
15167 },
15168 ..Default::default()
15169 });
15170 let agent = AgentBuilder::new()
15171 .system_prompt("Reject corrupted native history.")
15172 .llm(Arc::new(TerminalHistoryProvider {
15173 calls: Arc::clone(&calls),
15174 }))
15175 .recovery_manager(recovery)
15176 .build()
15177 .unwrap();
15178
15179 let error = agent.chat("continue").await.unwrap_err();
15180
15181 assert!(
15182 error
15183 .to_string()
15184 .contains("native history integrity failure")
15185 );
15186 assert_eq!(calls.load(Ordering::SeqCst), 1);
15187 }
15188
15189 #[tokio::test]
15190 async fn prompt_fallback_uses_one_corrective_retry() {
15191 let mut mock = MockLLMProvider::new("prompt-required");
15192 mock.set_tool_choice(Some(ToolChoice::Required));
15193 mock.set_native_tool_support(false);
15194 mock.set_responses(
15195 vec![
15196 "I can calculate that.".to_string(),
15197 r#"{"tool":"calculator","arguments":{"expression":"2 + 2"}}"#.to_string(),
15198 "The answer is 4.".to_string(),
15199 ],
15200 false,
15201 );
15202 let observed = mock.clone();
15203 let agent = AgentBuilder::new()
15204 .system_prompt("Use tools.")
15205 .llm(Arc::new(mock))
15206 .tool(Arc::new(CalculatorTool::new()))
15207 .build()
15208 .unwrap();
15209
15210 let response = agent.chat("What is 2 + 2?").await.unwrap();
15211
15212 assert_eq!(response.content, "The answer is 4.");
15213 assert_eq!(observed.call_count(), 3);
15214 let corrective = &observed.call_history()[1].messages;
15215 assert!(
15216 corrective
15217 .last()
15218 .unwrap()
15219 .content
15220 .contains("previous response")
15221 );
15222 }
15223
15224 #[tokio::test]
15225 async fn prompt_fallback_fails_after_one_noncompliant_retry() {
15226 let mut mock = MockLLMProvider::new("prompt-required-failure");
15227 mock.set_tool_choice(Some(ToolChoice::Required));
15228 mock.set_native_tool_support(false);
15229 mock.set_responses(
15230 vec!["No tool.".to_string(), "Still no tool.".to_string()],
15231 false,
15232 );
15233 let observed = mock.clone();
15234 let agent = AgentBuilder::new()
15235 .system_prompt("Use tools.")
15236 .llm(Arc::new(mock))
15237 .tool(Arc::new(CalculatorTool::new()))
15238 .build()
15239 .unwrap();
15240
15241 let error = agent.chat("What is 2 + 2?").await.unwrap_err();
15242
15243 assert!(error.to_string().contains("one corrective retry"));
15244 assert_eq!(observed.call_count(), 2);
15245 }
15246
15247 #[tokio::test]
15248 async fn specific_choice_cannot_widen_the_effective_grant() {
15249 let mut mock = MockLLMProvider::new("specific-outside-grant");
15250 mock.set_tool_choice(Some(ToolChoice::Specific("random".to_string())));
15251 let observed = mock.clone();
15252 let agent = AgentBuilder::new()
15253 .system_prompt("Use tools.")
15254 .llm(Arc::new(mock))
15255 .tool(Arc::new(CalculatorTool::new()))
15256 .build()
15257 .unwrap();
15258
15259 let error = agent.chat("Generate a value.").await.unwrap_err();
15260
15261 assert!(error.to_string().contains("is not registered"));
15262 assert_eq!(observed.call_count(), 0);
15263 }
15264
15265 #[tokio::test]
15266 async fn none_choice_exposes_no_tool_protocol() {
15267 let mut mock = MockLLMProvider::new("no-tools");
15268 mock.set_tool_choice(Some(ToolChoice::None));
15269 mock.set_response(r#"{"tool":"calculator","arguments":{"expression":"2 + 2"}}"#);
15270 let observed = mock.clone();
15271 let agent = AgentBuilder::new()
15272 .system_prompt("Answer directly.")
15273 .llm(Arc::new(mock))
15274 .tool(Arc::new(CalculatorTool::new()))
15275 .build()
15276 .unwrap();
15277
15278 let response = agent.chat("Hello").await.unwrap();
15279
15280 assert!(response.tool_calls.is_none());
15281 assert_eq!(observed.call_count(), 1);
15282 let call = observed.last_call().unwrap();
15283 assert!(call.request.is_none());
15284 assert!(
15285 call.messages
15286 .iter()
15287 .all(|message| !message.content.contains("Available tools:"))
15288 );
15289 }
15290
15291 struct RuntimeStorage {
15292 capabilities: Box<[StorageCapability]>,
15293 snapshots: RwLock<HashMap<String, AgentSnapshot>>,
15294 metadata: RwLock<HashMap<String, ai_agents_core::SessionMetadata>>,
15295 metadata_save_calls: AtomicU64,
15296 metadata_load_calls: AtomicU64,
15297 fail_metadata_save: AtomicBool,
15298 fail_metadata_load: AtomicBool,
15299 }
15300
15301 impl RuntimeStorage {
15302 fn new(capabilities: impl IntoIterator<Item = StorageCapability>) -> Self {
15303 Self {
15304 capabilities: capabilities.into_iter().collect(),
15305 snapshots: RwLock::new(HashMap::new()),
15306 metadata: RwLock::new(HashMap::new()),
15307 metadata_save_calls: AtomicU64::new(0),
15308 metadata_load_calls: AtomicU64::new(0),
15309 fail_metadata_save: AtomicBool::new(false),
15310 fail_metadata_load: AtomicBool::new(false),
15311 }
15312 }
15313 }
15314
15315 #[async_trait]
15316 impl AgentStorage for RuntimeStorage {
15317 fn supports(&self, capability: StorageCapability) -> bool {
15318 self.capabilities.contains(&capability)
15319 }
15320
15321 async fn save(&self, session_id: &str, snapshot: &AgentSnapshot) -> Result<()> {
15322 self.snapshots
15323 .write()
15324 .insert(session_id.to_string(), snapshot.clone());
15325 Ok(())
15326 }
15327
15328 async fn load(&self, session_id: &str) -> Result<Option<AgentSnapshot>> {
15329 Ok(self.snapshots.read().get(session_id).cloned())
15330 }
15331
15332 async fn delete(&self, session_id: &str) -> Result<()> {
15333 self.snapshots.write().remove(session_id);
15334 Ok(())
15335 }
15336
15337 async fn list_sessions(&self) -> Result<Vec<String>> {
15338 Ok(self.snapshots.read().keys().cloned().collect())
15339 }
15340
15341 async fn save_snapshot_with_metadata(
15342 &self,
15343 session_id: &str,
15344 snapshot: &AgentSnapshot,
15345 metadata: &ai_agents_core::SessionMetadata,
15346 ) -> Result<()> {
15347 self.metadata_save_calls.fetch_add(1, Ordering::SeqCst);
15348 if self.fail_metadata_save.load(Ordering::SeqCst) {
15349 return Err(AgentError::Persistence("metadata save failed".into()));
15350 }
15351 self.snapshots
15352 .write()
15353 .insert(session_id.to_string(), snapshot.clone());
15354 self.metadata
15355 .write()
15356 .insert(session_id.to_string(), metadata.clone());
15357 Ok(())
15358 }
15359
15360 async fn save_metadata(
15361 &self,
15362 session_id: &str,
15363 metadata: &ai_agents_core::SessionMetadata,
15364 ) -> Result<()> {
15365 self.metadata_save_calls.fetch_add(1, Ordering::SeqCst);
15366 if self.fail_metadata_save.load(Ordering::SeqCst) {
15367 return Err(AgentError::Persistence("metadata save failed".into()));
15368 }
15369 self.metadata
15370 .write()
15371 .insert(session_id.to_string(), metadata.clone());
15372 Ok(())
15373 }
15374
15375 async fn load_metadata(
15376 &self,
15377 session_id: &str,
15378 ) -> Result<Option<ai_agents_core::SessionMetadata>> {
15379 self.metadata_load_calls.fetch_add(1, Ordering::SeqCst);
15380 if self.fail_metadata_load.load(Ordering::SeqCst) {
15381 return Err(AgentError::Persistence("metadata load failed".into()));
15382 }
15383 Ok(self.metadata.read().get(session_id).cloned())
15384 }
15385 }
15386
15387 fn runtime_storage_agent() -> RuntimeAgent {
15388 AgentBuilder::new()
15389 .system_prompt("Test runtime storage integration.")
15390 .llm(Arc::new(mock_with_response("done")))
15391 .build()
15392 .unwrap()
15393 }
15394
15395 fn restore_spec(id: &str) -> crate::spec::AgentSpec {
15396 crate::spec::AgentSpec {
15397 name: id.to_string(),
15398 system_prompt: format!("Restore child {id}."),
15399 ..crate::spec::AgentSpec::default()
15400 }
15401 }
15402
15403 fn restore_entry(id: &str) -> ai_agents_core::SpawnedAgentEntry {
15404 ai_agents_core::SpawnedAgentEntry {
15405 id: id.to_string(),
15406 name: id.to_string(),
15407 spec_yaml: serde_yaml::to_string(&restore_spec(id)).unwrap(),
15408 }
15409 }
15410
15411 fn restore_spawner(
15412 storage: Arc<RuntimeStorage>,
15413 max_agents: usize,
15414 ) -> (
15415 Arc<crate::spawner::AgentSpawner>,
15416 Arc<crate::spawner::AgentRegistry>,
15417 ) {
15418 let mut llms = LLMRegistry::new();
15419 llms.register("default", Arc::new(mock_with_response("done")));
15420 (
15421 Arc::new(
15422 crate::spawner::AgentSpawner::new()
15423 .with_shared_llms(llms)
15424 .with_shared_storage(storage)
15425 .with_max_agents(max_agents),
15426 ),
15427 Arc::new(crate::spawner::AgentRegistry::new()),
15428 )
15429 }
15430
15431 async fn save_restore_target(
15432 parent: &RuntimeAgent,
15433 storage: &RuntimeStorage,
15434 session_id: &str,
15435 entries: Vec<ai_agents_core::SpawnedAgentEntry>,
15436 ) {
15437 let mut snapshot = parent.save_state().await.unwrap();
15438 snapshot.spawned_agents = Some(entries);
15439 storage.save(session_id, &snapshot).await.unwrap();
15440 storage
15441 .save_metadata(session_id, &ai_agents_core::SessionMetadata::default())
15442 .await
15443 .unwrap();
15444 }
15445
15446 #[tokio::test]
15447 async fn storage_init_requires_storage_for_actor_facts() {
15448 let facts = ai_agents_facts::FactsConfig {
15449 enabled: true,
15450 ..Default::default()
15451 };
15452 let agent = runtime_storage_agent().with_facts_config(None, Some(facts));
15453
15454 let error = agent.init_storage().await.unwrap_err();
15455 assert!(matches!(
15456 error,
15457 AgentError::Config(message)
15458 if message.contains("actor facts or actor memory")
15459 && message.contains("none is configured or injected")
15460 ));
15461 }
15462
15463 #[tokio::test]
15464 async fn storage_init_validates_actor_facts_capability() {
15465 let storage = Arc::new(RuntimeStorage::new([StorageCapability::Snapshot]));
15466 let actor_memory = ai_agents_facts::ActorMemoryConfig {
15467 enabled: true,
15468 ..Default::default()
15469 };
15470 let agent = runtime_storage_agent()
15471 .with_storage(storage)
15472 .with_facts_config(Some(actor_memory), None);
15473
15474 assert!(matches!(
15475 agent.init_storage().await,
15476 Err(AgentError::UnsupportedStorageCapability(
15477 StorageCapability::ActorFacts
15478 ))
15479 ));
15480 }
15481
15482 #[tokio::test]
15483 async fn blocking_chat_rejects_unsupported_required_storage() {
15484 let storage = Arc::new(RuntimeStorage::new([StorageCapability::Snapshot]));
15485 let facts = ai_agents_facts::FactsConfig {
15486 enabled: true,
15487 ..Default::default()
15488 };
15489 let agent = runtime_storage_agent()
15490 .with_storage(storage)
15491 .with_facts_config(None, Some(facts));
15492
15493 assert!(matches!(
15494 agent.chat("hello").await,
15495 Err(AgentError::UnsupportedStorageCapability(
15496 StorageCapability::ActorFacts
15497 ))
15498 ));
15499 }
15500
15501 #[tokio::test]
15502 async fn streaming_chat_rejects_unsupported_required_storage_before_stream_creation() {
15503 let storage = Arc::new(RuntimeStorage::new([StorageCapability::Snapshot]));
15504 let config = ai_agents_relationships::RelationshipConfig {
15505 enabled: true,
15506 ..Default::default()
15507 };
15508 let manager = Arc::new(RelationshipManager::from_config(config).unwrap());
15509 let agent = runtime_storage_agent()
15510 .with_storage(storage)
15511 .with_relationships(manager);
15512
15513 assert!(matches!(
15514 agent.chat_stream("hello").await,
15515 Err(AgentError::UnsupportedStorageCapability(
15516 StorageCapability::ActorRelationships
15517 ))
15518 ));
15519 }
15520
15521 #[tokio::test]
15522 async fn storage_init_completes_facts_for_injected_storage() {
15523 let storage = Arc::new(RuntimeStorage::new([
15524 StorageCapability::Snapshot,
15525 StorageCapability::ActorFacts,
15526 ]));
15527 let facts = ai_agents_facts::FactsConfig {
15528 enabled: true,
15529 ..Default::default()
15530 };
15531 let agent = runtime_storage_agent()
15532 .with_storage(storage)
15533 .with_facts_config(None, Some(facts));
15534
15535 agent.init_storage().await.unwrap();
15536 assert!(agent.fact_store().is_some());
15537 }
15538
15539 #[tokio::test]
15540 async fn storage_init_requires_storage_for_persistent_relationships() {
15541 let config = ai_agents_relationships::RelationshipConfig {
15542 enabled: true,
15543 ..Default::default()
15544 };
15545 let manager = Arc::new(RelationshipManager::from_config(config).unwrap());
15546 let agent = runtime_storage_agent().with_relationships(manager);
15547
15548 let error = agent.init_storage().await.unwrap_err();
15549 assert!(matches!(
15550 error,
15551 AgentError::Config(message)
15552 if message.contains("persistent relationships")
15553 && message.contains("none is configured or injected")
15554 ));
15555 }
15556
15557 #[tokio::test]
15558 async fn storage_init_validates_persistent_relationships_capability() {
15559 let storage = Arc::new(RuntimeStorage::new([StorageCapability::Snapshot]));
15560 let config = ai_agents_relationships::RelationshipConfig {
15561 enabled: true,
15562 ..Default::default()
15563 };
15564 let manager = Arc::new(RelationshipManager::from_config(config).unwrap());
15565 let agent = runtime_storage_agent()
15566 .with_storage(storage)
15567 .with_relationships(manager);
15568
15569 assert!(matches!(
15570 agent.init_storage().await,
15571 Err(AgentError::UnsupportedStorageCapability(
15572 StorageCapability::ActorRelationships
15573 ))
15574 ));
15575 }
15576
15577 #[tokio::test]
15578 async fn session_restore_updates_identity_and_clears_stale_actor_binding() {
15579 let storage = Arc::new(RuntimeStorage::new([
15580 StorageCapability::Snapshot,
15581 StorageCapability::SessionMetadata,
15582 ]));
15583 let agent = runtime_storage_agent().with_storage(storage.clone());
15584 agent.set_actor_id("old-actor").unwrap();
15585 agent.save_session("old").await.unwrap();
15586 storage
15587 .save("target", &agent.save_state().await.unwrap())
15588 .await
15589 .unwrap();
15590 storage
15591 .save_metadata("target", &ai_agents_core::SessionMetadata::default())
15592 .await
15593 .unwrap();
15594
15595 assert!(agent.load_session("target").await.unwrap());
15596
15597 assert_eq!(agent.current_session_id.read().as_deref(), Some("target"));
15598 assert_eq!(agent.actor_id(), None);
15599 }
15600
15601 #[tokio::test]
15602 async fn complete_restore_reconciles_growth_shrink_and_empty_topologies() {
15603 let storage = Arc::new(RuntimeStorage::new([
15604 StorageCapability::Snapshot,
15605 StorageCapability::SessionMetadata,
15606 ]));
15607 let (spawner, registry) = restore_spawner(storage.clone(), 3);
15608 let parent = runtime_storage_agent()
15609 .with_storage(storage.clone())
15610 .with_spawner_handles(Arc::clone(&spawner), Arc::clone(®istry));
15611
15612 for id in ["a", "b"] {
15613 let spawned = spawner
15614 .spawn_with_id(id.to_string(), restore_spec(id))
15615 .await
15616 .unwrap();
15617 spawned.agent.save_session("grow").await.unwrap();
15618 registry.register(spawned).await.unwrap();
15619 }
15620 let staged_c = crate::spawner::storage::NamespacedStorage::new(storage.clone(), "c");
15621 staged_c
15622 .save("grow", &AgentSnapshot::new("c".into()))
15623 .await
15624 .unwrap();
15625 staged_c
15626 .save_metadata("grow", &ai_agents_core::SessionMetadata::default())
15627 .await
15628 .unwrap();
15629 save_restore_target(
15630 &parent,
15631 storage.as_ref(),
15632 "grow",
15633 vec![restore_entry("a"), restore_entry("b"), restore_entry("c")],
15634 )
15635 .await;
15636
15637 assert_eq!(parent.restore_session_full("grow").await.unwrap(), 3);
15638 assert_eq!(registry.count(), 3);
15639 assert!(registry.contains("c"));
15640 assert_eq!(spawner.spawned_count(), 3);
15641
15642 for id in ["a", "b"] {
15643 registry
15644 .get(id)
15645 .unwrap()
15646 .save_session("shrink")
15647 .await
15648 .unwrap();
15649 }
15650 save_restore_target(
15651 &parent,
15652 storage.as_ref(),
15653 "shrink",
15654 vec![restore_entry("a"), restore_entry("b")],
15655 )
15656 .await;
15657
15658 assert_eq!(parent.restore_session_full("shrink").await.unwrap(), 2);
15659 assert_eq!(registry.count(), 2);
15660 assert!(!registry.contains("c"));
15661 assert_eq!(spawner.spawned_count(), 2);
15662
15663 save_restore_target(&parent, storage.as_ref(), "empty", Vec::new()).await;
15664
15665 assert_eq!(parent.restore_session_full("empty").await.unwrap(), 0);
15666 assert_eq!(registry.count(), 0);
15667 assert_eq!(spawner.spawned_count(), 0);
15668 assert_eq!(parent.current_session_id.read().as_deref(), Some("empty"));
15669 }
15670
15671 #[tokio::test]
15672 async fn storage_session_metadata_is_called_only_when_advertised() {
15673 let storage = Arc::new(RuntimeStorage::new([StorageCapability::Snapshot]));
15674 storage.fail_metadata_save.store(true, Ordering::SeqCst);
15675 storage.fail_metadata_load.store(true, Ordering::SeqCst);
15676 let agent = runtime_storage_agent().with_storage(storage.clone());
15677
15678 agent.save_session("session").await.unwrap();
15679 assert!(agent.load_session("session").await.unwrap());
15680 assert_eq!(storage.metadata_save_calls.load(Ordering::SeqCst), 0);
15681 assert_eq!(storage.metadata_load_calls.load(Ordering::SeqCst), 0);
15682 }
15683
15684 #[cfg(feature = "sqlite")]
15685 #[tokio::test]
15686 async fn sqlite_runtime_save_filter_reopen_and_reload_stay_consistent() {
15687 let directory =
15688 std::env::temp_dir().join(format!("ai-agents-runtime-sqlite-{}", uuid::Uuid::new_v4()));
15689 let path = directory.join("sessions.sqlite");
15690 let path_string = path.to_string_lossy().into_owned();
15691 let storage = Arc::new(
15692 ai_agents_storage::SqliteStorage::new(&path_string)
15693 .await
15694 .unwrap(),
15695 );
15696 let agent = runtime_storage_agent().with_storage(storage.clone());
15697 agent.set_session_metadata(ai_agents_core::SessionMetadata {
15698 tags: vec!["initial".into()],
15699 ..Default::default()
15700 });
15701 agent.chat("persist this turn").await.unwrap();
15702 agent.save_session("session").await.unwrap();
15703
15704 agent.set_session_metadata(ai_agents_core::SessionMetadata {
15705 tags: vec!["updated".into()],
15706 ..Default::default()
15707 });
15708 agent.save_session("session").await.unwrap();
15709 assert!(
15710 agent
15711 .list_sessions_filtered(&ai_agents_core::SessionFilter {
15712 tags: Some(vec!["initial".into()]),
15713 ..Default::default()
15714 })
15715 .await
15716 .unwrap()
15717 .is_empty()
15718 );
15719 assert_eq!(
15720 agent
15721 .list_sessions_filtered(&ai_agents_core::SessionFilter {
15722 tags: Some(vec!["updated".into()]),
15723 ..Default::default()
15724 })
15725 .await
15726 .unwrap()
15727 .len(),
15728 1
15729 );
15730 drop(agent);
15731 storage.close().await;
15732 drop(storage);
15733
15734 let reopened_storage = Arc::new(
15735 ai_agents_storage::SqliteStorage::new(&path_string)
15736 .await
15737 .unwrap(),
15738 );
15739 let restored = runtime_storage_agent().with_storage(reopened_storage.clone());
15740 assert!(restored.load_session("session").await.unwrap());
15741 assert_eq!(restored.session_metadata().tags, vec!["updated"]);
15742 assert_eq!(
15743 restored.current_session_id.read().as_deref(),
15744 Some("session")
15745 );
15746 assert!(restored.save_state().await.unwrap().memory.messages.len() >= 2);
15747 assert_eq!(
15748 restored
15749 .list_sessions_filtered(&ai_agents_core::SessionFilter {
15750 tags: Some(vec!["updated".into()]),
15751 ..Default::default()
15752 })
15753 .await
15754 .unwrap()
15755 .len(),
15756 1
15757 );
15758
15759 drop(restored);
15760 reopened_storage.close().await;
15761 drop(reopened_storage);
15762 crate::remove_sqlite_test_directory(&directory)
15763 .await
15764 .unwrap();
15765 }
15766
15767 #[tokio::test]
15768 async fn storage_session_metadata_backend_failures_propagate() {
15769 let storage = Arc::new(RuntimeStorage::new([
15770 StorageCapability::Snapshot,
15771 StorageCapability::SessionMetadata,
15772 ]));
15773 let agent = runtime_storage_agent().with_storage(storage.clone());
15774
15775 agent.save_session("session").await.unwrap();
15776 storage
15777 .save("target", &agent.save_state().await.unwrap())
15778 .await
15779 .unwrap();
15780 storage.fail_metadata_load.store(true, Ordering::SeqCst);
15781 assert!(matches!(
15782 agent.load_session("target").await,
15783 Err(AgentError::Persistence(message)) if message == "metadata load failed"
15784 ));
15785 assert_eq!(agent.current_session_id.read().as_deref(), Some("session"));
15786
15787 storage.fail_metadata_save.store(true, Ordering::SeqCst);
15788 assert!(matches!(
15789 agent.save_session("session").await,
15790 Err(AgentError::Persistence(message)) if message == "metadata save failed"
15791 ));
15792 }
15793
15794 struct ProviderFutureDropSignal {
15795 dropped: Arc<AtomicBool>,
15796 }
15797
15798 impl Drop for ProviderFutureDropSignal {
15799 fn drop(&mut self) {
15800 self.dropped.store(true, Ordering::SeqCst);
15801 }
15802 }
15803
15804 struct BufferedLockingProvider {
15805 lock: Arc<tokio::sync::Mutex<()>>,
15806 stream_started: Arc<tokio::sync::Notify>,
15807 stream_dropped: Arc<AtomicBool>,
15808 committed_after_drop: Arc<AtomicBool>,
15809 }
15810
15811 #[async_trait]
15812 impl LLMProvider for BufferedLockingProvider {
15813 async fn complete(
15814 &self,
15815 _messages: &[ChatMessage],
15816 _config: Option<&LLMConfig>,
15817 ) -> std::result::Result<LLMResponse, LLMError> {
15818 let _guard = self.lock.lock().await;
15819 self.committed_after_drop
15820 .store(self.stream_dropped.load(Ordering::SeqCst), Ordering::SeqCst);
15821 Ok(LLMResponse::new(
15822 "Committed technical response.",
15823 FinishReason::Stop,
15824 ))
15825 }
15826
15827 async fn complete_stream(
15828 &self,
15829 _messages: &[ChatMessage],
15830 _config: Option<&LLMConfig>,
15831 ) -> std::result::Result<
15832 Box<dyn Stream<Item = std::result::Result<LLMChunk, LLMError>> + Unpin + Send>,
15833 LLMError,
15834 > {
15835 let _guard = self.lock.lock().await;
15836 let _drop_signal = ProviderFutureDropSignal {
15837 dropped: Arc::clone(&self.stream_dropped),
15838 };
15839 self.stream_started.notify_one();
15840 std::future::pending().await
15841 }
15842
15843 fn provider_name(&self) -> &str {
15844 "buffered-locking"
15845 }
15846
15847 fn supports(&self, _feature: LLMFeature) -> bool {
15848 false
15849 }
15850 }
15851
15852 struct PendingDropStream {
15853 dropped: Arc<AtomicBool>,
15854 dropped_notify: Arc<tokio::sync::Notify>,
15855 }
15856
15857 impl Stream for PendingDropStream {
15858 type Item = std::result::Result<LLMChunk, LLMError>;
15859
15860 fn poll_next(
15861 self: Pin<&mut Self>,
15862 _cx: &mut std::task::Context<'_>,
15863 ) -> std::task::Poll<Option<Self::Item>> {
15864 std::task::Poll::Pending
15865 }
15866 }
15867
15868 impl Drop for PendingDropStream {
15869 fn drop(&mut self) {
15870 self.dropped.store(true, Ordering::SeqCst);
15871 self.dropped_notify.notify_one();
15872 }
15873 }
15874
15875 struct EstablishedStreamProvider {
15876 stream_started: Arc<tokio::sync::Notify>,
15877 stream_dropped: Arc<AtomicBool>,
15878 stream_dropped_notify: Arc<tokio::sync::Notify>,
15879 committed_after_drop: Arc<AtomicBool>,
15880 }
15881
15882 #[async_trait]
15883 impl LLMProvider for EstablishedStreamProvider {
15884 async fn complete(
15885 &self,
15886 _messages: &[ChatMessage],
15887 _config: Option<&LLMConfig>,
15888 ) -> std::result::Result<LLMResponse, LLMError> {
15889 if !self.stream_dropped.load(Ordering::SeqCst) {
15890 self.stream_dropped_notify.notified().await;
15891 }
15892 self.committed_after_drop
15893 .store(self.stream_dropped.load(Ordering::SeqCst), Ordering::SeqCst);
15894 Ok(LLMResponse::new(
15895 "Committed technical response.",
15896 FinishReason::Stop,
15897 ))
15898 }
15899
15900 async fn complete_stream(
15901 &self,
15902 _messages: &[ChatMessage],
15903 _config: Option<&LLMConfig>,
15904 ) -> std::result::Result<
15905 Box<dyn Stream<Item = std::result::Result<LLMChunk, LLMError>> + Unpin + Send>,
15906 LLMError,
15907 > {
15908 self.stream_started.notify_one();
15909 Ok(Box::new(PendingDropStream {
15910 dropped: Arc::clone(&self.stream_dropped),
15911 dropped_notify: Arc::clone(&self.stream_dropped_notify),
15912 }))
15913 }
15914
15915 fn provider_name(&self) -> &str {
15916 "established-stream"
15917 }
15918
15919 fn supports(&self, _feature: LLMFeature) -> bool {
15920 false
15921 }
15922 }
15923
15924 struct FirstCallLockingProvider {
15925 lock: Arc<tokio::sync::Mutex<()>>,
15926 first_started: Arc<tokio::sync::Notify>,
15927 first_dropped: Arc<AtomicBool>,
15928 committed_after_drop: Arc<AtomicBool>,
15929 calls: AtomicU64,
15930 }
15931
15932 #[async_trait]
15933 impl LLMProvider for FirstCallLockingProvider {
15934 async fn complete(
15935 &self,
15936 _messages: &[ChatMessage],
15937 _config: Option<&LLMConfig>,
15938 ) -> std::result::Result<LLMResponse, LLMError> {
15939 let _guard = self.lock.lock().await;
15940 let call = self.calls.fetch_add(1, Ordering::SeqCst);
15941 if call == 0 {
15942 let _drop_signal = ProviderFutureDropSignal {
15943 dropped: Arc::clone(&self.first_dropped),
15944 };
15945 self.first_started.notify_one();
15946 return std::future::pending().await;
15947 }
15948 self.committed_after_drop
15949 .store(self.first_dropped.load(Ordering::SeqCst), Ordering::SeqCst);
15950 Ok(LLMResponse::new(
15951 "Committed technical response.",
15952 FinishReason::Stop,
15953 ))
15954 }
15955
15956 async fn complete_stream(
15957 &self,
15958 _messages: &[ChatMessage],
15959 _config: Option<&LLMConfig>,
15960 ) -> std::result::Result<
15961 Box<dyn Stream<Item = std::result::Result<LLMChunk, LLMError>> + Unpin + Send>,
15962 LLMError,
15963 > {
15964 Err(LLMError::Other(
15965 "streaming is not used in this test".to_string(),
15966 ))
15967 }
15968
15969 fn provider_name(&self) -> &str {
15970 "first-call-locking"
15971 }
15972
15973 fn supports(&self, _feature: LLMFeature) -> bool {
15974 false
15975 }
15976 }
15977
15978 struct RoutingAfterProviderStart {
15979 provider_started: Arc<tokio::sync::Notify>,
15980 }
15981
15982 #[async_trait]
15983 impl LLMProvider for RoutingAfterProviderStart {
15984 async fn complete(
15985 &self,
15986 _messages: &[ChatMessage],
15987 _config: Option<&LLMConfig>,
15988 ) -> std::result::Result<LLMResponse, LLMError> {
15989 self.provider_started.notified().await;
15990 Ok(LLMResponse::new("1", FinishReason::Stop))
15991 }
15992
15993 async fn complete_stream(
15994 &self,
15995 _messages: &[ChatMessage],
15996 _config: Option<&LLMConfig>,
15997 ) -> std::result::Result<
15998 Box<dyn Stream<Item = std::result::Result<LLMChunk, LLMError>> + Unpin + Send>,
15999 LLMError,
16000 > {
16001 Err(LLMError::Other(
16002 "streaming is not used in this test".to_string(),
16003 ))
16004 }
16005
16006 fn provider_name(&self) -> &str {
16007 "routing-after-start"
16008 }
16009
16010 fn supports(&self, _feature: LLMFeature) -> bool {
16011 false
16012 }
16013 }
16014
16015 struct ResponseCountingHooks {
16017 responses: Arc<std::sync::atomic::AtomicUsize>,
16018 }
16019
16020 struct RootTurnProbeProvider {
16022 complete_entered: tokio::sync::mpsc::UnboundedSender<()>,
16023 }
16024
16025 struct ResponseChatHooks {
16027 target: parking_lot::Mutex<Option<Weak<RuntimeAgent>>>,
16028 invoked: AtomicBool,
16029 nested_result: parking_lot::Mutex<Option<std::result::Result<String, String>>>,
16030 }
16031
16032 struct ConcurrentResponseHooks {
16034 registry: Weak<crate::spawner::AgentRegistry>,
16035 child_id: String,
16036 invoked: AtomicBool,
16037 nested_result: parking_lot::Mutex<Option<std::result::Result<String, String>>>,
16038 }
16039
16040 struct RetryDeadlineTool {
16042 calls: Arc<std::sync::atomic::AtomicUsize>,
16043 deadlines: Arc<parking_lot::Mutex<Vec<chrono::DateTime<chrono::Utc>>>>,
16044 remaining_ms: Arc<parking_lot::Mutex<Vec<i64>>>,
16045 }
16046
16047 struct ToolLifecycleRecordingHooks {
16049 events: parking_lot::Mutex<Vec<String>>,
16050 records: parking_lot::Mutex<Vec<ToolExecutionRecord>>,
16051 }
16052
16053 impl ToolLifecycleRecordingHooks {
16054 fn new() -> Self {
16056 Self {
16057 events: parking_lot::Mutex::new(Vec::new()),
16058 records: parking_lot::Mutex::new(Vec::new()),
16059 }
16060 }
16061
16062 fn events(&self) -> Vec<String> {
16064 self.events.lock().clone()
16065 }
16066
16067 fn records(&self) -> Vec<ToolExecutionRecord> {
16069 self.records.lock().clone()
16070 }
16071 }
16072
16073 struct ContextEchoTool;
16075
16076 #[async_trait]
16077 impl LLMProvider for RootTurnProbeProvider {
16078 async fn complete(
16079 &self,
16080 _messages: &[ChatMessage],
16081 _config: Option<&LLMConfig>,
16082 ) -> std::result::Result<LLMResponse, LLMError> {
16083 let _ = self.complete_entered.send(());
16084 Ok(LLMResponse::new("blocking complete", FinishReason::Stop))
16085 }
16086
16087 async fn complete_stream(
16088 &self,
16089 _messages: &[ChatMessage],
16090 _config: Option<&LLMConfig>,
16091 ) -> std::result::Result<
16092 Box<dyn Stream<Item = std::result::Result<LLMChunk, LLMError>> + Unpin + Send>,
16093 LLMError,
16094 > {
16095 Ok(Box::new(futures::stream::iter(vec![Ok(
16096 LLMChunk::final_chunk("stream complete", FinishReason::Stop, None),
16097 )])))
16098 }
16099
16100 fn provider_name(&self) -> &str {
16101 "root-turn-probe"
16102 }
16103
16104 fn supports(&self, feature: LLMFeature) -> bool {
16105 matches!(feature, LLMFeature::Streaming)
16106 }
16107 }
16108
16109 #[async_trait]
16110 impl ai_agents_core::Tool for ContextEchoTool {
16111 fn id(&self) -> &str {
16112 "context_echo"
16113 }
16114
16115 fn name(&self) -> &str {
16116 "Context Echo"
16117 }
16118
16119 fn description(&self) -> &str {
16120 "Returns selected execution context fields."
16121 }
16122
16123 fn input_schema(&self) -> Value {
16124 serde_json::json!({"type": "object"})
16125 }
16126
16127 fn policy_bindings(&self) -> ai_agents_core::ToolPolicyBindings {
16128 ai_agents_core::ToolPolicyBindings {
16129 path_fields: vec![ai_agents_core::PathPolicyBinding::read("path")],
16130 result_limit_fields: vec![ai_agents_core::ResultLimitBinding::new(
16131 "max_results",
16132 ai_agents_core::ResultLimitKind::MaxResults,
16133 )],
16134 ..Default::default()
16135 }
16136 }
16137
16138 async fn execute(
16139 &self,
16140 _args: Value,
16141 ctx: ai_agents_core::ToolExecutionContext,
16142 ) -> ToolResult {
16143 ToolResult::ok(
16144 serde_json::json!({
16145 "requested_name": ctx.requested_name,
16146 "canonical_id": ctx.canonical_id,
16147 "display_name": ctx.display_name,
16148 "max_results": ctx.limits.max_results,
16149 "custom_config": ctx.custom_config,
16150 })
16151 .to_string(),
16152 )
16153 }
16154 }
16155
16156 #[async_trait]
16157 impl ai_agents_core::Tool for RetryDeadlineTool {
16158 fn id(&self) -> &str {
16159 "retry_deadline"
16160 }
16161
16162 fn name(&self) -> &str {
16163 "Retry Deadline"
16164 }
16165
16166 fn description(&self) -> &str {
16167 "Records one deadline per retry invocation."
16168 }
16169
16170 fn input_schema(&self) -> Value {
16171 serde_json::json!({"type": "object"})
16172 }
16173
16174 fn safety_metadata(&self) -> ai_agents_core::ToolSafetyMetadata {
16175 ai_agents_core::ToolSafetyMetadata::compute()
16176 }
16177
16178 fn classify_call(&self, _args: &Value) -> ai_agents_core::ToolCallClassification {
16179 let mut classification =
16180 ai_agents_core::ToolCallClassification::from_metadata(&self.safety_metadata());
16181 classification.timeout_ms = Some(1_000);
16182 classification.safely_retryable = true;
16183 classification
16184 }
16185
16186 async fn execute(
16188 &self,
16189 _args: Value,
16190 ctx: ai_agents_core::ToolExecutionContext,
16191 ) -> ToolResult {
16192 let deadline = ctx
16193 .deadline
16194 .expect("each invocation must receive a deadline");
16195 self.remaining_ms.lock().push(
16196 deadline
16197 .signed_duration_since(chrono::Utc::now())
16198 .num_milliseconds(),
16199 );
16200 self.deadlines.lock().push(deadline);
16201 let call = self.calls.fetch_add(1, Ordering::SeqCst);
16202 if call == 0 {
16203 tokio::time::sleep(std::time::Duration::from_millis(20)).await;
16204 ToolResult::error("retry")
16205 } else {
16206 ToolResult::ok("done")
16207 }
16208 }
16209 }
16210
16211 struct ClassifiedTimeoutTool {
16213 id: &'static str,
16214 calls: Arc<std::sync::atomic::AtomicUsize>,
16215 timeout_ms: u64,
16216 sleep_ms: u64,
16217 requires_approval: bool,
16218 remaining_ms: Arc<parking_lot::Mutex<Vec<i64>>>,
16219 }
16220
16221 struct ApprovalModifiedTimeoutTool {
16223 calls: Arc<std::sync::atomic::AtomicUsize>,
16224 }
16225
16226 struct SlowTool;
16228
16229 struct FlakyWriteTool {
16231 calls: Arc<std::sync::atomic::AtomicUsize>,
16232 }
16233
16234 struct LockedWriteTool {
16236 active: Arc<std::sync::atomic::AtomicUsize>,
16237 max_active: Arc<std::sync::atomic::AtomicUsize>,
16238 }
16239
16240 struct MultiResourceWriteTool {
16241 active: Arc<std::sync::atomic::AtomicUsize>,
16242 max_active: Arc<std::sync::atomic::AtomicUsize>,
16243 }
16244
16245 #[derive(Clone)]
16246 struct PathMutationGate {
16247 entered: Arc<AtomicBool>,
16248 entered_notify: Arc<tokio::sync::Notify>,
16249 release: Arc<tokio::sync::Notify>,
16250 }
16251
16252 impl PathMutationGate {
16253 fn new() -> Self {
16254 Self {
16255 entered: Arc::new(AtomicBool::new(false)),
16256 entered_notify: Arc::new(tokio::sync::Notify::new()),
16257 release: Arc::new(tokio::sync::Notify::new()),
16258 }
16259 }
16260
16261 async fn wait_until_entered(&self) {
16262 if !self.entered.load(Ordering::SeqCst) {
16263 self.entered_notify.notified().await;
16264 }
16265 }
16266
16267 fn release(&self) {
16268 self.release.notify_one();
16269 }
16270 }
16271
16272 struct BlockingPathMutationTool {
16273 id: &'static str,
16274 path_fields: Vec<ai_agents_core::PathPolicyBinding>,
16275 gate: PathMutationGate,
16276 }
16277
16278 struct NoBindingWriteTool {
16279 active: Arc<std::sync::atomic::AtomicUsize>,
16280 max_active: Arc<std::sync::atomic::AtomicUsize>,
16281 }
16282
16283 struct RecoveryTestTool {
16284 id: String,
16285 succeeds: bool,
16286 calls: Arc<std::sync::atomic::AtomicUsize>,
16287 max_output_chars: Option<usize>,
16288 }
16289
16290 struct BlockingApprovalHandler {
16291 entered: Arc<tokio::sync::Barrier>,
16292 release: Arc<tokio::sync::Notify>,
16293 result: ApprovalResult,
16294 }
16295
16296 struct CountingApprovalHandler {
16297 calls: Arc<std::sync::atomic::AtomicUsize>,
16298 }
16299
16300 struct DriftingFallbackProvider {
16302 refreshed: AtomicBool,
16303 primary_calls: Arc<std::sync::atomic::AtomicUsize>,
16304 secondary_calls: Arc<std::sync::atomic::AtomicUsize>,
16305 }
16306
16307 struct RefreshFallbackProviderHooks {
16309 agent: parking_lot::Mutex<Option<Weak<RuntimeAgent>>>,
16310 lifecycle: Arc<ToolLifecycleRecordingHooks>,
16311 }
16312
16313 struct RuntimeWebFetchTransport {
16314 calls: Arc<std::sync::atomic::AtomicUsize>,
16315 }
16316
16317 struct RuntimeWebFetchResolver;
16318
16319 struct ReentrantToolHooks {
16320 agent: parking_lot::Mutex<Option<Weak<RuntimeAgent>>>,
16321 invoked: AtomicBool,
16322 nested_success: AtomicBool,
16323 }
16324
16325 #[async_trait]
16326 impl ai_agents_core::Tool for ClassifiedTimeoutTool {
16327 fn id(&self) -> &str {
16329 self.id
16330 }
16331
16332 fn name(&self) -> &str {
16334 "Classified Timeout"
16335 }
16336
16337 fn description(&self) -> &str {
16339 "Records and waits under one call-level timeout."
16340 }
16341
16342 fn input_schema(&self) -> Value {
16344 serde_json::json!({"type": "object"})
16345 }
16346
16347 fn classify_call(&self, _args: &Value) -> ai_agents_core::ToolCallClassification {
16349 let mut classification =
16350 ai_agents_core::ToolCallClassification::from_metadata(&self.safety_metadata());
16351 classification.timeout_ms = Some(self.timeout_ms);
16352 classification.requires_approval = self.requires_approval;
16353 classification
16354 }
16355
16356 async fn execute(
16358 &self,
16359 _args: Value,
16360 ctx: ai_agents_core::ToolExecutionContext,
16361 ) -> ToolResult {
16362 self.calls.fetch_add(1, Ordering::SeqCst);
16363 let deadline = ctx
16364 .deadline
16365 .expect("each invocation must receive a deadline");
16366 self.remaining_ms.lock().push(
16367 deadline
16368 .signed_duration_since(chrono::Utc::now())
16369 .num_milliseconds(),
16370 );
16371 tokio::time::sleep(Duration::from_millis(self.sleep_ms)).await;
16372 ToolResult::ok("done")
16373 }
16374 }
16375
16376 #[async_trait]
16377 impl ai_agents_core::Tool for ApprovalModifiedTimeoutTool {
16378 fn id(&self) -> &str {
16380 "approval_modified_timeout"
16381 }
16382
16383 fn name(&self) -> &str {
16385 "Approval Modified Timeout"
16386 }
16387
16388 fn description(&self) -> &str {
16390 "Becomes invalid only after approval modifies its arguments."
16391 }
16392
16393 fn input_schema(&self) -> Value {
16395 serde_json::json!({"type": "object"})
16396 }
16397
16398 fn policy_bindings(&self) -> ai_agents_core::ToolPolicyBindings {
16400 ai_agents_core::ToolPolicyBindings {
16401 path_fields: vec![ai_agents_core::PathPolicyBinding::write("path")],
16402 ..Default::default()
16403 }
16404 }
16405
16406 fn safety_metadata(&self) -> ai_agents_core::ToolSafetyMetadata {
16408 ai_agents_core::ToolSafetyMetadata {
16409 read_only: false,
16410 concurrency_safe: false,
16411 operation: ai_agents_core::ToolOperationKind::Write,
16412 side_effect_level: ai_agents_core::ToolSideEffectLevel::LocalWrite,
16413 requires_network: false,
16414 destructive: false,
16415 open_world: false,
16416 host_dependent: false,
16417 requires_user_interaction: false,
16418 supports_cancellation: true,
16419 default_requires_approval: true,
16420 should_defer_schema: false,
16421 max_output_chars: Some(1024),
16422 max_result_size_chars: Some(1024),
16423 }
16424 }
16425
16426 fn classify_call(&self, args: &Value) -> ai_agents_core::ToolCallClassification {
16428 let mut classification =
16429 ai_agents_core::ToolCallClassification::from_metadata(&self.safety_metadata());
16430 classification.timeout_ms = Some(if args["invalid_timeout"].as_bool() == Some(true) {
16431 u64::MAX
16432 } else {
16433 1_000
16434 });
16435 classification
16436 }
16437
16438 async fn execute(
16440 &self,
16441 _args: Value,
16442 _ctx: ai_agents_core::ToolExecutionContext,
16443 ) -> ToolResult {
16444 self.calls.fetch_add(1, Ordering::SeqCst);
16445 ToolResult::ok("unexpected")
16446 }
16447 }
16448
16449 #[async_trait]
16450 impl ai_agents_core::Tool for SlowTool {
16451 fn id(&self) -> &str {
16452 "slow"
16453 }
16454
16455 fn name(&self) -> &str {
16456 "Slow"
16457 }
16458
16459 fn description(&self) -> &str {
16460 "Waits until cancelled or timed out."
16461 }
16462
16463 fn input_schema(&self) -> Value {
16464 serde_json::json!({"type": "object"})
16465 }
16466
16467 async fn execute(
16468 &self,
16469 _args: Value,
16470 _ctx: ai_agents_core::ToolExecutionContext,
16471 ) -> ToolResult {
16472 tokio::time::sleep(std::time::Duration::from_secs(5)).await;
16473 ToolResult::ok("done")
16474 }
16475 }
16476
16477 #[async_trait]
16478 impl ai_agents_core::Tool for FlakyWriteTool {
16479 fn id(&self) -> &str {
16480 "flaky_write"
16481 }
16482
16483 fn name(&self) -> &str {
16484 "Flaky Write"
16485 }
16486
16487 fn description(&self) -> &str {
16488 "Fails on the first write attempt."
16489 }
16490
16491 fn input_schema(&self) -> Value {
16492 serde_json::json!({"type": "object", "properties": {"path": {"type": "string"}}})
16493 }
16494
16495 fn policy_bindings(&self) -> ai_agents_core::ToolPolicyBindings {
16496 ai_agents_core::ToolPolicyBindings {
16497 path_fields: vec![ai_agents_core::PathPolicyBinding::write("path")],
16498 ..Default::default()
16499 }
16500 }
16501
16502 fn safety_metadata(&self) -> ai_agents_core::ToolSafetyMetadata {
16503 ai_agents_core::ToolSafetyMetadata {
16504 read_only: false,
16505 concurrency_safe: false,
16506 operation: ai_agents_core::ToolOperationKind::Write,
16507 side_effect_level: ai_agents_core::ToolSideEffectLevel::LocalWrite,
16508 requires_network: false,
16509 destructive: false,
16510 open_world: false,
16511 host_dependent: false,
16512 requires_user_interaction: false,
16513 supports_cancellation: true,
16514 default_requires_approval: false,
16515 should_defer_schema: false,
16516 max_output_chars: Some(1024),
16517 max_result_size_chars: Some(1024),
16518 }
16519 }
16520
16521 fn classify_call(&self, _args: &Value) -> ai_agents_core::ToolCallClassification {
16522 let mut classification =
16523 ai_agents_core::ToolCallClassification::from_metadata(&self.safety_metadata());
16524 classification.safely_retryable = false;
16525 classification
16526 }
16527
16528 async fn execute(
16529 &self,
16530 _args: Value,
16531 _ctx: ai_agents_core::ToolExecutionContext,
16532 ) -> ToolResult {
16533 let call = self.calls.fetch_add(1, Ordering::SeqCst);
16534 if call == 0 {
16535 ToolResult::error("first failure")
16536 } else {
16537 ToolResult::ok("second success")
16538 }
16539 }
16540 }
16541
16542 #[async_trait]
16543 impl ai_agents_core::Tool for LockedWriteTool {
16544 fn id(&self) -> &str {
16545 "locked_write"
16546 }
16547
16548 fn name(&self) -> &str {
16549 "Locked Write"
16550 }
16551
16552 fn description(&self) -> &str {
16553 "Tracks concurrent execution on one resource."
16554 }
16555
16556 fn input_schema(&self) -> Value {
16557 serde_json::json!({"type": "object", "properties": {"path": {"type": "string"}}})
16558 }
16559
16560 fn policy_bindings(&self) -> ai_agents_core::ToolPolicyBindings {
16561 ai_agents_core::ToolPolicyBindings {
16562 path_fields: vec![ai_agents_core::PathPolicyBinding::write("path")],
16563 ..Default::default()
16564 }
16565 }
16566
16567 fn safety_metadata(&self) -> ai_agents_core::ToolSafetyMetadata {
16568 ai_agents_core::ToolSafetyMetadata {
16569 read_only: false,
16570 concurrency_safe: false,
16571 operation: ai_agents_core::ToolOperationKind::Write,
16572 side_effect_level: ai_agents_core::ToolSideEffectLevel::LocalWrite,
16573 requires_network: false,
16574 destructive: false,
16575 open_world: false,
16576 host_dependent: false,
16577 requires_user_interaction: false,
16578 supports_cancellation: true,
16579 default_requires_approval: false,
16580 should_defer_schema: false,
16581 max_output_chars: Some(1024),
16582 max_result_size_chars: Some(1024),
16583 }
16584 }
16585
16586 async fn execute(
16587 &self,
16588 _args: Value,
16589 _ctx: ai_agents_core::ToolExecutionContext,
16590 ) -> ToolResult {
16591 let active = self.active.fetch_add(1, Ordering::SeqCst) + 1;
16592 loop {
16593 let current_max = self.max_active.load(Ordering::SeqCst);
16594 if active <= current_max {
16595 break;
16596 }
16597 if self
16598 .max_active
16599 .compare_exchange(current_max, active, Ordering::SeqCst, Ordering::SeqCst)
16600 .is_ok()
16601 {
16602 break;
16603 }
16604 }
16605 tokio::time::sleep(std::time::Duration::from_millis(100)).await;
16606 self.active.fetch_sub(1, Ordering::SeqCst);
16607 ToolResult::ok("done")
16608 }
16609 }
16610
16611 #[async_trait]
16612 impl ai_agents_core::Tool for MultiResourceWriteTool {
16613 fn id(&self) -> &str {
16614 "multi_resource_write"
16615 }
16616
16617 fn name(&self) -> &str {
16618 "Multi Resource Write"
16619 }
16620
16621 fn description(&self) -> &str {
16622 "Tracks concurrent execution across source and destination resources."
16623 }
16624
16625 fn input_schema(&self) -> Value {
16626 serde_json::json!({"type": "object"})
16627 }
16628
16629 fn policy_bindings(&self) -> ai_agents_core::ToolPolicyBindings {
16630 ai_agents_core::ToolPolicyBindings {
16631 path_fields: vec![
16632 ai_agents_core::PathPolicyBinding::read_write("source_path"),
16633 ai_agents_core::PathPolicyBinding::write("destination_path"),
16634 ],
16635 ..Default::default()
16636 }
16637 }
16638
16639 fn safety_metadata(&self) -> ai_agents_core::ToolSafetyMetadata {
16640 LockedWriteTool {
16641 active: Arc::clone(&self.active),
16642 max_active: Arc::clone(&self.max_active),
16643 }
16644 .safety_metadata()
16645 }
16646
16647 async fn execute(
16648 &self,
16649 _args: Value,
16650 _ctx: ai_agents_core::ToolExecutionContext,
16651 ) -> ToolResult {
16652 let active = self.active.fetch_add(1, Ordering::SeqCst) + 1;
16653 self.max_active.fetch_max(active, Ordering::SeqCst);
16654 tokio::time::sleep(std::time::Duration::from_millis(75)).await;
16655 self.active.fetch_sub(1, Ordering::SeqCst);
16656 ToolResult::ok("done")
16657 }
16658 }
16659
16660 #[async_trait]
16661 impl ai_agents_core::Tool for BlockingPathMutationTool {
16662 fn id(&self) -> &str {
16663 self.id
16664 }
16665
16666 fn name(&self) -> &str {
16667 self.id
16668 }
16669
16670 fn description(&self) -> &str {
16671 "Blocks a path mutation until the test releases it."
16672 }
16673
16674 fn input_schema(&self) -> Value {
16675 serde_json::json!({"type": "object"})
16676 }
16677
16678 fn policy_bindings(&self) -> ai_agents_core::ToolPolicyBindings {
16679 ai_agents_core::ToolPolicyBindings {
16680 path_fields: self.path_fields.clone(),
16681 ..Default::default()
16682 }
16683 }
16684
16685 fn safety_metadata(&self) -> ai_agents_core::ToolSafetyMetadata {
16686 ai_agents_core::ToolSafetyMetadata {
16687 read_only: false,
16688 concurrency_safe: false,
16689 operation: ai_agents_core::ToolOperationKind::Write,
16690 side_effect_level: ai_agents_core::ToolSideEffectLevel::LocalWrite,
16691 requires_network: false,
16692 destructive: false,
16693 open_world: false,
16694 host_dependent: false,
16695 requires_user_interaction: false,
16696 supports_cancellation: true,
16697 default_requires_approval: false,
16698 should_defer_schema: false,
16699 max_output_chars: Some(1024),
16700 max_result_size_chars: Some(1024),
16701 }
16702 }
16703
16704 async fn execute(
16705 &self,
16706 _args: Value,
16707 _ctx: ai_agents_core::ToolExecutionContext,
16708 ) -> ToolResult {
16709 self.gate.entered.store(true, Ordering::SeqCst);
16710 self.gate.entered_notify.notify_one();
16711 self.gate.release.notified().await;
16712 ToolResult::ok("done")
16713 }
16714 }
16715
16716 #[async_trait]
16717 impl ai_agents_core::Tool for NoBindingWriteTool {
16718 fn id(&self) -> &str {
16719 "no_binding_write"
16720 }
16721
16722 fn name(&self) -> &str {
16723 "No Binding Write"
16724 }
16725
16726 fn description(&self) -> &str {
16727 "Tracks concurrent execution without resource bindings."
16728 }
16729
16730 fn input_schema(&self) -> Value {
16731 serde_json::json!({"type": "object"})
16732 }
16733
16734 fn safety_metadata(&self) -> ai_agents_core::ToolSafetyMetadata {
16735 LockedWriteTool {
16736 active: Arc::clone(&self.active),
16737 max_active: Arc::clone(&self.max_active),
16738 }
16739 .safety_metadata()
16740 }
16741
16742 async fn execute(
16743 &self,
16744 _args: Value,
16745 _ctx: ai_agents_core::ToolExecutionContext,
16746 ) -> ToolResult {
16747 let active = self.active.fetch_add(1, Ordering::SeqCst) + 1;
16748 self.max_active.fetch_max(active, Ordering::SeqCst);
16749 tokio::time::sleep(std::time::Duration::from_millis(75)).await;
16750 self.active.fetch_sub(1, Ordering::SeqCst);
16751 ToolResult::ok("done")
16752 }
16753 }
16754
16755 #[async_trait]
16756 impl ai_agents_core::Tool for RecoveryTestTool {
16757 fn id(&self) -> &str {
16758 &self.id
16759 }
16760
16761 fn name(&self) -> &str {
16762 &self.id
16763 }
16764
16765 fn description(&self) -> &str {
16766 "Records recovery execution and returns a configured result."
16767 }
16768
16769 fn input_schema(&self) -> Value {
16770 serde_json::json!({"type": "object"})
16771 }
16772
16773 fn policy_bindings(&self) -> ai_agents_core::ToolPolicyBindings {
16774 ai_agents_core::ToolPolicyBindings {
16775 path_fields: vec![ai_agents_core::PathPolicyBinding::write("path")],
16776 ..Default::default()
16777 }
16778 }
16779
16780 fn safety_metadata(&self) -> ai_agents_core::ToolSafetyMetadata {
16782 ai_agents_core::ToolSafetyMetadata {
16783 read_only: false,
16784 concurrency_safe: false,
16785 operation: ai_agents_core::ToolOperationKind::Write,
16786 side_effect_level: ai_agents_core::ToolSideEffectLevel::LocalWrite,
16787 requires_network: false,
16788 destructive: false,
16789 open_world: false,
16790 host_dependent: false,
16791 requires_user_interaction: false,
16792 supports_cancellation: true,
16793 default_requires_approval: false,
16794 should_defer_schema: false,
16795 max_output_chars: Some(self.max_output_chars.unwrap_or(1024)),
16796 max_result_size_chars: Some(1024),
16797 }
16798 }
16799
16800 async fn execute(
16802 &self,
16803 _args: Value,
16804 _ctx: ai_agents_core::ToolExecutionContext,
16805 ) -> ToolResult {
16806 self.calls.fetch_add(1, Ordering::SeqCst);
16807 let mut result = if self.succeeds {
16808 ToolResult::ok(format!("{} succeeded", self.id))
16809 } else {
16810 ToolResult::error(format!("{} failed", self.id))
16811 };
16812 result.metadata = Some(HashMap::from([(
16813 "recovery_test_tool".to_string(),
16814 Value::String(self.id.clone()),
16815 )]));
16816 result
16817 }
16818 }
16819
16820 #[async_trait]
16821 impl WebFetchTransport for RuntimeWebFetchTransport {
16822 async fn send(
16824 &self,
16825 _request: WebFetchTransportRequest,
16826 ) -> std::result::Result<WebFetchTransportResponse, String> {
16827 Err("validated addresses are required".to_string())
16828 }
16829
16830 async fn send_validated(
16832 &self,
16833 _request: WebFetchTransportRequest,
16834 _addresses: &[std::net::SocketAddr],
16835 ) -> std::result::Result<WebFetchTransportResponse, String> {
16836 self.calls.fetch_add(1, Ordering::SeqCst);
16837 Ok(WebFetchTransportResponse {
16838 status: 200,
16839 content_type: Some("text/plain".to_string()),
16840 location: None,
16841 body: b"approved".to_vec(),
16842 })
16843 }
16844 }
16845
16846 #[async_trait]
16847 impl WebFetchResolver for RuntimeWebFetchResolver {
16848 async fn resolve(
16850 &self,
16851 _host: &str,
16852 _port: u16,
16853 ) -> std::result::Result<Vec<std::net::IpAddr>, String> {
16854 Ok(vec![std::net::IpAddr::V4(std::net::Ipv4Addr::new(
16855 93, 184, 216, 34,
16856 ))])
16857 }
16858 }
16859
16860 #[async_trait]
16861 impl ToolProvider for DriftingFallbackProvider {
16862 fn id(&self) -> &str {
16864 "drifting_fallback"
16865 }
16866
16867 fn name(&self) -> &str {
16869 "Drifting Fallback"
16870 }
16871
16872 fn provider_type(&self) -> ToolProviderType {
16874 ToolProviderType::Custom
16875 }
16876
16877 async fn list_tools(&self) -> Vec<ToolDescriptor> {
16879 let alias = ToolAliases::new().with_name("en", "fallback alias");
16880 let mut primary = ToolDescriptor::new(
16881 "primary",
16882 "Primary",
16883 "Fails before fallback.",
16884 serde_json::json!({"type": "object"}),
16885 );
16886 let mut secondary = ToolDescriptor::new(
16887 "secondary",
16888 "Secondary",
16889 "Must not execute after final canonical drift.",
16890 serde_json::json!({"type": "object"}),
16891 );
16892 if self.refreshed.load(Ordering::SeqCst) {
16893 primary = primary.with_aliases(alias);
16894 } else {
16895 secondary = secondary.with_aliases(alias);
16896 }
16897 vec![primary, secondary]
16898 }
16899
16900 async fn get_tool(&self, tool_id: &str) -> Option<Arc<dyn Tool>> {
16902 let calls = match tool_id {
16903 "primary" => Arc::clone(&self.primary_calls),
16904 "secondary" => Arc::clone(&self.secondary_calls),
16905 _ => return None,
16906 };
16907 Some(Arc::new(RecoveryTestTool {
16908 id: tool_id.to_string(),
16909 succeeds: false,
16910 calls,
16911 max_output_chars: None,
16912 }))
16913 }
16914
16915 fn supports_refresh(&self) -> bool {
16917 true
16918 }
16919
16920 async fn refresh(&self) -> std::result::Result<(), ToolProviderError> {
16922 self.refreshed.store(true, Ordering::SeqCst);
16923 Ok(())
16924 }
16925 }
16926
16927 #[async_trait]
16928 impl AgentHooks for RefreshFallbackProviderHooks {
16929 async fn on_tool_start(&self, tool: &str, args: &Value) {
16931 self.lifecycle.on_tool_start(tool, args).await;
16932 if tool != "secondary" {
16933 return;
16934 }
16935 let agent = self.agent.lock().as_ref().and_then(Weak::upgrade);
16936 if let Some(agent) = agent {
16937 agent
16938 .tools
16939 .refresh_provider("drifting_fallback")
16940 .await
16941 .unwrap();
16942 }
16943 }
16944
16945 async fn on_tool_complete(&self, tool: &str, result: &ToolResult, duration_ms: u64) {
16946 self.lifecycle
16947 .on_tool_complete(tool, result, duration_ms)
16948 .await;
16949 }
16950
16951 async fn on_tool_execution_record(&self, record: &ToolExecutionRecord) {
16952 self.lifecycle.on_tool_execution_record(record).await;
16953 }
16954
16955 async fn on_error(&self, error: &AgentError) {
16956 self.lifecycle.on_error(error).await;
16957 }
16958 }
16959
16960 #[async_trait]
16961 impl ApprovalHandler for BlockingApprovalHandler {
16962 async fn request_approval(
16963 &self,
16964 _request: ai_agents_hitl::ApprovalRequest,
16965 ) -> ApprovalResult {
16966 self.entered.wait().await;
16967 self.release.notified().await;
16968 self.result.clone()
16969 }
16970 }
16971
16972 #[async_trait]
16973 impl ApprovalHandler for CountingApprovalHandler {
16974 async fn request_approval(
16975 &self,
16976 _request: ai_agents_hitl::ApprovalRequest,
16977 ) -> ApprovalResult {
16978 self.calls.fetch_add(1, Ordering::SeqCst);
16979 ApprovalResult::Approved
16980 }
16981 }
16982
16983 #[async_trait]
16984 impl AgentHooks for ReentrantToolHooks {
16985 async fn on_tool_complete(&self, tool: &str, _result: &ToolResult, _duration_ms: u64) {
16986 if tool != "reentrant_write" || self.invoked.swap(true, Ordering::SeqCst) {
16987 return;
16988 }
16989 let agent = self.agent.lock().as_ref().and_then(Weak::upgrade);
16990 if let Some(agent) = agent {
16991 let result = agent
16992 .invoke_tool(ToolExecutionRequest::new(
16993 "nested-hook-call",
16994 "reentrant_write",
16995 serde_json::json!({"path": "./hook.txt"}),
16996 ToolCallSource::Manual,
16997 ))
16998 .await;
16999 self.nested_success
17000 .store(result.is_ok_and(|record| record.success), Ordering::SeqCst);
17001 }
17002 }
17003 }
17004
17005 #[async_trait]
17006 impl AgentHooks for ResponseCountingHooks {
17007 async fn on_response(&self, _response: &AgentResponse) {
17008 self.responses.fetch_add(1, Ordering::SeqCst);
17009 }
17010 }
17011
17012 #[async_trait]
17013 impl AgentHooks for ResponseChatHooks {
17014 async fn on_response(&self, _response: &AgentResponse) {
17016 if self.invoked.swap(true, Ordering::SeqCst) {
17017 return;
17018 }
17019 let target = self.target.lock().as_ref().and_then(Weak::upgrade);
17020 let result = if let Some(target) = target {
17021 target
17022 .chat("nested response hook call")
17023 .await
17024 .map(|response| response.content)
17025 .map_err(|error| error.to_string())
17026 } else {
17027 Err("response hook target is unavailable".to_string())
17028 };
17029 *self.nested_result.lock() = Some(result);
17030 }
17031 }
17032
17033 #[async_trait]
17034 impl AgentHooks for ConcurrentResponseHooks {
17035 async fn on_response(&self, _response: &AgentResponse) {
17037 if self.invoked.swap(true, Ordering::SeqCst) {
17038 return;
17039 }
17040 let Some(registry) = self.registry.upgrade() else {
17041 *self.nested_result.lock() =
17042 Some(Err("concurrent registry is unavailable".to_string()));
17043 return;
17044 };
17045 let agents = [ai_agents_state::ConcurrentAgentRef::Id(
17046 self.child_id.clone(),
17047 )];
17048 let aggregation = ai_agents_state::AggregationConfig {
17049 strategy: ai_agents_state::AggregationStrategy::FirstWins,
17050 synthesizer_llm: None,
17051 synthesizer_prompt: None,
17052 vote: None,
17053 };
17054 let result = crate::orchestration::concurrent(
17055 ®istry,
17056 "nested concurrent response hook call",
17057 &agents,
17058 &aggregation,
17059 None,
17060 Some(1),
17061 None,
17062 ai_agents_state::PartialFailureAction::Abort,
17063 None,
17064 )
17065 .await
17066 .map(|result| result.response.content)
17067 .map_err(|error| error.to_string());
17068 *self.nested_result.lock() = Some(result);
17069 }
17070 }
17071
17072 #[async_trait]
17073 impl AgentHooks for ToolLifecycleRecordingHooks {
17074 async fn on_tool_start(&self, tool: &str, _args: &Value) {
17075 self.events.lock().push(format!("start:{tool}"));
17076 }
17077
17078 async fn on_tool_complete(&self, tool: &str, result: &ToolResult, _duration_ms: u64) {
17079 self.events
17080 .lock()
17081 .push(format!("complete:{tool}:{}", result.success));
17082 }
17083
17084 async fn on_tool_execution_record(&self, record: &ToolExecutionRecord) {
17085 self.events.lock().push(format!(
17086 "record:{}:{}",
17087 record.canonical_id, record.executed
17088 ));
17089 self.records.lock().push(record.clone());
17090 }
17091
17092 async fn on_error(&self, _error: &AgentError) {
17094 self.events.lock().push("error".to_string());
17095 }
17096 }
17097
17098 struct ApprovalRecordingHooks {
17099 events: parking_lot::Mutex<Vec<String>>,
17100 }
17101
17102 impl ApprovalRecordingHooks {
17103 fn new() -> Self {
17104 Self {
17105 events: parking_lot::Mutex::new(Vec::new()),
17106 }
17107 }
17108
17109 fn events(&self) -> Vec<String> {
17110 self.events.lock().clone()
17111 }
17112 }
17113
17114 #[async_trait]
17115 impl AgentHooks for ApprovalRecordingHooks {
17116 async fn on_approval_result(&self, request_id: &str, result: &ApprovalResult) {
17117 self.events.lock().push(format!(
17118 "raw:{}:{}",
17119 request_id,
17120 approval_result_name(result)
17121 ));
17122 }
17123
17124 async fn on_approval_resolved(
17125 &self,
17126 request: &ai_agents_hitl::ApprovalRequest,
17127 raw_result: &ApprovalResult,
17128 outcome: &ApprovalResolvedOutcome,
17129 ) {
17130 self.events.lock().push(format!(
17131 "resolved:{}:{}:{}",
17132 request.id,
17133 approval_result_name(raw_result),
17134 approval_outcome_name(outcome)
17135 ));
17136 }
17137 }
17138
17139 fn approval_result_name(result: &ApprovalResult) -> &'static str {
17140 match result {
17141 ApprovalResult::Approved => "approved",
17142 ApprovalResult::Rejected { .. } => "rejected",
17143 ApprovalResult::Modified { .. } => "modified",
17144 ApprovalResult::Timeout => "timeout",
17145 }
17146 }
17147
17148 fn approval_outcome_name(outcome: &ApprovalResolvedOutcome) -> &'static str {
17149 match outcome {
17150 ApprovalResolvedOutcome::Approved => "approved",
17151 ApprovalResolvedOutcome::Rejected { .. } => "rejected",
17152 ApprovalResolvedOutcome::Modified { .. } => "modified",
17153 ApprovalResolvedOutcome::Error { .. } => "error",
17154 }
17155 }
17156
17157 fn assert_correlated_approval_events(
17158 events: &[String],
17159 raw_status: &str,
17160 outcome_status: &str,
17161 ) {
17162 assert_eq!(events.len(), 2);
17163 let raw: Vec<_> = events[0].split(':').collect();
17164 let resolved: Vec<_> = events[1].split(':').collect();
17165 assert_eq!(raw[0], "raw");
17166 assert_eq!(resolved[0], "resolved");
17167 assert_eq!(raw[1], resolved[1]);
17168 assert_eq!(raw[2], raw_status);
17169 assert_eq!(resolved[2], raw_status);
17170 assert_eq!(resolved[3], outcome_status);
17171 }
17172
17173 fn approval_security_config(policy_enabled: bool) -> ToolSecurityConfig {
17174 let mut security = ToolSecurityConfig {
17175 enabled: true,
17176 fail_closed: true,
17177 ..Default::default()
17178 };
17179 let policy = ai_agents_tools::ToolPolicyConfig {
17180 enabled: policy_enabled,
17181 write_paths: vec![".".to_string()],
17182 require_confirmation: true,
17183 ..Default::default()
17184 };
17185 security.tools.insert("locked_write".to_string(), policy);
17186 security
17187 }
17188
17189 struct MutationTestWorkspace {
17190 root: std::path::PathBuf,
17191 }
17192
17193 impl MutationTestWorkspace {
17194 fn new() -> Self {
17195 let root = std::env::temp_dir().join(format!(
17196 "ai-agents-runtime-mutation-{}",
17197 uuid::Uuid::new_v4()
17198 ));
17199 std::fs::create_dir_all(&root).unwrap();
17200 Self { root }
17201 }
17202 }
17203
17204 impl Drop for MutationTestWorkspace {
17205 fn drop(&mut self) {
17206 let _ = std::fs::remove_dir_all(&self.root);
17207 }
17208 }
17209
17210 async fn wait_for_resource_lock_strong_count(locks: &ToolResourceLocks, minimum: usize) {
17211 tokio::time::timeout(std::time::Duration::from_secs(2), async {
17212 loop {
17213 let strong_count = locks
17214 .read()
17215 .get("path-mutation:global")
17216 .map_or(0, |lock| lock.strong_count());
17217 if strong_count >= minimum {
17218 break;
17219 }
17220 tokio::task::yield_now().await;
17221 }
17222 })
17223 .await
17224 .expect("path mutation call did not reach the shared lock");
17225 }
17226
17227 async fn assert_path_mutation_pair_serialized(
17228 first_id: &'static str,
17229 first_fields: Vec<ai_agents_core::PathPolicyBinding>,
17230 first_args: Value,
17231 second_id: &'static str,
17232 second_fields: Vec<ai_agents_core::PathPolicyBinding>,
17233 second_args: Value,
17234 ) {
17235 let locks = new_tool_resource_locks();
17236 let first_gate = PathMutationGate::new();
17237 let second_gate = PathMutationGate::new();
17238 second_gate.release();
17239 let agent = Arc::new(
17240 AgentBuilder::new()
17241 .system_prompt("Test global path mutation locking.")
17242 .llm(Arc::new(mock_with_response("done")))
17243 .tool(Arc::new(BlockingPathMutationTool {
17244 id: first_id,
17245 path_fields: first_fields,
17246 gate: first_gate.clone(),
17247 }))
17248 .tool(Arc::new(BlockingPathMutationTool {
17249 id: second_id,
17250 path_fields: second_fields,
17251 gate: second_gate.clone(),
17252 }))
17253 .build()
17254 .unwrap()
17255 .with_shared_resource_locks(Arc::clone(&locks)),
17256 );
17257
17258 let first = {
17259 let agent = Arc::clone(&agent);
17260 tokio::spawn(async move {
17261 agent
17262 .invoke_tool(ToolExecutionRequest::new(
17263 format!("{}-first", first_id),
17264 first_id,
17265 first_args,
17266 ToolCallSource::Manual,
17267 ))
17268 .await
17269 .unwrap()
17270 })
17271 };
17272 first_gate.wait_until_entered().await;
17273
17274 let second = {
17275 let agent = Arc::clone(&agent);
17276 tokio::spawn(async move {
17277 agent
17278 .invoke_tool(ToolExecutionRequest::new(
17279 format!("{}-second", second_id),
17280 second_id,
17281 second_args,
17282 ToolCallSource::Manual,
17283 ))
17284 .await
17285 .unwrap()
17286 })
17287 };
17288 wait_for_resource_lock_strong_count(&locks, 2).await;
17289 assert!(!second_gate.entered.load(Ordering::SeqCst));
17290 assert!(!second.is_finished());
17291
17292 first_gate.release();
17293 let (first, second) = tokio::time::timeout(std::time::Duration::from_secs(2), async {
17294 tokio::join!(first, second)
17295 })
17296 .await
17297 .expect("serialized path mutation calls did not finish");
17298 assert!(first.unwrap().success);
17299 assert!(second.unwrap().success);
17300 assert!(second_gate.entered.load(Ordering::SeqCst));
17301 assert!(locks.read().is_empty());
17302 }
17303
17304 #[derive(Clone, Copy)]
17305 enum MutationDenial {
17306 Policy,
17307 Approval,
17308 }
17309
17310 fn mutation_denial_security_config(
17311 tool_id: &str,
17312 workspace: &std::path::Path,
17313 denial: MutationDenial,
17314 ) -> ToolSecurityConfig {
17315 let workspace = workspace.to_string_lossy().into_owned();
17316 let mut policy = ai_agents_tools::ToolPolicyConfig {
17317 read_paths: vec![workspace.clone()],
17318 write_paths: vec![workspace.clone()],
17319 ..Default::default()
17320 };
17321 match denial {
17322 MutationDenial::Policy => policy.blocked_paths = vec![workspace],
17323 MutationDenial::Approval => policy.require_confirmation = true,
17324 }
17325
17326 let mut security = ToolSecurityConfig {
17327 enabled: true,
17328 fail_closed: true,
17329 ..Default::default()
17330 };
17331 security.tools.insert(tool_id.to_string(), policy);
17332 security
17333 }
17334
17335 async fn assert_path_mutation_denied(tool: Arc<dyn Tool>, denial: MutationDenial) {
17336 let workspace = MutationTestWorkspace::new();
17337 let tool_id = tool.id().to_string();
17338 let preserved = workspace.root.join(format!("{}-preserved.txt", tool_id));
17339 let destination = workspace.root.join(format!("{}-destination.txt", tool_id));
17340 std::fs::write(&preserved, "preserved").unwrap();
17341 let arguments = match tool_id.as_str() {
17342 "copy_path" | "move_path" => serde_json::json!({
17343 "source_path": preserved.to_string_lossy(),
17344 "destination_path": destination.to_string_lossy(),
17345 "dry_run": false
17346 }),
17347 "delete_path" => serde_json::json!({
17348 "path": preserved.to_string_lossy(),
17349 "recursive": false,
17350 "dry_run": false
17351 }),
17352 _ => panic!("unsupported mutation tool: {}", tool_id),
17353 };
17354 let security = mutation_denial_security_config(&tool_id, &workspace.root, denial);
17355 let builder = AgentBuilder::new()
17356 .system_prompt("Test mutation denial.")
17357 .llm(Arc::new(mock_with_response("done")))
17358 .tool(tool)
17359 .tool_security(ToolSecurityEngine::new(security));
17360 let builder = match denial {
17361 MutationDenial::Policy => builder,
17362 MutationDenial::Approval => builder
17363 .hitl_engine(HITLEngine::new(ai_agents_hitl::HITLConfig::default()))
17364 .approval_handler(Arc::new(RejectAllHandler::new())),
17365 };
17366 let agent = builder.build().unwrap();
17367
17368 let record = agent
17369 .invoke_tool(ToolExecutionRequest::new(
17370 format!("{}-denied", tool_id),
17371 tool_id.clone(),
17372 arguments,
17373 ToolCallSource::Manual,
17374 ))
17375 .await
17376 .unwrap();
17377
17378 assert!(!record.executed, "{} must not be invoked", tool_id);
17379 assert!(!record.success);
17380 match denial {
17381 MutationDenial::Policy => {
17382 assert_eq!(record.policy.outcome, PermissionOutcome::Deny);
17383 assert!(record.approval.as_ref().is_some_and(|approval| matches!(
17384 &approval.status,
17385 ToolApprovalStatus::NotRequired
17386 )));
17387 }
17388 MutationDenial::Approval => {
17389 assert_eq!(record.policy.outcome, PermissionOutcome::RequiresApproval);
17390 assert!(record.approval.as_ref().is_some_and(|approval| matches!(
17391 &approval.status,
17392 ToolApprovalStatus::Rejected
17393 )));
17394 }
17395 }
17396 assert_eq!(std::fs::read_to_string(&preserved).unwrap(), "preserved");
17397 assert!(!destination.exists());
17398 }
17399
17400 fn recovery_manager_with_fallbacks(
17401 fallbacks: impl IntoIterator<Item = (String, String)>,
17402 ) -> RecoveryManager {
17403 use ai_agents_recovery::{ErrorRecoveryConfig, ToolRecoveryConfig, ToolRetryConfig};
17404
17405 let per_tool = fallbacks
17406 .into_iter()
17407 .map(|(tool, fallback_tool)| {
17408 (
17409 tool,
17410 ToolRetryConfig {
17411 max_retries: 0,
17412 timeout_ms: Some(1_000),
17413 on_failure: ToolFailureAction::Fallback { fallback_tool },
17414 },
17415 )
17416 })
17417 .collect();
17418 RecoveryManager::new(ErrorRecoveryConfig {
17419 tools: ToolRecoveryConfig {
17420 per_tool,
17421 ..Default::default()
17422 },
17423 ..Default::default()
17424 })
17425 }
17426
17427 fn approval_check() -> HITLCheckResult {
17428 HITLCheckResult::required(
17429 ApprovalTrigger::tool("test", serde_json::json!({})),
17430 HashMap::new(),
17431 "Approve?",
17432 None,
17433 )
17434 }
17435
17436 fn agent_with_approval_result(
17437 raw_result: ApprovalResult,
17438 timeout_action: TimeoutAction,
17439 hooks: Arc<ApprovalRecordingHooks>,
17440 ) -> RuntimeAgent {
17441 use ai_agents_hitl::{CallbackHandler, HITLConfig};
17442
17443 let config = HITLConfig {
17444 on_timeout: timeout_action,
17445 ..Default::default()
17446 };
17447 let handler = CallbackHandler::new(move |_| raw_result.clone());
17448 AgentBuilder::new()
17449 .system_prompt("Test HITL hooks.")
17450 .llm(Arc::new(mock_with_response("done")))
17451 .build()
17452 .unwrap()
17453 .with_hooks(hooks)
17454 .with_hitl(HITLEngine::new(config), Arc::new(handler))
17455 }
17456
17457 #[tokio::test]
17458 async fn approval_hooks_expose_direct_effective_decisions_after_raw_results() {
17459 let cases = vec![
17460 (ApprovalResult::Approved, "approved"),
17461 (
17462 ApprovalResult::Rejected {
17463 reason: Some("denied".to_string()),
17464 },
17465 "rejected",
17466 ),
17467 (
17468 ApprovalResult::Modified {
17469 changes: HashMap::from([("value".to_string(), serde_json::json!(2))]),
17470 },
17471 "modified",
17472 ),
17473 ];
17474
17475 for (raw_result, expected) in cases {
17476 let hooks = Arc::new(ApprovalRecordingHooks::new());
17477 let agent =
17478 agent_with_approval_result(raw_result, TimeoutAction::Reject, hooks.clone());
17479
17480 let result = agent.request_hitl_approval(approval_check()).await.unwrap();
17481
17482 assert_eq!(approval_result_name(&result), expected);
17483 assert_correlated_approval_events(&hooks.events(), expected, expected);
17484 }
17485 }
17486
17487 #[tokio::test]
17488 async fn approval_hooks_expose_timeout_policy_decisions() {
17489 for (timeout_action, expected) in [
17490 (TimeoutAction::Approve, "approved"),
17491 (TimeoutAction::Reject, "rejected"),
17492 ] {
17493 let hooks = Arc::new(ApprovalRecordingHooks::new());
17494 let agent =
17495 agent_with_approval_result(ApprovalResult::Timeout, timeout_action, hooks.clone());
17496
17497 let result = agent.request_hitl_approval(approval_check()).await.unwrap();
17498
17499 assert_eq!(approval_result_name(&result), expected);
17500 assert_correlated_approval_events(&hooks.events(), "timeout", expected);
17501 }
17502 }
17503
17504 #[tokio::test]
17505 async fn timeout_error_fires_correlated_resolved_error_before_returning() {
17506 let hooks = Arc::new(ApprovalRecordingHooks::new());
17507 let agent = agent_with_approval_result(
17508 ApprovalResult::Timeout,
17509 TimeoutAction::Error,
17510 hooks.clone(),
17511 );
17512
17513 let error = agent
17514 .request_hitl_approval(approval_check())
17515 .await
17516 .unwrap_err();
17517
17518 assert!(error.to_string().contains("HITL approval timeout"));
17519 assert_correlated_approval_events(&hooks.events(), "timeout", "error");
17520 }
17521
17522 #[tokio::test]
17524 async fn test_integration_yaml_to_chat_basic() {
17525 let mock = mock_with_response("Hello! How can I help you?");
17526 let agent = AgentBuilder::new()
17527 .system_prompt("You are a test assistant.")
17528 .llm(Arc::new(mock))
17529 .build()
17530 .unwrap();
17531
17532 let response = agent.chat("Hi").await.unwrap();
17533 assert!(!response.content.is_empty());
17534 assert_eq!(response.content, "Hello! How can I help you?");
17535 }
17536
17537 #[tokio::test]
17538 async fn stream_events_emit_one_authoritative_final_without_legacy_done() {
17539 let agent = AgentBuilder::new()
17540 .system_prompt("You are a test assistant.")
17541 .llm(Arc::new(mock_with_response(
17542 "Hello from the final response.",
17543 )))
17544 .build()
17545 .unwrap();
17546
17547 let mut stream = agent.chat_stream_events("Hi").await.unwrap();
17548 let mut final_responses = Vec::new();
17549 let mut legacy_done = 0;
17550 while let Some(event) = stream.next().await {
17551 match event {
17552 AgentStreamEvent::Chunk(StreamChunk::Done {}) => legacy_done += 1,
17553 AgentStreamEvent::Chunk(StreamChunk::Error { message }) => {
17554 panic!("unexpected stream error: {message}")
17555 }
17556 AgentStreamEvent::Final(response) => final_responses.push(response),
17557 AgentStreamEvent::Chunk(_) => {}
17558 }
17559 }
17560
17561 assert_eq!(legacy_done, 0);
17562 assert_eq!(final_responses.len(), 1);
17563 let response = final_responses.pop().unwrap();
17564 assert_eq!(response.content, "Hello from the final response.");
17565 assert!(
17566 response
17567 .metadata
17568 .as_ref()
17569 .is_some_and(|metadata| { metadata.contains_key("reasoning") })
17570 );
17571 }
17572
17573 #[tokio::test]
17574 async fn stream_final_content_includes_output_processing_after_provisional_chunks() {
17575 let yaml = r#"
17576name: ProcessedStreamAgent
17577system_prompt: "Answer directly."
17578process:
17579 output:
17580 - type: format
17581 config:
17582 template: "{{ response }} [finalized]"
17583streaming:
17584 enabled: true
17585"#;
17586 let agent = AgentBuilder::from_yaml(yaml)
17587 .unwrap()
17588 .llm(Arc::new(mock_with_response("provisional answer")))
17589 .auto_configure_features()
17590 .unwrap()
17591 .build()
17592 .unwrap();
17593
17594 let mut stream = agent.chat_stream_events("Hi").await.unwrap();
17595 let mut provisional = String::new();
17596 let mut final_content = None;
17597 while let Some(event) = stream.next().await {
17598 match event {
17599 AgentStreamEvent::Chunk(StreamChunk::Content { text }) => {
17600 provisional.push_str(&text)
17601 }
17602 AgentStreamEvent::Chunk(StreamChunk::Error { message }) => {
17603 panic!("unexpected stream error: {message}")
17604 }
17605 AgentStreamEvent::Final(response) => final_content = Some(response.content),
17606 AgentStreamEvent::Chunk(_) => {}
17607 }
17608 }
17609
17610 assert_eq!(provisional, "provisional answer");
17611 assert_eq!(
17612 final_content.as_deref(),
17613 Some("provisional answer [finalized]")
17614 );
17615 }
17616
17617 #[tokio::test]
17618 async fn stream_events_preserve_tool_progress_and_final_tool_calls() {
17619 let agent = AgentBuilder::new()
17620 .system_prompt("Use the echo tool once, then answer.")
17621 .llm(Arc::new(mock_with_responses(vec![
17622 r#"{"tool":"echo","arguments":{"message":"hello"}}"#,
17623 "Echo completed.",
17624 ])))
17625 .tool(Arc::new(ai_agents_tools::EchoTool::new()))
17626 .build()
17627 .unwrap();
17628
17629 let mut stream = agent.chat_stream_events("echo hello").await.unwrap();
17630 let mut starts = 0;
17631 let mut results = 0;
17632 let mut ends = 0;
17633 let mut final_response = None;
17634 while let Some(event) = stream.next().await {
17635 match event {
17636 AgentStreamEvent::Chunk(StreamChunk::ToolCallStart { name, .. }) => {
17637 assert_eq!(name, "echo");
17638 starts += 1;
17639 }
17640 AgentStreamEvent::Chunk(StreamChunk::ToolResult { name, success, .. }) => {
17641 assert_eq!(name, "echo");
17642 assert!(success);
17643 results += 1;
17644 }
17645 AgentStreamEvent::Chunk(StreamChunk::ToolCallEnd { .. }) => ends += 1,
17646 AgentStreamEvent::Chunk(StreamChunk::Error { message }) => {
17647 panic!("unexpected stream error: {message}")
17648 }
17649 AgentStreamEvent::Final(response) => final_response = Some(response),
17650 AgentStreamEvent::Chunk(_) => {}
17651 }
17652 }
17653
17654 assert_eq!((starts, results, ends), (1, 1, 1));
17655 let response = final_response.expect("tool stream must finalize");
17656 assert_eq!(response.content, "Echo completed.");
17657 assert_eq!(
17658 response.tool_calls.as_ref().map(|calls| calls
17659 .iter()
17660 .map(|call| call.name.as_str())
17661 .collect::<Vec<_>>()),
17662 Some(vec!["echo"])
17663 );
17664 }
17665
17666 #[tokio::test]
17667 async fn legacy_stream_still_emits_one_done_chunk() {
17668 let agent = AgentBuilder::new()
17669 .system_prompt("You are a test assistant.")
17670 .llm(Arc::new(mock_with_response(
17671 "Hello from the legacy stream.",
17672 )))
17673 .build()
17674 .unwrap();
17675
17676 let mut stream = agent.chat_stream("Hi").await.unwrap();
17677 let mut done = 0;
17678 while let Some(chunk) = stream.next().await {
17679 match chunk {
17680 StreamChunk::Done {} => done += 1,
17681 StreamChunk::Error { message } => panic!("unexpected stream error: {message}"),
17682 _ => {}
17683 }
17684 }
17685
17686 assert_eq!(done, 1);
17687 }
17688
17689 #[tokio::test]
17691 async fn test_integration_multi_turn_conversation() {
17692 let mock = mock_with_responses(vec![
17693 "Hello! I'm your assistant.",
17694 "The weather is sunny today.",
17695 "Goodbye!",
17696 ]);
17697 let agent = AgentBuilder::new()
17698 .system_prompt("You are helpful.")
17699 .llm(Arc::new(mock))
17700 .build()
17701 .unwrap();
17702
17703 let r1 = agent.chat("Hi").await.unwrap();
17704 assert_eq!(r1.content, "Hello! I'm your assistant.");
17705
17706 let r2 = agent.chat("What's the weather?").await.unwrap();
17707 assert_eq!(r2.content, "The weather is sunny today.");
17708
17709 let r3 = agent.chat("Bye").await.unwrap();
17710 assert_eq!(r3.content, "Goodbye!");
17711
17712 let messages = agent.memory.get_messages(None).await.unwrap();
17714 assert_eq!(messages.len(), 6);
17716 }
17717
17718 #[test]
17719 fn later_approval_preserves_modified_evidence() {
17720 let arguments = serde_json::json!({"dry_run": true});
17721 let mut record = Some(ToolApprovalRecord {
17722 status: ToolApprovalStatus::Modified,
17723 reason: None,
17724 modified_arguments: Some(arguments.clone()),
17725 });
17726
17727 merge_approved_record(&mut record);
17728
17729 let record = record.unwrap();
17730 assert!(matches!(record.status, ToolApprovalStatus::Modified));
17731 assert_eq!(record.modified_arguments, Some(arguments));
17732 }
17733
17734 #[test]
17735 fn approval_binding_rejects_replaced_tool_implementation() {
17736 let reviewed_tool: Arc<dyn ai_agents_core::Tool> = Arc::new(ContextEchoTool);
17737 let same_tool = Arc::clone(&reviewed_tool);
17738 let replacement_tool: Arc<dyn ai_agents_core::Tool> = Arc::new(ContextEchoTool);
17739 let arguments = serde_json::json!({"path": "."});
17740 let versions = ToolDecisionVersions {
17741 policy: 2,
17742 registry: 3,
17743 runtime_control: 4,
17744 state: Some(5),
17745 };
17746 let binding = ToolApprovalBinding {
17747 canonical_id: "context_echo".to_string(),
17748 arguments: arguments.clone(),
17749 confirmation_required: true,
17750 policy_version: versions.policy,
17751 runtime_control_version: versions.runtime_control,
17752 state_generation: versions.state,
17753 reviewed_tool,
17754 };
17755
17756 assert!(!binding.is_stale("context_echo", &arguments, true, versions, &same_tool,));
17757 assert!(binding.is_stale(
17758 "context_echo",
17759 &arguments,
17760 true,
17761 versions,
17762 &replacement_tool,
17763 ));
17764 }
17765
17766 #[tokio::test]
17767 async fn approved_mutation_to_dry_run_remains_executable() {
17768 use ai_agents_hitl::CallbackHandler;
17769
17770 let handler = CallbackHandler::new(|_| ApprovalResult::Modified {
17771 changes: HashMap::from([("dry_run".to_string(), serde_json::json!(true))]),
17772 });
17773 let agent = AgentBuilder::new()
17774 .system_prompt("Test safer approval modifications.")
17775 .llm(Arc::new(mock_with_response("done")))
17776 .tool(Arc::new(ai_agents_tools::FileWriteTool::new()))
17777 .hitl_engine(HITLEngine::new(ai_agents_hitl::HITLConfig::default()))
17778 .approval_handler(Arc::new(handler))
17779 .build()
17780 .unwrap();
17781
17782 let record = agent
17783 .invoke_tool(ToolExecutionRequest::new(
17784 "approved-dry-run",
17785 "file_write",
17786 serde_json::json!({
17787 "path": "./approval-dry-run.txt",
17788 "content": "not written"
17789 }),
17790 ToolCallSource::Manual,
17791 ))
17792 .await
17793 .unwrap();
17794
17795 assert!(record.executed);
17796 assert!(record.success);
17797 assert_eq!(record.executed_arguments["dry_run"], true);
17798 assert!(matches!(
17799 record.approval.as_ref().map(|approval| &approval.status),
17800 Some(ToolApprovalStatus::Modified)
17801 ));
17802 let output: Value = serde_json::from_str(&record.output).unwrap();
17803 assert_eq!(output["mutation_performed"], false);
17804 }
17805
17806 #[tokio::test]
17808 async fn shared_executor_approval_reaches_web_fetch_transport() {
17809 use ai_agents_hitl::{CallbackHandler, HITLConfig};
17810 use ai_agents_tools::{DomainPolicyConfig, ToolPolicyConfig};
17811
17812 let calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
17813 let tool = WebFetchTool::with_transport_and_resolver(
17814 Arc::new(RuntimeWebFetchTransport {
17815 calls: Arc::clone(&calls),
17816 }),
17817 Arc::new(RuntimeWebFetchResolver),
17818 );
17819 let mut security = ToolSecurityConfig {
17820 enabled: true,
17821 fail_closed: true,
17822 ..Default::default()
17823 };
17824 security.tools.insert(
17825 "web_fetch".to_string(),
17826 ToolPolicyConfig {
17827 domains: DomainPolicyConfig {
17828 requires_approval: vec!["approval.test".to_string()],
17829 ..Default::default()
17830 },
17831 allowed_schemes: vec!["https".to_string()],
17832 allowed_ports: vec![443],
17833 ..Default::default()
17834 },
17835 );
17836 let handler = CallbackHandler::new(|_| ApprovalResult::Approved);
17837 let agent = AgentBuilder::new()
17838 .system_prompt("Test approved web fetch execution.")
17839 .llm(Arc::new(mock_with_response("done")))
17840 .tool(Arc::new(tool))
17841 .tool_security(ToolSecurityEngine::new(security))
17842 .build()
17843 .unwrap()
17844 .with_hitl(HITLEngine::new(HITLConfig::default()), Arc::new(handler));
17845
17846 let record = agent
17847 .invoke_tool(ToolExecutionRequest::new(
17848 "approved-web-fetch",
17849 "web_fetch",
17850 serde_json::json!({
17851 "url": "https://approval.test/page",
17852 "cache_ttl_seconds": 0
17853 }),
17854 ToolCallSource::Manual,
17855 ))
17856 .await
17857 .unwrap();
17858
17859 assert!(record.success);
17860 assert!(
17861 record
17862 .approval
17863 .as_ref()
17864 .is_some_and(|approval| matches!(approval.status, ToolApprovalStatus::Approved))
17865 );
17866 assert_eq!(calls.load(Ordering::SeqCst), 1);
17867 }
17868
17869 #[tokio::test]
17870 async fn context_preserves_requested_and_canonical_identity() {
17871 let mock = mock_with_response("hello");
17872 let mut tools = ai_agents_tools::ToolRegistry::new();
17873 tools.register(Arc::new(ContextEchoTool)).unwrap();
17874
17875 let mut security = ToolSecurityConfig {
17876 enabled: true,
17877 fail_closed: true,
17878 ..Default::default()
17879 };
17880 let mut policy = ai_agents_tools::ToolPolicyConfig {
17881 read_paths: vec![".".to_string()],
17882 max_results: Some(7),
17883 ..Default::default()
17884 };
17885 policy
17886 .config
17887 .insert("backend".to_string(), serde_json::json!("memory"));
17888 security.tools.insert("context_echo".to_string(), policy);
17889
17890 let agent = AgentBuilder::new()
17891 .system_prompt("You are helpful.")
17892 .llm(Arc::new(mock))
17893 .tools(tools)
17894 .tool_security(ToolSecurityEngine::new(security))
17895 .build()
17896 .unwrap();
17897
17898 let record = agent
17899 .invoke_tool(ToolExecutionRequest::new(
17900 "ctx-call",
17901 "Context Echo",
17902 serde_json::json!({"path": ".", "max_results": 99}),
17903 ToolCallSource::Manual,
17904 ))
17905 .await
17906 .unwrap();
17907
17908 assert!(record.success);
17909 assert!(matches!(&record.source, ToolCallSource::Manual));
17910 assert_eq!(record.requested_name, "Context Echo");
17911 assert_eq!(record.canonical_id, "context_echo");
17912 assert_eq!(record.policy.outcome, PermissionOutcome::Allow);
17913 assert_eq!(record.executed_arguments["max_results"], 7);
17914 let output: Value = serde_json::from_str(&record.output).unwrap();
17915 assert_eq!(output["requested_name"], "Context Echo");
17916 assert_eq!(output["canonical_id"], "context_echo");
17917 assert_eq!(output["max_results"], 7);
17918 assert_eq!(output["custom_config"]["backend"], "memory");
17919 assert!(record.metadata.contains_key("effective_limits"));
17920 assert!(record.metadata.contains_key("policy_snapshot"));
17921 }
17922
17923 #[tokio::test]
17924 async fn test_runtime_control_cancels_active_tool_call() {
17925 let mock = mock_with_response("hello");
17926 let agent = Arc::new(
17927 AgentBuilder::new()
17928 .system_prompt("You are helpful.")
17929 .llm(Arc::new(mock))
17930 .tool(Arc::new(SlowTool))
17931 .build()
17932 .unwrap(),
17933 );
17934 let control = agent.runtime_control();
17935 let running_agent = Arc::clone(&agent);
17936 let handle = tokio::spawn(async move {
17937 running_agent
17938 .invoke_tool(ToolExecutionRequest::new(
17939 "slow-call",
17940 "slow",
17941 serde_json::json!({}),
17942 ToolCallSource::Manual,
17943 ))
17944 .await
17945 .unwrap()
17946 });
17947
17948 tokio::time::sleep(std::time::Duration::from_millis(100)).await;
17949 control.cancel_all();
17950 let record = handle.await.unwrap();
17951
17952 assert!(record.executed);
17953 assert!(record.cancelled);
17954 assert!(!record.success);
17955 assert!(record.cancellation_reason.is_some());
17956 }
17957
17958 #[tokio::test]
17960 async fn cancelled_tool_does_not_enter_fallback() {
17961 let fallback_calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
17962 let agent = Arc::new(
17963 AgentBuilder::new()
17964 .system_prompt("Test cancellation before fallback.")
17965 .llm(Arc::new(mock_with_response("done")))
17966 .tool(Arc::new(SlowTool))
17967 .tool(Arc::new(RecoveryTestTool {
17968 id: "fallback".to_string(),
17969 succeeds: true,
17970 calls: Arc::clone(&fallback_calls),
17971 max_output_chars: None,
17972 }))
17973 .recovery_manager(recovery_manager_with_fallbacks([(
17974 "slow".to_string(),
17975 "fallback".to_string(),
17976 )]))
17977 .build()
17978 .unwrap(),
17979 );
17980 let control = agent.runtime_control();
17981 let running_agent = Arc::clone(&agent);
17982 let handle = tokio::spawn(async move {
17983 running_agent
17984 .invoke_tool(ToolExecutionRequest::new(
17985 "cancelled-fallback-call",
17986 "slow",
17987 serde_json::json!({}),
17988 ToolCallSource::Manual,
17989 ))
17990 .await
17991 .unwrap()
17992 });
17993
17994 tokio::time::sleep(Duration::from_millis(100)).await;
17995 control.cancel_all();
17996 let record = handle.await.unwrap();
17997
17998 assert!(record.executed);
17999 assert!(record.cancelled);
18000 assert!(!record.success);
18001 assert_eq!(record.canonical_id, "slow");
18002 assert_eq!(fallback_calls.load(Ordering::SeqCst), 0);
18003 assert_eq!(agent.tool_call_history().len(), 1);
18004 }
18005
18006 #[tokio::test]
18007 async fn non_idempotent_tool_calls_are_not_retried() {
18008 use ai_agents_recovery::{ErrorRecoveryConfig, ToolRecoveryConfig, ToolRetryConfig};
18009
18010 let mock = mock_with_response("hello");
18011 let calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
18012 let agent = AgentBuilder::new()
18013 .system_prompt("You are helpful.")
18014 .llm(Arc::new(mock))
18015 .tool(Arc::new(FlakyWriteTool {
18016 calls: Arc::clone(&calls),
18017 }))
18018 .recovery_manager(RecoveryManager::new(ErrorRecoveryConfig {
18019 tools: ToolRecoveryConfig {
18020 default: ToolRetryConfig {
18021 max_retries: 2,
18022 ..Default::default()
18023 },
18024 ..Default::default()
18025 },
18026 ..Default::default()
18027 }))
18028 .build()
18029 .unwrap();
18030
18031 let record = agent
18032 .invoke_tool(ToolExecutionRequest::new(
18033 "flaky-call",
18034 "flaky_write",
18035 serde_json::json!({"path": "./tmp.txt"}),
18036 ToolCallSource::Manual,
18037 ))
18038 .await
18039 .unwrap();
18040
18041 assert!(!record.success);
18042 assert_eq!(calls.load(Ordering::SeqCst), 1);
18043 }
18044
18045 #[tokio::test]
18046 async fn safely_retryable_tool_receives_a_fresh_deadline_per_attempt() {
18047 use ai_agents_recovery::{ErrorRecoveryConfig, ToolRecoveryConfig, ToolRetryConfig};
18048
18049 let calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
18050 let deadlines = Arc::new(parking_lot::Mutex::new(Vec::new()));
18051 let remaining_ms = Arc::new(parking_lot::Mutex::new(Vec::new()));
18052 let agent = AgentBuilder::new()
18053 .system_prompt("Test retry deadlines.")
18054 .llm(Arc::new(mock_with_response("done")))
18055 .tool(Arc::new(RetryDeadlineTool {
18056 calls: Arc::clone(&calls),
18057 deadlines: Arc::clone(&deadlines),
18058 remaining_ms: Arc::clone(&remaining_ms),
18059 }))
18060 .recovery_manager(RecoveryManager::new(ErrorRecoveryConfig {
18061 tools: ToolRecoveryConfig {
18062 per_tool: HashMap::from([(
18063 "retry_deadline".to_string(),
18064 ToolRetryConfig {
18065 max_retries: 1,
18066 ..Default::default()
18067 },
18068 )]),
18069 ..Default::default()
18070 },
18071 ..Default::default()
18072 }))
18073 .build()
18074 .unwrap();
18075
18076 let record = agent
18077 .invoke_tool(ToolExecutionRequest::new(
18078 "retry-deadline-call",
18079 "retry_deadline",
18080 serde_json::json!({}),
18081 ToolCallSource::Manual,
18082 ))
18083 .await
18084 .unwrap();
18085
18086 assert!(record.executed);
18087 assert!(record.success);
18088 assert_eq!(calls.load(Ordering::SeqCst), 2);
18089 let deadlines = deadlines.lock();
18090 assert_eq!(deadlines.len(), 2);
18091 assert!(
18092 deadlines[1] > deadlines[0],
18093 "retry inherited the first invocation deadline"
18094 );
18095 let remaining_ms = remaining_ms.lock();
18096 assert_eq!(remaining_ms.len(), 2);
18097 assert!(
18098 remaining_ms
18099 .iter()
18100 .all(|remaining| (800..=1_000).contains(remaining))
18101 );
18102 }
18103
18104 #[tokio::test]
18106 async fn call_classification_timeout_controls_deadline_and_timer() {
18107 let calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
18108 let remaining_ms = Arc::new(parking_lot::Mutex::new(Vec::new()));
18109 let agent = AgentBuilder::new()
18110 .system_prompt("Test call-level timeout.")
18111 .llm(Arc::new(mock_with_response("done")))
18112 .tool(Arc::new(ClassifiedTimeoutTool {
18113 id: "classified_timeout",
18114 calls: Arc::clone(&calls),
18115 timeout_ms: 100,
18116 sleep_ms: 150,
18117 requires_approval: false,
18118 remaining_ms: Arc::clone(&remaining_ms),
18119 }))
18120 .build()
18121 .unwrap();
18122
18123 let started = Instant::now();
18124 let record = agent
18125 .invoke_tool(ToolExecutionRequest::new(
18126 "classified-timeout-call",
18127 "classified_timeout",
18128 serde_json::json!({}),
18129 ToolCallSource::Manual,
18130 ))
18131 .await
18132 .unwrap();
18133
18134 assert!(record.executed);
18135 assert!(record.timed_out);
18136 assert!(!record.success);
18137 assert_eq!(calls.load(Ordering::SeqCst), 1);
18138 assert!(started.elapsed() < Duration::from_secs(1));
18139 let remaining_ms = remaining_ms.lock();
18140 assert_eq!(remaining_ms.len(), 1);
18141 assert!((1..=100).contains(&remaining_ms[0]));
18142 }
18143
18144 #[tokio::test]
18146 async fn recovery_timeout_only_lowers_call_and_policy_timeouts() {
18147 use ai_agents_recovery::{ErrorRecoveryConfig, ToolRecoveryConfig, ToolRetryConfig};
18148
18149 let calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
18150 let remaining_ms = Arc::new(parking_lot::Mutex::new(Vec::new()));
18151 let agent = AgentBuilder::new()
18152 .system_prompt("Test recovery timeout.")
18153 .llm(Arc::new(mock_with_response("done")))
18154 .tool(Arc::new(ClassifiedTimeoutTool {
18155 id: "recovery_timeout",
18156 calls: Arc::clone(&calls),
18157 timeout_ms: 1_000,
18158 sleep_ms: 150,
18159 requires_approval: false,
18160 remaining_ms: Arc::clone(&remaining_ms),
18161 }))
18162 .recovery_manager(RecoveryManager::new(ErrorRecoveryConfig {
18163 tools: ToolRecoveryConfig {
18164 per_tool: HashMap::from([(
18165 "recovery_timeout".to_string(),
18166 ToolRetryConfig {
18167 timeout_ms: Some(100),
18168 ..Default::default()
18169 },
18170 )]),
18171 ..Default::default()
18172 },
18173 ..Default::default()
18174 }))
18175 .build()
18176 .unwrap();
18177
18178 let started = Instant::now();
18179 let record = agent
18180 .invoke_tool(ToolExecutionRequest::new(
18181 "recovery-timeout-call",
18182 "recovery_timeout",
18183 serde_json::json!({}),
18184 ToolCallSource::Manual,
18185 ))
18186 .await
18187 .unwrap();
18188
18189 assert!(record.executed);
18190 assert!(record.timed_out);
18191 assert!(!record.success);
18192 assert_eq!(calls.load(Ordering::SeqCst), 1);
18193 assert!(started.elapsed() < Duration::from_secs(1));
18194 assert_eq!(record.metadata["effective_limits"]["timeout_ms"], 100);
18195 let remaining_ms = remaining_ms.lock();
18196 assert_eq!(remaining_ms.len(), 1);
18197 assert!((1..=100).contains(&remaining_ms[0]));
18198 }
18199
18200 #[tokio::test]
18202 async fn recovery_default_timeout_controls_deadline_and_timer() {
18203 use ai_agents_recovery::{ErrorRecoveryConfig, ToolRecoveryConfig, ToolRetryConfig};
18204
18205 let calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
18206 let remaining_ms = Arc::new(parking_lot::Mutex::new(Vec::new()));
18207 let agent = AgentBuilder::new()
18208 .system_prompt("Test default recovery timeout.")
18209 .llm(Arc::new(mock_with_response("done")))
18210 .tool(Arc::new(ClassifiedTimeoutTool {
18211 id: "default_recovery_timeout",
18212 calls: Arc::clone(&calls),
18213 timeout_ms: 1_000,
18214 sleep_ms: 150,
18215 requires_approval: false,
18216 remaining_ms: Arc::clone(&remaining_ms),
18217 }))
18218 .recovery_manager(RecoveryManager::new(ErrorRecoveryConfig {
18219 tools: ToolRecoveryConfig {
18220 default: ToolRetryConfig {
18221 timeout_ms: Some(100),
18222 ..Default::default()
18223 },
18224 ..Default::default()
18225 },
18226 ..Default::default()
18227 }))
18228 .build()
18229 .unwrap();
18230
18231 let started = Instant::now();
18232 let record = agent
18233 .invoke_tool(ToolExecutionRequest::new(
18234 "default-recovery-timeout-call",
18235 "default_recovery_timeout",
18236 serde_json::json!({}),
18237 ToolCallSource::Manual,
18238 ))
18239 .await
18240 .unwrap();
18241
18242 assert!(record.executed);
18243 assert!(record.timed_out);
18244 assert!(!record.success);
18245 assert_eq!(calls.load(Ordering::SeqCst), 1);
18246 assert!(started.elapsed() < Duration::from_secs(1));
18247 assert_eq!(record.metadata["effective_limits"]["timeout_ms"], 100);
18248 let remaining_ms = remaining_ms.lock();
18249 assert_eq!(remaining_ms.len(), 1);
18250 assert!((1..=100).contains(&remaining_ms[0]));
18251 }
18252
18253 #[test]
18255 fn recovery_timeout_cannot_widen_security_baseline() {
18256 let security_engine = ToolSecurityEngine::new(ToolSecurityConfig {
18257 default_timeout_ms: 100,
18258 ..Default::default()
18259 });
18260 let safety = ToolSafetyMetadata::compute();
18261 let mut classification = ToolCallClassification::from_metadata(&safety);
18262 classification.timeout_ms = Some(500);
18263
18264 let (limits, timeout) = RuntimeAgent::effective_tool_limits(
18265 &security_engine,
18266 "recovery_cannot_widen",
18267 &safety,
18268 &classification,
18269 Some(1_000),
18270 )
18271 .unwrap();
18272
18273 assert_eq!(limits.timeout_ms, Some(100));
18274 assert_eq!(timeout.timer, Duration::from_millis(100));
18275 }
18276
18277 #[tokio::test]
18279 async fn invalid_call_timeout_stops_before_approval_or_tool_invocation() {
18280 let tool_calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
18281 let approval_calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
18282 let remaining_ms = Arc::new(parking_lot::Mutex::new(Vec::new()));
18283 let mut security = ToolSecurityConfig {
18284 enabled: true,
18285 ..Default::default()
18286 };
18287 security.tools.insert(
18288 "invalid_call_timeout".to_string(),
18289 ai_agents_tools::ToolPolicyConfig {
18290 require_confirmation: true,
18291 ..Default::default()
18292 },
18293 );
18294 let agent = AgentBuilder::new()
18295 .system_prompt("Test invalid call timeout.")
18296 .llm(Arc::new(mock_with_response("done")))
18297 .tool(Arc::new(ClassifiedTimeoutTool {
18298 id: "invalid_call_timeout",
18299 calls: Arc::clone(&tool_calls),
18300 timeout_ms: u64::MAX,
18301 sleep_ms: 0,
18302 requires_approval: false,
18303 remaining_ms,
18304 }))
18305 .tool_security(ToolSecurityEngine::new(security))
18306 .hitl_engine(HITLEngine::new(ai_agents_hitl::HITLConfig::default()))
18307 .approval_handler(Arc::new(CountingApprovalHandler {
18308 calls: Arc::clone(&approval_calls),
18309 }))
18310 .build()
18311 .unwrap();
18312
18313 let error = agent
18314 .invoke_tool(ToolExecutionRequest::new(
18315 "invalid-call-timeout",
18316 "invalid_call_timeout",
18317 serde_json::json!({}),
18318 ToolCallSource::Manual,
18319 ))
18320 .await
18321 .unwrap_err();
18322
18323 assert!(error.to_string().contains(
18324 "effective tool timeout_ms must be no greater than 3153600000000000 milliseconds"
18325 ));
18326 assert_eq!(approval_calls.load(Ordering::SeqCst), 0);
18327 assert_eq!(tool_calls.load(Ordering::SeqCst), 0);
18328 }
18329
18330 #[tokio::test]
18332 async fn invalid_modified_call_timeout_stops_before_lock_or_invocation() {
18333 use ai_agents_hitl::CallbackHandler;
18334
18335 let blocker_gate = PathMutationGate::new();
18336 let tool_calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
18337 let hooks = Arc::new(ToolLifecycleRecordingHooks::new());
18338 let handler = CallbackHandler::new(|_| ApprovalResult::Modified {
18339 changes: HashMap::from([("invalid_timeout".to_string(), Value::Bool(true))]),
18340 });
18341 let agent = Arc::new(
18342 AgentBuilder::new()
18343 .system_prompt("Test final call timeout validation.")
18344 .llm(Arc::new(mock_with_response("done")))
18345 .tool(Arc::new(BlockingPathMutationTool {
18346 id: "timeout_lock_blocker",
18347 path_fields: vec![ai_agents_core::PathPolicyBinding::write("path")],
18348 gate: blocker_gate.clone(),
18349 }))
18350 .tool(Arc::new(ApprovalModifiedTimeoutTool {
18351 calls: Arc::clone(&tool_calls),
18352 }))
18353 .hitl_engine(HITLEngine::new(ai_agents_hitl::HITLConfig::default()))
18354 .approval_handler(Arc::new(handler))
18355 .hooks(hooks.clone())
18356 .build()
18357 .unwrap(),
18358 );
18359 let blocking_agent = Arc::clone(&agent);
18360 let blocker = tokio::spawn(async move {
18361 blocking_agent
18362 .invoke_tool(ToolExecutionRequest::new(
18363 "timeout-lock-blocker",
18364 "timeout_lock_blocker",
18365 serde_json::json!({"path": "./shared-timeout.txt"}),
18366 ToolCallSource::Manual,
18367 ))
18368 .await
18369 .unwrap()
18370 });
18371 blocker_gate.wait_until_entered().await;
18372
18373 let record = tokio::time::timeout(
18374 Duration::from_millis(500),
18375 agent.invoke_tool(ToolExecutionRequest::new(
18376 "invalid-modified-timeout",
18377 "approval_modified_timeout",
18378 serde_json::json!({
18379 "path": "./shared-timeout.txt",
18380 "invalid_timeout": false
18381 }),
18382 ToolCallSource::Manual,
18383 )),
18384 )
18385 .await
18386 .expect("final timeout validation must not wait for the held path lock")
18387 .unwrap();
18388
18389 blocker_gate.release();
18390 assert!(blocker.await.unwrap().success);
18391 assert!(!record.executed);
18392 assert!(!record.success);
18393 assert_eq!(record.policy.outcome, PermissionOutcome::Deny);
18394 assert!(record.output.contains(
18395 "effective tool timeout_ms must be no greater than 3153600000000000 milliseconds"
18396 ));
18397 assert_eq!(tool_calls.load(Ordering::SeqCst), 0);
18398 let invalid_request_events = hooks
18399 .events()
18400 .into_iter()
18401 .filter(|event| event.contains("approval_modified_timeout") || event == "error")
18402 .collect::<Vec<_>>();
18403 assert_eq!(
18404 invalid_request_events,
18405 vec![
18406 "start:approval_modified_timeout",
18407 "complete:approval_modified_timeout:false",
18408 "record:approval_modified_timeout:false",
18409 "error"
18410 ]
18411 );
18412 }
18413
18414 #[tokio::test]
18415 async fn side_effecting_tools_are_serialized_per_resource() {
18416 let mock = mock_with_response("hello");
18417 let active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
18418 let max_active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
18419 let agent = Arc::new(
18420 AgentBuilder::new()
18421 .system_prompt("You are helpful.")
18422 .llm(Arc::new(mock))
18423 .tool(Arc::new(LockedWriteTool {
18424 active: Arc::clone(&active),
18425 max_active: Arc::clone(&max_active),
18426 }))
18427 .build()
18428 .unwrap(),
18429 );
18430
18431 let left = {
18432 let agent = Arc::clone(&agent);
18433 tokio::spawn(async move {
18434 agent
18435 .invoke_tool(ToolExecutionRequest::new(
18436 "lock-1",
18437 "locked_write",
18438 serde_json::json!({"path": "./same.txt"}),
18439 ToolCallSource::Manual,
18440 ))
18441 .await
18442 .unwrap()
18443 })
18444 };
18445 let right = {
18446 let agent = Arc::clone(&agent);
18447 tokio::spawn(async move {
18448 agent
18449 .invoke_tool(ToolExecutionRequest::new(
18450 "lock-2",
18451 "locked_write",
18452 serde_json::json!({"path": "./same.txt"}),
18453 ToolCallSource::Manual,
18454 ))
18455 .await
18456 .unwrap()
18457 })
18458 };
18459
18460 let left = left.await.unwrap();
18461 let right = right.await.unwrap();
18462 assert!(left.success);
18463 assert!(right.success);
18464 assert_eq!(max_active.load(Ordering::SeqCst), 1);
18465 }
18466
18467 #[tokio::test]
18468 async fn path_resources_use_shared_global_lock_and_cleanup() {
18469 let active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
18470 let max_active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
18471 let bindings = ai_agents_core::ToolPolicyBindings {
18472 path_fields: vec![
18473 ai_agents_core::PathPolicyBinding::read_write("source_path"),
18474 ai_agents_core::PathPolicyBinding::write("destination_path"),
18475 ],
18476 ..Default::default()
18477 };
18478 let classification = ai_agents_core::ToolCallClassification::from_metadata(
18479 &MultiResourceWriteTool {
18480 active: Arc::clone(&active),
18481 max_active: Arc::clone(&max_active),
18482 }
18483 .safety_metadata(),
18484 );
18485 let left_args = serde_json::json!({
18486 "source_path": "./a/../first.txt",
18487 "destination_path": "./second.txt"
18488 });
18489 let right_args = serde_json::json!({
18490 "source_path": "./second.txt",
18491 "destination_path": "./first.txt"
18492 });
18493 let left_keys = tool_resource_lock_keys(
18494 "multi_resource_write",
18495 &left_args,
18496 &bindings,
18497 &classification,
18498 );
18499 let right_keys = tool_resource_lock_keys(
18500 "multi_resource_write",
18501 &right_args,
18502 &bindings,
18503 &classification,
18504 );
18505 assert_eq!(left_keys, right_keys);
18506 assert_eq!(left_keys, vec!["path-mutation:global".to_string()]);
18507
18508 let locks = new_tool_resource_locks();
18509 let build_agent = || {
18510 AgentBuilder::new()
18511 .system_prompt("Test shared resource locks.")
18512 .llm(Arc::new(mock_with_response("done")))
18513 .tool(Arc::new(MultiResourceWriteTool {
18514 active: Arc::clone(&active),
18515 max_active: Arc::clone(&max_active),
18516 }))
18517 .build()
18518 .unwrap()
18519 .with_shared_resource_locks(Arc::clone(&locks))
18520 };
18521 let left_agent = Arc::new(build_agent());
18522 let right_agent = Arc::new(build_agent());
18523 let left = tokio::spawn(async move {
18524 left_agent
18525 .invoke_tool(ToolExecutionRequest::new(
18526 "multi-left",
18527 "multi_resource_write",
18528 left_args,
18529 ToolCallSource::Manual,
18530 ))
18531 .await
18532 .unwrap()
18533 });
18534 let right = tokio::spawn(async move {
18535 right_agent
18536 .invoke_tool(ToolExecutionRequest::new(
18537 "multi-right",
18538 "multi_resource_write",
18539 right_args,
18540 ToolCallSource::Manual,
18541 ))
18542 .await
18543 .unwrap()
18544 });
18545 let (left, right) = tokio::time::timeout(std::time::Duration::from_secs(2), async {
18546 tokio::join!(left, right)
18547 })
18548 .await
18549 .expect("reversed resource acquisition must not deadlock");
18550
18551 assert!(left.unwrap().success);
18552 assert!(right.unwrap().success);
18553 assert_eq!(max_active.load(Ordering::SeqCst), 1);
18554 assert!(locks.read().is_empty());
18555 }
18556
18557 #[tokio::test]
18558 async fn global_path_lock_serializes_copy_destination_with_file_write() {
18559 assert_path_mutation_pair_serialized(
18560 "copy_path",
18561 CopyPathTool::new().policy_bindings().path_fields,
18562 serde_json::json!({
18563 "source_path": "./source.txt",
18564 "destination_path": "./shared.txt"
18565 }),
18566 "file_write",
18567 FileWriteTool::new().policy_bindings().path_fields,
18568 serde_json::json!({"path": "./shared.txt"}),
18569 )
18570 .await;
18571 }
18572
18573 #[tokio::test]
18574 async fn parent_and_spawned_runtime_share_global_path_lock() {
18575 let workspace = MutationTestWorkspace::new();
18576 let destination = workspace.root.join("spawned.txt");
18577 let parent_gate = PathMutationGate::new();
18578 let parent = Arc::new(
18579 AgentBuilder::from_yaml(
18580 r#"
18581name: LockParent
18582system_prompt: parent
18583llm:
18584 default: default
18585tools:
18586 - parent_path_write
18587spawner:
18588 shared_llms: true
18589"#,
18590 )
18591 .unwrap()
18592 .llm(Arc::new(mock_with_response("done")))
18593 .auto_configure_spawner()
18594 .await
18595 .unwrap()
18596 .tool(Arc::new(BlockingPathMutationTool {
18597 id: "parent_path_write",
18598 path_fields: vec![ai_agents_core::PathPolicyBinding::write("path")],
18599 gate: parent_gate.clone(),
18600 }))
18601 .build()
18602 .unwrap(),
18603 );
18604
18605 let mut child_spec = crate::spec::AgentSpec {
18606 name: "LockChild".to_string(),
18607 system_prompt: "child".to_string(),
18608 tools: Some(vec![crate::spec::ToolEntry::Simple(
18609 "file_write".to_string(),
18610 )]),
18611 ..Default::default()
18612 };
18613 child_spec.tool_security.enabled = true;
18614 child_spec.tool_security.fail_closed = true;
18615 let file_write_policy = ai_agents_tools::ToolPolicyConfig {
18616 write_paths: vec![workspace.root.to_string_lossy().into_owned()],
18617 allow_without_confirmation: true,
18618 ..Default::default()
18619 };
18620 child_spec
18621 .tool_security
18622 .tools
18623 .insert("file_write".to_string(), file_write_policy);
18624 let spawned = parent
18625 .spawner()
18626 .unwrap()
18627 .spawn_from_spec(child_spec)
18628 .await
18629 .unwrap();
18630 assert!(Arc::ptr_eq(
18631 &parent.resource_locks,
18632 &spawned.agent.resource_locks
18633 ));
18634 assert!(!Arc::ptr_eq(
18635 &parent.runtime_control,
18636 &spawned.agent.runtime_control
18637 ));
18638
18639 let parent_call = {
18640 let parent = Arc::clone(&parent);
18641 let destination = destination.clone();
18642 tokio::spawn(async move {
18643 parent
18644 .invoke_tool(ToolExecutionRequest::new(
18645 "parent-lock-holder",
18646 "parent_path_write",
18647 serde_json::json!({"path": destination}),
18648 ToolCallSource::Manual,
18649 ))
18650 .await
18651 .unwrap()
18652 })
18653 };
18654 parent_gate.wait_until_entered().await;
18655
18656 let child_call = {
18657 let child = Arc::clone(&spawned.agent);
18658 let destination = destination.clone();
18659 tokio::spawn(async move {
18660 child
18661 .invoke_tool(ToolExecutionRequest::new(
18662 "spawned-file-write",
18663 "file_write",
18664 serde_json::json!({
18665 "path": destination,
18666 "content": "spawned",
18667 "dry_run": false
18668 }),
18669 ToolCallSource::Manual,
18670 ))
18671 .await
18672 .unwrap()
18673 })
18674 };
18675 wait_for_resource_lock_strong_count(&parent.resource_locks, 2).await;
18676 assert!(!child_call.is_finished());
18677
18678 parent_gate.release();
18679 let (parent_record, child_record) =
18680 tokio::time::timeout(std::time::Duration::from_secs(2), async {
18681 tokio::join!(parent_call, child_call)
18682 })
18683 .await
18684 .expect("parent and spawned path mutations did not finish");
18685 assert!(parent_record.unwrap().success);
18686 assert!(child_record.unwrap().success);
18687 assert_eq!(std::fs::read_to_string(destination).unwrap(), "spawned");
18688 assert!(parent.resource_locks.read().is_empty());
18689 }
18690
18691 #[tokio::test]
18692 async fn cancelled_global_path_lock_waiter_does_not_retain_weak_entry() {
18693 let locks = new_tool_resource_locks();
18694 let holder_gate = PathMutationGate::new();
18695 let waiter_gate = PathMutationGate::new();
18696 waiter_gate.release();
18697 let holder = Arc::new(
18698 AgentBuilder::new()
18699 .system_prompt("Hold the global path lock.")
18700 .llm(Arc::new(mock_with_response("done")))
18701 .tool(Arc::new(BlockingPathMutationTool {
18702 id: "holder_write",
18703 path_fields: vec![ai_agents_core::PathPolicyBinding::write("path")],
18704 gate: holder_gate.clone(),
18705 }))
18706 .build()
18707 .unwrap()
18708 .with_shared_resource_locks(Arc::clone(&locks)),
18709 );
18710 let waiter = Arc::new(
18711 AgentBuilder::new()
18712 .system_prompt("Wait for the global path lock.")
18713 .llm(Arc::new(mock_with_response("done")))
18714 .tool(Arc::new(BlockingPathMutationTool {
18715 id: "waiter_write",
18716 path_fields: vec![ai_agents_core::PathPolicyBinding::write("path")],
18717 gate: waiter_gate.clone(),
18718 }))
18719 .build()
18720 .unwrap()
18721 .with_shared_resource_locks(Arc::clone(&locks)),
18722 );
18723
18724 let holder_call = {
18725 let holder = Arc::clone(&holder);
18726 tokio::spawn(async move {
18727 holder
18728 .invoke_tool(ToolExecutionRequest::new(
18729 "holder-call",
18730 "holder_write",
18731 serde_json::json!({"path": "./shared.txt"}),
18732 ToolCallSource::Manual,
18733 ))
18734 .await
18735 .unwrap()
18736 })
18737 };
18738 holder_gate.wait_until_entered().await;
18739
18740 let waiter_call = {
18741 let waiter = Arc::clone(&waiter);
18742 tokio::spawn(async move {
18743 waiter
18744 .invoke_tool(ToolExecutionRequest::new(
18745 "waiter-call",
18746 "waiter_write",
18747 serde_json::json!({"path": "./shared.txt"}),
18748 ToolCallSource::Manual,
18749 ))
18750 .await
18751 .unwrap()
18752 })
18753 };
18754 wait_for_resource_lock_strong_count(&locks, 2).await;
18755 waiter.runtime_control().cancel_all();
18756
18757 let waiter_record = tokio::time::timeout(std::time::Duration::from_secs(2), waiter_call)
18758 .await
18759 .expect("cancelled lock waiter did not finish")
18760 .unwrap();
18761 assert!(!waiter_record.success);
18762 assert!(!waiter_record.executed);
18763 assert!(waiter_record.cancelled);
18764 assert_eq!(
18765 waiter_record.cancellation_reason.as_deref(),
18766 Some("runtime control cancellation")
18767 );
18768 assert!(!waiter_gate.entered.load(Ordering::SeqCst));
18769 assert_eq!(
18770 locks
18771 .read()
18772 .get("path-mutation:global")
18773 .map_or(0, |lock| lock.strong_count()),
18774 1
18775 );
18776
18777 holder_gate.release();
18778 let holder_record = tokio::time::timeout(std::time::Duration::from_secs(2), holder_call)
18779 .await
18780 .expect("lock holder did not finish")
18781 .unwrap();
18782 assert!(holder_record.success);
18783 assert!(locks.read().is_empty());
18784 }
18785
18786 #[tokio::test]
18787 async fn path_mutation_policy_and_approval_denials_do_not_invoke_tools() {
18788 for denial in [MutationDenial::Policy, MutationDenial::Approval] {
18789 let tools: [Arc<dyn Tool>; 3] = [
18790 Arc::new(CopyPathTool::new()),
18791 Arc::new(MovePathTool::new()),
18792 Arc::new(DeletePathTool::new()),
18793 ];
18794 for tool in tools {
18795 assert_path_mutation_denied(tool, denial).await;
18796 }
18797 }
18798 }
18799
18800 #[tokio::test]
18801 async fn policy_denial_keeps_executor_hook_lifecycle_and_record_authority() {
18802 let workspace = MutationTestWorkspace::new();
18803 let target = workspace.root.join("denied.txt");
18804 let hooks = Arc::new(ToolLifecycleRecordingHooks::new());
18805 let agent = AgentBuilder::new()
18806 .system_prompt("Test denied tool hooks.")
18807 .llm(Arc::new(mock_with_response("done")))
18808 .tool(Arc::new(FileWriteTool::new()))
18809 .tool_security(ToolSecurityEngine::new(mutation_denial_security_config(
18810 "file_write",
18811 &workspace.root,
18812 MutationDenial::Policy,
18813 )))
18814 .hooks(hooks.clone())
18815 .build()
18816 .unwrap();
18817
18818 let record = agent
18819 .invoke_tool(ToolExecutionRequest::new(
18820 "denied-hook-call",
18821 "file_write",
18822 serde_json::json!({
18823 "path": target.to_string_lossy(),
18824 "content": "blocked"
18825 }),
18826 ToolCallSource::Manual,
18827 ))
18828 .await
18829 .unwrap();
18830
18831 assert!(!record.executed);
18832 assert!(!record.success);
18833 assert_eq!(record.policy.outcome, PermissionOutcome::Deny);
18834 assert_eq!(
18835 hooks.events(),
18836 vec![
18837 "start:file_write",
18838 "complete:file_write:false",
18839 "record:file_write:false",
18840 "error"
18841 ]
18842 );
18843 assert!(!target.exists());
18844 }
18845
18846 #[tokio::test]
18847 async fn approval_argument_changes_are_rechecked_against_final_scope() {
18848 let active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
18849 let max_active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
18850 let entered = Arc::new(tokio::sync::Barrier::new(2));
18851 let release = Arc::new(tokio::sync::Notify::new());
18852 let handler = Arc::new(BlockingApprovalHandler {
18853 entered: Arc::clone(&entered),
18854 release: Arc::clone(&release),
18855 result: ApprovalResult::Modified {
18856 changes: HashMap::from([(
18857 "path".to_string(),
18858 Value::String("./after-approval.txt".to_string()),
18859 )]),
18860 },
18861 });
18862 let agent = Arc::new(
18863 AgentBuilder::new()
18864 .system_prompt("Test final scope validation.")
18865 .llm(Arc::new(mock_with_response("done")))
18866 .tool(Arc::new(LockedWriteTool {
18867 active: Arc::clone(&active),
18868 max_active: Arc::clone(&max_active),
18869 }))
18870 .tool_security(ToolSecurityEngine::new(approval_security_config(true)))
18871 .hitl_engine(HITLEngine::new(ai_agents_hitl::HITLConfig::default()))
18872 .approval_handler(handler)
18873 .build()
18874 .unwrap(),
18875 );
18876 let control = agent.runtime_control();
18877 let running = Arc::clone(&agent);
18878 let call = tokio::spawn(async move {
18879 running
18880 .invoke_tool(ToolExecutionRequest::new(
18881 "approval-scope",
18882 "locked_write",
18883 serde_json::json!({"path": "./before-approval.txt"}),
18884 ToolCallSource::Manual,
18885 ))
18886 .await
18887 .unwrap()
18888 });
18889 entered.wait().await;
18890 let expected_version = control.set_tool_scope(Vec::new());
18891 release.notify_one();
18892 let record = call.await.unwrap();
18893
18894 assert!(!record.executed);
18895 assert!(!record.success);
18896 assert_eq!(record.runtime_config_version, expected_version);
18897 assert_eq!(record.executed_arguments["path"], "./after-approval.txt");
18898 assert_eq!(max_active.load(Ordering::SeqCst), 0);
18899 assert_eq!(
18900 record.metadata["runtime_scope_snapshot"],
18901 serde_json::json!([])
18902 );
18903 }
18904
18905 #[tokio::test]
18906 async fn approval_is_rechecked_against_final_policy_snapshot() {
18907 let active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
18908 let max_active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
18909 let entered = Arc::new(tokio::sync::Barrier::new(2));
18910 let release = Arc::new(tokio::sync::Notify::new());
18911 let handler = Arc::new(BlockingApprovalHandler {
18912 entered: Arc::clone(&entered),
18913 release: Arc::clone(&release),
18914 result: ApprovalResult::Approved,
18915 });
18916 let agent = Arc::new(
18917 AgentBuilder::new()
18918 .system_prompt("Test final policy validation.")
18919 .llm(Arc::new(mock_with_response("done")))
18920 .tool(Arc::new(LockedWriteTool {
18921 active: Arc::clone(&active),
18922 max_active: Arc::clone(&max_active),
18923 }))
18924 .tool_security(ToolSecurityEngine::new(approval_security_config(true)))
18925 .hitl_engine(HITLEngine::new(ai_agents_hitl::HITLConfig::default()))
18926 .approval_handler(handler)
18927 .build()
18928 .unwrap(),
18929 );
18930 let control = agent.runtime_control();
18931 let running = Arc::clone(&agent);
18932 let call = tokio::spawn(async move {
18933 running
18934 .invoke_tool(ToolExecutionRequest::new(
18935 "approval-policy",
18936 "locked_write",
18937 serde_json::json!({"path": "./policy.txt"}),
18938 ToolCallSource::Manual,
18939 ))
18940 .await
18941 .unwrap()
18942 });
18943 entered.wait().await;
18944 let expected_version = control.set_tool_security(approval_security_config(false));
18945 release.notify_one();
18946 let record = call.await.unwrap();
18947
18948 assert!(!record.executed);
18949 assert!(!record.success);
18950 assert_eq!(record.runtime_config_version, expected_version);
18951 assert_eq!(record.policy.outcome, PermissionOutcome::Unavailable);
18952 assert_eq!(max_active.load(Ordering::SeqCst), 0);
18953 assert!(record.metadata.contains_key("policy_snapshot"));
18954 }
18955
18956 #[test]
18957 fn invalid_live_policy_does_not_replace_snapshot_or_generation() {
18958 let agent = AgentBuilder::new()
18959 .system_prompt("Test runtime policy validation.")
18960 .llm(Arc::new(mock_with_response("done")))
18961 .build()
18962 .unwrap();
18963 let control = agent.runtime_control();
18964 let mut valid = ToolSecurityConfig::default();
18965 valid.tools.insert(
18966 "web_search".to_string(),
18967 ai_agents_tools::ToolPolicyConfig {
18968 max_results: Some(5),
18969 ..Default::default()
18970 },
18971 );
18972 let generation = control.try_set_tool_security(valid).unwrap();
18973
18974 let mut invalid = ToolSecurityConfig::default();
18975 invalid.tools.insert(
18976 "web_search".to_string(),
18977 ai_agents_tools::ToolPolicyConfig {
18978 max_results: Some(0),
18979 ..Default::default()
18980 },
18981 );
18982 let error = control.try_set_tool_security(invalid).unwrap_err();
18983
18984 assert!(
18985 error
18986 .to_string()
18987 .contains("max_results must be greater than 0")
18988 );
18989 assert_eq!(control.version(), generation);
18990 assert_eq!(
18991 control
18992 .state
18993 .tool_security_override
18994 .read()
18995 .as_ref()
18996 .unwrap()
18997 .config()
18998 .tools["web_search"]
18999 .max_results,
19000 Some(5)
19001 );
19002 }
19003
19004 #[test]
19006 fn invalid_timeout_config_stops_before_approval_or_tool_invocation() {
19007 let tool_calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19008 let approval_calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19009 let spec = crate::spec::AgentSpec {
19010 tool_security: ToolSecurityConfig {
19011 enabled: true,
19012 default_timeout_ms: u64::MAX,
19013 ..Default::default()
19014 },
19015 ..Default::default()
19016 };
19017
19018 let result = AgentBuilder::from_spec(spec)
19019 .llm(Arc::new(mock_with_response("done")))
19020 .tool(Arc::new(FlakyWriteTool {
19021 calls: Arc::clone(&tool_calls),
19022 }))
19023 .hitl_engine(HITLEngine::new(ai_agents_hitl::HITLConfig::default()))
19024 .approval_handler(Arc::new(CountingApprovalHandler {
19025 calls: Arc::clone(&approval_calls),
19026 }))
19027 .build();
19028
19029 assert!(result.is_err());
19030 assert_eq!(approval_calls.load(Ordering::SeqCst), 0);
19031 assert_eq!(tool_calls.load(Ordering::SeqCst), 0);
19032 }
19033
19034 #[test]
19036 fn invalid_recovery_timeout_config_stops_before_approval_or_tool_invocation() {
19037 use ai_agents_recovery::{ErrorRecoveryConfig, ToolRecoveryConfig, ToolRetryConfig};
19038
19039 let tool_calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19040 let approval_calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19041 let spec = crate::spec::AgentSpec {
19042 error_recovery: ErrorRecoveryConfig {
19043 tools: ToolRecoveryConfig {
19044 default: ToolRetryConfig {
19045 timeout_ms: Some(u64::MAX),
19046 ..Default::default()
19047 },
19048 ..Default::default()
19049 },
19050 ..Default::default()
19051 },
19052 ..Default::default()
19053 };
19054
19055 let result = AgentBuilder::from_spec(spec)
19056 .llm(Arc::new(mock_with_response("done")))
19057 .tool(Arc::new(FlakyWriteTool {
19058 calls: Arc::clone(&tool_calls),
19059 }))
19060 .hitl_engine(HITLEngine::new(ai_agents_hitl::HITLConfig::default()))
19061 .approval_handler(Arc::new(CountingApprovalHandler {
19062 calls: Arc::clone(&approval_calls),
19063 }))
19064 .build();
19065
19066 assert!(result.is_err());
19067 assert_eq!(approval_calls.load(Ordering::SeqCst), 0);
19068 assert_eq!(tool_calls.load(Ordering::SeqCst), 0);
19069 }
19070
19071 #[test]
19073 fn invalid_timeout_policy_does_not_replace_snapshot_or_generation() {
19074 let agent = AgentBuilder::new()
19075 .system_prompt("Test runtime timeout policy validation.")
19076 .llm(Arc::new(mock_with_response("done")))
19077 .build()
19078 .unwrap();
19079 let control = agent.runtime_control();
19080 let valid = ToolSecurityConfig {
19081 default_timeout_ms: 5_000,
19082 ..Default::default()
19083 };
19084 let generation = control.try_set_tool_security(valid).unwrap();
19085
19086 let invalid = ToolSecurityConfig {
19087 default_timeout_ms: MAX_TOOL_TIMEOUT_MS + 1,
19088 ..Default::default()
19089 };
19090 let error = control.try_set_tool_security(invalid).unwrap_err();
19091
19092 assert!(error.to_string().contains(&format!(
19093 "tool_security.default_timeout_ms must be no greater than {MAX_TOOL_TIMEOUT_MS} milliseconds"
19094 )));
19095 assert_eq!(control.version(), generation);
19096 assert_eq!(
19097 control
19098 .state
19099 .tool_security_override
19100 .read()
19101 .as_ref()
19102 .unwrap()
19103 .config()
19104 .default_timeout_ms,
19105 5_000
19106 );
19107 }
19108
19109 #[test]
19111 fn runtime_tool_timeout_conversion_enforces_the_stable_boundary() {
19112 let timeout = RuntimeAgent::validated_tool_timeout(MAX_TOOL_TIMEOUT_MS).unwrap();
19113 assert_eq!(timeout.timer, Duration::from_millis(MAX_TOOL_TIMEOUT_MS));
19114 assert_eq!(
19115 timeout.deadline_delta,
19116 chrono::Duration::milliseconds(MAX_TOOL_TIMEOUT_MS as i64)
19117 );
19118
19119 for timeout_ms in [MAX_TOOL_TIMEOUT_MS + 1, u64::MAX] {
19120 let error = RuntimeAgent::validated_tool_timeout(timeout_ms).unwrap_err();
19121 assert!(error.to_string().contains(&format!(
19122 "effective tool timeout_ms must be no greater than {MAX_TOOL_TIMEOUT_MS} milliseconds"
19123 )));
19124 }
19125 }
19126
19127 #[tokio::test]
19128 async fn persistent_override_preserves_rate_history_within_generation() {
19129 let calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19130 let agent = AgentBuilder::new()
19131 .system_prompt("Test persistent policy overrides.")
19132 .llm(Arc::new(mock_with_response("done")))
19133 .tool(Arc::new(RecoveryTestTool {
19134 id: "limited_override".to_string(),
19135 succeeds: true,
19136 calls: Arc::clone(&calls),
19137 max_output_chars: None,
19138 }))
19139 .build()
19140 .unwrap();
19141 let mut security = ToolSecurityConfig {
19142 enabled: true,
19143 fail_closed: true,
19144 ..Default::default()
19145 };
19146 let policy = ai_agents_tools::ToolPolicyConfig {
19147 write_paths: vec![".".to_string()],
19148 rate_limit: Some(1),
19149 ..Default::default()
19150 };
19151 security
19152 .tools
19153 .insert("limited_override".to_string(), policy);
19154 let generation = agent.runtime_control().set_tool_security(security);
19155
19156 let first = agent
19157 .invoke_tool(ToolExecutionRequest::new(
19158 "limited-first",
19159 "limited_override",
19160 serde_json::json!({"path": "./limited.txt"}),
19161 ToolCallSource::Manual,
19162 ))
19163 .await
19164 .unwrap();
19165 let second = agent
19166 .invoke_tool(ToolExecutionRequest::new(
19167 "limited-second",
19168 "limited_override",
19169 serde_json::json!({"path": "./limited.txt"}),
19170 ToolCallSource::Manual,
19171 ))
19172 .await
19173 .unwrap();
19174
19175 assert!(first.success);
19176 assert_eq!(first.policy_version, generation);
19177 assert!(!second.executed);
19178 assert!(second.output.contains("Rate limit exceeded"));
19179 assert_eq!(second.policy_version, generation);
19180 assert_eq!(calls.load(Ordering::SeqCst), 1);
19181 }
19182
19183 #[tokio::test]
19184 async fn concurrent_rate_admission_consumes_capacity_atomically() {
19185 let calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19186 let tool = Arc::new(RecoveryTestTool {
19187 id: "atomic_rate".to_string(),
19188 succeeds: true,
19189 calls: Arc::clone(&calls),
19190 max_output_chars: None,
19191 });
19192 let arguments = serde_json::json!({"path": "./atomic-rate.txt"});
19193 let bindings = tool.policy_bindings();
19194 let classification = tool.classify_call(&arguments);
19195 let resource_keys =
19196 tool_resource_lock_keys(tool.id(), &arguments, &bindings, &classification);
19197 let mut security = ToolSecurityConfig {
19198 enabled: true,
19199 fail_closed: true,
19200 ..Default::default()
19201 };
19202 let policy = ai_agents_tools::ToolPolicyConfig {
19203 write_paths: vec![".".to_string()],
19204 rate_limit: Some(1),
19205 ..Default::default()
19206 };
19207 security.tools.insert(tool.id().to_string(), policy);
19208 let agent = Arc::new(
19209 AgentBuilder::new()
19210 .system_prompt("Test atomic rate admission.")
19211 .llm(Arc::new(mock_with_response("done")))
19212 .tool(tool)
19213 .tool_security(ToolSecurityEngine::new(security))
19214 .build()
19215 .unwrap(),
19216 );
19217 let held = agent
19218 .acquire_tool_resource_locks(&resource_keys)
19219 .await
19220 .unwrap();
19221 let left = {
19222 let agent = Arc::clone(&agent);
19223 let arguments = arguments.clone();
19224 tokio::spawn(async move {
19225 agent
19226 .invoke_tool(ToolExecutionRequest::new(
19227 "atomic-rate-left",
19228 "atomic_rate",
19229 arguments,
19230 ToolCallSource::Manual,
19231 ))
19232 .await
19233 .unwrap()
19234 })
19235 };
19236 let right = {
19237 let agent = Arc::clone(&agent);
19238 tokio::spawn(async move {
19239 agent
19240 .invoke_tool(ToolExecutionRequest::new(
19241 "atomic-rate-right",
19242 "atomic_rate",
19243 arguments,
19244 ToolCallSource::Manual,
19245 ))
19246 .await
19247 .unwrap()
19248 })
19249 };
19250 tokio::time::sleep(std::time::Duration::from_millis(25)).await;
19251 drop(held);
19252 let (left, right) = tokio::join!(left, right);
19253 let records = [left.unwrap(), right.unwrap()];
19254
19255 assert_eq!(records.iter().filter(|record| record.success).count(), 1);
19256 assert_eq!(records.iter().filter(|record| record.executed).count(), 1);
19257 assert!(
19258 records.iter().any(|record| {
19259 !record.executed && record.output.contains("Rate limit exceeded")
19260 })
19261 );
19262 assert_eq!(calls.load(Ordering::SeqCst), 1);
19263 }
19264
19265 #[tokio::test]
19266 async fn changed_policy_generation_invalidates_pending_approval() {
19267 let active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19268 let max_active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19269 let entered = Arc::new(tokio::sync::Barrier::new(2));
19270 let release = Arc::new(tokio::sync::Notify::new());
19271 let handler = Arc::new(BlockingApprovalHandler {
19272 entered: Arc::clone(&entered),
19273 release: Arc::clone(&release),
19274 result: ApprovalResult::Approved,
19275 });
19276 let agent = Arc::new(
19277 AgentBuilder::new()
19278 .system_prompt("Test stale approval denial.")
19279 .llm(Arc::new(mock_with_response("done")))
19280 .tool(Arc::new(LockedWriteTool {
19281 active: Arc::clone(&active),
19282 max_active: Arc::clone(&max_active),
19283 }))
19284 .tool_security(ToolSecurityEngine::new(approval_security_config(true)))
19285 .hitl_engine(HITLEngine::new(ai_agents_hitl::HITLConfig::default()))
19286 .approval_handler(handler)
19287 .build()
19288 .unwrap(),
19289 );
19290 let running = Arc::clone(&agent);
19291 let call = tokio::spawn(async move {
19292 running
19293 .invoke_tool(ToolExecutionRequest::new(
19294 "stale-approval",
19295 "locked_write",
19296 serde_json::json!({"path": "./stale.txt"}),
19297 ToolCallSource::Manual,
19298 ))
19299 .await
19300 .unwrap()
19301 });
19302 entered.wait().await;
19303 let generation = agent
19304 .runtime_control()
19305 .set_tool_security(approval_security_config(true));
19306 release.notify_one();
19307 let record = call.await.unwrap();
19308
19309 assert!(!record.executed);
19310 assert!(record.output.contains("Approval became stale"));
19311 assert_eq!(record.policy_version, generation);
19312 assert_eq!(max_active.load(Ordering::SeqCst), 0);
19313 }
19314
19315 #[tokio::test]
19316 async fn final_policy_reapplies_argument_caps_after_approval_changes() {
19317 use ai_agents_hitl::CallbackHandler;
19318
19319 let mut security = ToolSecurityConfig {
19320 enabled: true,
19321 fail_closed: true,
19322 ..Default::default()
19323 };
19324 let policy = ai_agents_tools::ToolPolicyConfig {
19325 read_paths: vec![".".to_string()],
19326 max_results: Some(5),
19327 require_confirmation: true,
19328 ..Default::default()
19329 };
19330 security.tools.insert("context_echo".to_string(), policy);
19331 let handler = CallbackHandler::new(|_| ApprovalResult::Modified {
19332 changes: HashMap::from([("max_results".to_string(), serde_json::json!(99))]),
19333 });
19334 let agent = AgentBuilder::new()
19335 .system_prompt("Test final argument caps.")
19336 .llm(Arc::new(mock_with_response("done")))
19337 .tool(Arc::new(ContextEchoTool))
19338 .tool_security(ToolSecurityEngine::new(security))
19339 .hitl_engine(HITLEngine::new(ai_agents_hitl::HITLConfig::default()))
19340 .approval_handler(Arc::new(handler))
19341 .build()
19342 .unwrap();
19343
19344 let record = agent
19345 .invoke_tool(ToolExecutionRequest::new(
19346 "final-cap",
19347 "context_echo",
19348 serde_json::json!({"path": ".", "max_results": 1}),
19349 ToolCallSource::Manual,
19350 ))
19351 .await
19352 .unwrap();
19353
19354 assert!(record.success);
19355 assert_eq!(record.executed_arguments["max_results"], 5);
19356 assert_eq!(
19357 record.approval.unwrap().modified_arguments.unwrap()["max_results"],
19358 5
19359 );
19360 }
19361
19362 #[tokio::test]
19363 async fn no_binding_writes_use_canonical_fallback_lock() {
19364 let active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19365 let max_active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19366 let agent = Arc::new(
19367 AgentBuilder::new()
19368 .system_prompt("Test fallback resource locks.")
19369 .llm(Arc::new(mock_with_response("done")))
19370 .tool(Arc::new(NoBindingWriteTool {
19371 active: Arc::clone(&active),
19372 max_active: Arc::clone(&max_active),
19373 }))
19374 .build()
19375 .unwrap(),
19376 );
19377 let left = {
19378 let agent = Arc::clone(&agent);
19379 tokio::spawn(async move {
19380 agent
19381 .invoke_tool(ToolExecutionRequest::new(
19382 "no-binding-left",
19383 "no_binding_write",
19384 serde_json::json!({}),
19385 ToolCallSource::Manual,
19386 ))
19387 .await
19388 .unwrap()
19389 })
19390 };
19391 let right = {
19392 let agent = Arc::clone(&agent);
19393 tokio::spawn(async move {
19394 agent
19395 .invoke_tool(ToolExecutionRequest::new(
19396 "no-binding-right",
19397 "no_binding_write",
19398 serde_json::json!({}),
19399 ToolCallSource::Manual,
19400 ))
19401 .await
19402 .unwrap()
19403 })
19404 };
19405 let (left, right) = tokio::join!(left, right);
19406
19407 assert!(left.unwrap().success);
19408 assert!(right.unwrap().success);
19409 assert_eq!(max_active.load(Ordering::SeqCst), 1);
19410 }
19411
19412 #[tokio::test]
19413 async fn parent_and_child_paths_share_a_resource_lock() {
19414 let active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19415 let max_active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19416 let agent = Arc::new(
19417 AgentBuilder::new()
19418 .system_prompt("Test parent child resource locks.")
19419 .llm(Arc::new(mock_with_response("done")))
19420 .tool(Arc::new(LockedWriteTool {
19421 active: Arc::clone(&active),
19422 max_active: Arc::clone(&max_active),
19423 }))
19424 .build()
19425 .unwrap(),
19426 );
19427 let parent = format!("./lock-parent-{}", uuid::Uuid::new_v4());
19428 let child = format!("{}/child.txt", parent);
19429 let left = {
19430 let agent = Arc::clone(&agent);
19431 tokio::spawn(async move {
19432 agent
19433 .invoke_tool(ToolExecutionRequest::new(
19434 "parent-lock",
19435 "locked_write",
19436 serde_json::json!({"path": parent}),
19437 ToolCallSource::Manual,
19438 ))
19439 .await
19440 .unwrap()
19441 })
19442 };
19443 let right = {
19444 let agent = Arc::clone(&agent);
19445 tokio::spawn(async move {
19446 agent
19447 .invoke_tool(ToolExecutionRequest::new(
19448 "child-lock",
19449 "locked_write",
19450 serde_json::json!({"path": child}),
19451 ToolCallSource::Manual,
19452 ))
19453 .await
19454 .unwrap()
19455 })
19456 };
19457 let (left, right) = tokio::join!(left, right);
19458
19459 assert!(left.unwrap().success);
19460 assert!(right.unwrap().success);
19461 assert_eq!(max_active.load(Ordering::SeqCst), 1);
19462 }
19463
19464 #[tokio::test]
19465 async fn tool_hooks_can_reenter_after_resource_guards_are_dropped() {
19466 let calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19467 let hooks = Arc::new(ReentrantToolHooks {
19468 agent: parking_lot::Mutex::new(None),
19469 invoked: AtomicBool::new(false),
19470 nested_success: AtomicBool::new(false),
19471 });
19472 let agent = Arc::new(
19473 AgentBuilder::new()
19474 .system_prompt("Test hook reentrancy.")
19475 .llm(Arc::new(mock_with_response("done")))
19476 .tool(Arc::new(RecoveryTestTool {
19477 id: "reentrant_write".to_string(),
19478 succeeds: true,
19479 calls: Arc::clone(&calls),
19480 max_output_chars: None,
19481 }))
19482 .hooks(hooks.clone())
19483 .build()
19484 .unwrap(),
19485 );
19486 *hooks.agent.lock() = Some(Arc::downgrade(&agent));
19487 let record = tokio::time::timeout(
19488 std::time::Duration::from_secs(2),
19489 agent.invoke_tool(ToolExecutionRequest::new(
19490 "outer-hook-call",
19491 "reentrant_write",
19492 serde_json::json!({"path": "./hook.txt"}),
19493 ToolCallSource::Manual,
19494 )),
19495 )
19496 .await
19497 .expect("tool completion hook must not retain resource guards")
19498 .unwrap();
19499
19500 assert!(record.success);
19501 assert!(hooks.nested_success.load(Ordering::SeqCst));
19502 assert_eq!(calls.load(Ordering::SeqCst), 2);
19503 }
19504
19505 #[tokio::test]
19507 async fn fallback_finalizes_original_record_before_shared_execution() {
19508 let primary_calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19509 let fallback_calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19510 let hooks = Arc::new(ToolLifecycleRecordingHooks::new());
19511 let agent = AgentBuilder::new()
19512 .system_prompt("Test fallback execution.")
19513 .llm(Arc::new(mock_with_response("done")))
19514 .tool(Arc::new(RecoveryTestTool {
19515 id: "primary".to_string(),
19516 succeeds: false,
19517 calls: Arc::clone(&primary_calls),
19518 max_output_chars: None,
19519 }))
19520 .tool(Arc::new(RecoveryTestTool {
19521 id: "fallback".to_string(),
19522 succeeds: true,
19523 calls: Arc::clone(&fallback_calls),
19524 max_output_chars: None,
19525 }))
19526 .recovery_manager(recovery_manager_with_fallbacks([(
19527 "primary".to_string(),
19528 "fallback".to_string(),
19529 )]))
19530 .hooks(hooks.clone())
19531 .build()
19532 .unwrap();
19533 let record = tokio::time::timeout(
19534 std::time::Duration::from_secs(2),
19535 agent.invoke_tool(ToolExecutionRequest::new(
19536 "fallback-call",
19537 "primary",
19538 serde_json::json!({"path": "./shared.txt"}),
19539 ToolCallSource::Manual,
19540 )),
19541 )
19542 .await
19543 .expect("fallback must not retain the primary resource guard")
19544 .unwrap();
19545
19546 assert_eq!(
19547 hooks.events(),
19548 vec![
19549 "start:primary",
19550 "complete:primary:false",
19551 "record:primary:true",
19552 "error",
19553 "start:fallback",
19554 "complete:fallback:true",
19555 "record:fallback:true",
19556 ]
19557 );
19558 let records = hooks.records();
19559 assert_eq!(records.len(), 2);
19560 let original = &records[0];
19561 assert_eq!(original.canonical_id, "primary");
19562 assert!(matches!(original.source, ToolCallSource::Manual));
19563 assert!(original.executed);
19564 assert!(!original.success);
19565
19566 let fallback = &records[1];
19567 assert_eq!(fallback.canonical_id, "fallback");
19568 assert_eq!(fallback.call_id, "fallback-call");
19569 assert!(matches!(
19570 &fallback.source,
19571 ToolCallSource::Fallback { original_tool } if original_tool == "primary"
19572 ));
19573 assert!(fallback.executed);
19574 assert!(fallback.success);
19575 assert_eq!(record.canonical_id, fallback.canonical_id);
19576 assert_eq!(record.output, fallback.output);
19577
19578 let history = agent.tool_call_history();
19579 assert_eq!(
19580 history
19581 .iter()
19582 .map(|entry| entry.tool_id.as_str())
19583 .collect::<Vec<_>>(),
19584 vec!["primary", "fallback"]
19585 );
19586 assert_eq!(history[0].result.get("success"), Some(&Value::Bool(false)));
19587 assert_eq!(primary_calls.load(Ordering::SeqCst), 1);
19588 assert_eq!(fallback_calls.load(Ordering::SeqCst), 1);
19589 }
19590
19591 #[tokio::test]
19593 async fn self_fallback_cycle_is_denied_before_reinvocation() {
19594 let calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19595 let hooks = Arc::new(ToolLifecycleRecordingHooks::new());
19596 let agent = AgentBuilder::new()
19597 .system_prompt("Test self-fallback cycle admission.")
19598 .llm(Arc::new(mock_with_response("done")))
19599 .tool(Arc::new(RecoveryTestTool {
19600 id: "primary".to_string(),
19601 succeeds: false,
19602 calls: Arc::clone(&calls),
19603 max_output_chars: None,
19604 }))
19605 .recovery_manager(recovery_manager_with_fallbacks([(
19606 "primary".to_string(),
19607 "primary".to_string(),
19608 )]))
19609 .hooks(hooks.clone())
19610 .build()
19611 .unwrap();
19612
19613 let record = tokio::time::timeout(
19614 std::time::Duration::from_secs(2),
19615 agent.invoke_tool(ToolExecutionRequest::new(
19616 "self-fallback-call",
19617 "primary",
19618 serde_json::json!({"path": "./shared.txt"}),
19619 ToolCallSource::Manual,
19620 )),
19621 )
19622 .await
19623 .expect("self fallback must terminate without recursive execution")
19624 .unwrap();
19625
19626 assert_eq!(calls.load(Ordering::SeqCst), 1);
19627 assert_eq!(record.canonical_id, "primary");
19628 assert!(!record.executed);
19629 assert!(!record.success);
19630 assert_eq!(record.policy.outcome, PermissionOutcome::Deny);
19631 assert!(record.output.contains("fallback cycle"));
19632 assert!(matches!(
19633 record.source,
19634 ToolCallSource::Fallback { ref original_tool } if original_tool == "primary"
19635 ));
19636 assert_eq!(
19637 record.metadata.get("fallback_chain"),
19638 Some(&serde_json::json!(["primary"]))
19639 );
19640 assert_eq!(
19641 hooks.events(),
19642 vec![
19643 "start:primary",
19644 "complete:primary:false",
19645 "record:primary:true",
19646 "error",
19647 "complete:primary:false",
19648 "record:primary:false",
19649 "error",
19650 ]
19651 );
19652 assert_eq!(agent.tool_call_history().len(), 2);
19653 }
19654
19655 #[tokio::test]
19657 async fn alias_mediated_fallback_cycle_is_denied_canonically() {
19658 let primary_calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19659 let secondary_calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19660 let hooks = Arc::new(ToolLifecycleRecordingHooks::new());
19661 let agent = AgentBuilder::new()
19662 .system_prompt("Test canonical fallback cycle admission.")
19663 .llm(Arc::new(mock_with_response("done")))
19664 .tool(Arc::new(RecoveryTestTool {
19665 id: "primary".to_string(),
19666 succeeds: false,
19667 calls: Arc::clone(&primary_calls),
19668 max_output_chars: None,
19669 }))
19670 .tool(Arc::new(RecoveryTestTool {
19671 id: "secondary".to_string(),
19672 succeeds: false,
19673 calls: Arc::clone(&secondary_calls),
19674 max_output_chars: None,
19675 }))
19676 .recovery_manager(recovery_manager_with_fallbacks([
19677 ("primary".to_string(), "secondary".to_string()),
19678 ("secondary".to_string(), "primary alias".to_string()),
19679 ]))
19680 .hooks(hooks.clone())
19681 .build()
19682 .unwrap();
19683 agent.tools.set_tool_aliases(
19684 "primary",
19685 ToolAliases::new().with_name("en", "primary alias"),
19686 );
19687
19688 let record = tokio::time::timeout(
19689 std::time::Duration::from_secs(2),
19690 agent.invoke_tool(ToolExecutionRequest::new(
19691 "alias-fallback-call",
19692 "primary",
19693 serde_json::json!({"path": "./shared.txt"}),
19694 ToolCallSource::Manual,
19695 )),
19696 )
19697 .await
19698 .expect("alias-mediated fallback cycle must terminate")
19699 .unwrap();
19700
19701 assert_eq!(primary_calls.load(Ordering::SeqCst), 1);
19702 assert_eq!(secondary_calls.load(Ordering::SeqCst), 1);
19703 assert_eq!(record.requested_name, "primary alias");
19704 assert_eq!(record.canonical_id, "primary");
19705 assert!(!record.executed);
19706 assert!(record.output.contains("fallback cycle"));
19707 assert_eq!(
19708 record.metadata.get("fallback_chain"),
19709 Some(&serde_json::json!(["primary", "secondary"]))
19710 );
19711 assert_eq!(hooks.records().len(), 3);
19712 assert_eq!(agent.tool_call_history().len(), 3);
19713 }
19714
19715 #[tokio::test]
19717 async fn final_canonical_drift_cannot_bypass_fallback_ancestry() {
19718 let primary_calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19719 let secondary_calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19720 let provider = Arc::new(DriftingFallbackProvider {
19721 refreshed: AtomicBool::new(false),
19722 primary_calls: Arc::clone(&primary_calls),
19723 secondary_calls: Arc::clone(&secondary_calls),
19724 });
19725 let registry = ToolRegistry::new();
19726 registry.register_provider(provider).await.unwrap();
19727 let lifecycle = Arc::new(ToolLifecycleRecordingHooks::new());
19728 let hooks = Arc::new(RefreshFallbackProviderHooks {
19729 agent: parking_lot::Mutex::new(None),
19730 lifecycle: Arc::clone(&lifecycle),
19731 });
19732 let agent = Arc::new(
19733 AgentBuilder::new()
19734 .system_prompt("Test final canonical fallback admission.")
19735 .llm(Arc::new(mock_with_response("done")))
19736 .tools(registry)
19737 .recovery_manager(recovery_manager_with_fallbacks([(
19738 "primary".to_string(),
19739 "fallback alias".to_string(),
19740 )]))
19741 .hooks(hooks.clone())
19742 .build()
19743 .unwrap(),
19744 );
19745 *hooks.agent.lock() = Some(Arc::downgrade(&agent));
19746
19747 let record = agent
19748 .invoke_tool(ToolExecutionRequest::new(
19749 "drifting-fallback-call",
19750 "primary",
19751 serde_json::json!({"path": "./shared.txt"}),
19752 ToolCallSource::Manual,
19753 ))
19754 .await
19755 .unwrap();
19756
19757 assert_eq!(primary_calls.load(Ordering::SeqCst), 1);
19758 assert_eq!(secondary_calls.load(Ordering::SeqCst), 0);
19759 assert_eq!(record.canonical_id, "secondary");
19760 assert!(!record.executed);
19761 assert_eq!(record.policy.outcome, PermissionOutcome::Deny);
19762 assert!(record.output.contains("fallback cycle"));
19763 assert_eq!(
19764 record.metadata.get("fallback_chain"),
19765 Some(&serde_json::json!(["primary", "secondary"]))
19766 );
19767 assert_eq!(
19768 record.metadata.get("final_resolved_canonical_id"),
19769 Some(&serde_json::json!("primary"))
19770 );
19771 assert_eq!(
19772 lifecycle.events(),
19773 vec![
19774 "start:primary",
19775 "complete:primary:false",
19776 "record:primary:true",
19777 "error",
19778 "start:secondary",
19779 "complete:secondary:false",
19780 "record:secondary:false",
19781 "error",
19782 ]
19783 );
19784 let records = lifecycle.records();
19785 assert_eq!(records.len(), 2);
19786 assert_eq!(records[1].canonical_id, "secondary");
19787 assert_eq!(
19788 records[1].metadata.get("final_resolved_canonical_id"),
19789 Some(&serde_json::json!("primary"))
19790 );
19791 let history = agent.tool_call_history();
19792 assert_eq!(
19793 history
19794 .iter()
19795 .map(|entry| entry.tool_id.as_str())
19796 .collect::<Vec<_>>(),
19797 vec!["primary", "secondary"]
19798 );
19799 }
19800
19801 #[tokio::test]
19803 async fn acyclic_fallback_chain_is_denied_after_the_hop_limit() {
19804 let tool_count = MAX_TOOL_FALLBACK_HOPS + 2;
19805 let calls = (0..tool_count)
19806 .map(|_| Arc::new(std::sync::atomic::AtomicUsize::new(0)))
19807 .collect::<Vec<_>>();
19808 let mut builder = AgentBuilder::new()
19809 .system_prompt("Test bounded acyclic fallback admission.")
19810 .llm(Arc::new(mock_with_response("done")));
19811 for (index, counter) in calls.iter().enumerate() {
19812 builder = builder.tool(Arc::new(RecoveryTestTool {
19813 id: format!("fallback_{index}"),
19814 succeeds: false,
19815 calls: Arc::clone(counter),
19816 max_output_chars: None,
19817 }));
19818 }
19819 let fallbacks = (0..tool_count - 1).map(|index| {
19820 (
19821 format!("fallback_{index}"),
19822 format!("fallback_{}", index + 1),
19823 )
19824 });
19825 let agent = builder
19826 .recovery_manager(recovery_manager_with_fallbacks(fallbacks))
19827 .build()
19828 .unwrap();
19829
19830 let record = tokio::time::timeout(
19831 std::time::Duration::from_secs(2),
19832 agent.invoke_tool(ToolExecutionRequest::new(
19833 "bounded-fallback-call",
19834 "fallback_0",
19835 serde_json::json!({"path": "./shared.txt"}),
19836 ToolCallSource::Manual,
19837 )),
19838 )
19839 .await
19840 .expect("bounded fallback chain must terminate")
19841 .unwrap();
19842
19843 for counter in calls.iter().take(MAX_TOOL_FALLBACK_HOPS + 1) {
19844 assert_eq!(counter.load(Ordering::SeqCst), 1);
19845 }
19846 assert_eq!(calls[MAX_TOOL_FALLBACK_HOPS + 1].load(Ordering::SeqCst), 0);
19847 assert_eq!(
19848 record.canonical_id,
19849 format!("fallback_{}", MAX_TOOL_FALLBACK_HOPS + 1)
19850 );
19851 assert!(!record.executed);
19852 assert!(record.output.contains("maximum of 16 hops"));
19853 assert_eq!(agent.tool_call_history().len(), tool_count);
19854 }
19855
19856 #[tokio::test]
19857 async fn diagnostics_without_provider_records_unavailable_without_execution() {
19858 let mock = mock_with_response("hello");
19859 let yaml = r#"
19860name: DiagnosticsNoProviderAgent
19861system_prompt: "Review diagnostics."
19862tools: [diagnostics]
19863"#;
19864 let agent = AgentBuilder::from_yaml(yaml)
19865 .unwrap()
19866 .llm(Arc::new(mock))
19867 .auto_configure_features()
19868 .unwrap()
19869 .build()
19870 .unwrap();
19871
19872 let record = agent
19873 .invoke_tool(ToolExecutionRequest::new(
19874 "diagnostics-call",
19875 "diagnostics",
19876 serde_json::json!({}),
19877 ToolCallSource::Manual,
19878 ))
19879 .await
19880 .unwrap();
19881
19882 assert!(!record.executed);
19883 assert!(!record.success);
19884 assert_eq!(record.policy.outcome, PermissionOutcome::Unavailable);
19885 }
19886
19887 #[tokio::test]
19888 async fn web_search_without_provider_records_unavailable_without_execution() {
19889 let mock = mock_with_response("hello");
19890 let yaml = r#"
19891name: WebSearchNoProviderAgent
19892system_prompt: "You search the web."
19893tools: [web_search]
19894"#;
19895 let agent = AgentBuilder::from_yaml(yaml)
19896 .unwrap()
19897 .llm(Arc::new(mock))
19898 .auto_configure_features()
19899 .unwrap()
19900 .build()
19901 .unwrap();
19902
19903 let record = agent
19904 .invoke_tool(ToolExecutionRequest::new(
19905 "web-search-call",
19906 "web_search",
19907 serde_json::json!({"query": "rust async"}),
19908 ToolCallSource::Manual,
19909 ))
19910 .await
19911 .unwrap();
19912
19913 assert!(!record.executed);
19914 assert!(!record.success);
19915 assert_eq!(record.policy.outcome, PermissionOutcome::Unavailable);
19916 }
19917
19918 #[tokio::test]
19919 async fn unavailable_host_tool_does_not_request_approval() {
19920 let approvals = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19921 let handler = Arc::new(CountingApprovalHandler {
19922 calls: Arc::clone(&approvals),
19923 });
19924 let mut security = ToolSecurityConfig {
19925 enabled: true,
19926 fail_closed: true,
19927 ..Default::default()
19928 };
19929 security.tools.insert(
19930 "web_search".to_string(),
19931 ai_agents_tools::ToolPolicyConfig {
19932 enabled: true,
19933 require_confirmation: true,
19934 ..Default::default()
19935 },
19936 );
19937 let yaml = r#"
19938name: UnavailableApprovalAgent
19939system_prompt: "Search only with approval."
19940tools: [web_search]
19941"#;
19942 let agent = AgentBuilder::from_yaml(yaml)
19943 .unwrap()
19944 .llm(Arc::new(mock_with_response("done")))
19945 .auto_configure_features()
19946 .unwrap()
19947 .tool_security(ToolSecurityEngine::new(security))
19948 .hitl_engine(HITLEngine::new(ai_agents_hitl::HITLConfig::default()))
19949 .approval_handler(handler)
19950 .build()
19951 .unwrap();
19952
19953 let record = agent
19954 .invoke_tool(ToolExecutionRequest::new(
19955 "unavailable-before-approval",
19956 "web_search",
19957 serde_json::json!({"query": "rust async"}),
19958 ToolCallSource::Manual,
19959 ))
19960 .await
19961 .unwrap();
19962
19963 assert_eq!(approvals.load(Ordering::SeqCst), 0);
19964 assert!(!record.executed);
19965 assert!(!record.success);
19966 assert_eq!(record.policy.outcome, PermissionOutcome::Unavailable);
19967 assert!(
19968 record
19969 .approval
19970 .as_ref()
19971 .is_some_and(|approval| matches!(approval.status, ToolApprovalStatus::Unavailable))
19972 );
19973 }
19974
19975 #[tokio::test]
19976 async fn test_spawner_section_does_not_grant_core_tools_when_top_level_tools_omitted() {
19977 let mock = mock_with_response("hello");
19978 let yaml = r#"
19979name: SpawnerNoGrantAgent
19980system_prompt: "You manage agents."
19981spawner:
19982 max_agents: 2
19983"#;
19984 let agent = AgentBuilder::from_yaml(yaml)
19985 .unwrap()
19986 .llm(Arc::new(mock))
19987 .auto_configure_features()
19988 .unwrap()
19989 .auto_configure_spawner()
19990 .await
19991 .unwrap()
19992 .build()
19993 .unwrap();
19994
19995 let available = agent.get_available_tool_ids().await.unwrap();
19996 assert!(available.is_empty());
19997 }
19998
19999 #[tokio::test]
20000 async fn test_spawner_section_does_not_grant_core_tools_when_top_level_tools_empty() {
20001 let mock = mock_with_response("hello");
20002 let yaml = r#"
20003name: EmptySpawnerNoGrantAgent
20004system_prompt: "You manage agents."
20005tools: []
20006spawner:
20007 max_agents: 2
20008"#;
20009 let agent = AgentBuilder::from_yaml(yaml)
20010 .unwrap()
20011 .llm(Arc::new(mock))
20012 .auto_configure_features()
20013 .unwrap()
20014 .auto_configure_spawner()
20015 .await
20016 .unwrap()
20017 .build()
20018 .unwrap();
20019
20020 let available = agent.get_available_tool_ids().await.unwrap();
20021 assert!(available.is_empty());
20022 }
20023
20024 #[tokio::test]
20025 async fn test_management_tools_flag_grants_core_tools_when_top_level_tools_empty() {
20026 let mock = mock_with_response("hello");
20027 let yaml = r#"
20028name: ManagementGrantAgent
20029system_prompt: "You manage agents."
20030tools: []
20031spawner:
20032 management_tools: true
20033"#;
20034 let agent = AgentBuilder::from_yaml(yaml)
20035 .unwrap()
20036 .llm(Arc::new(mock))
20037 .auto_configure_features()
20038 .unwrap()
20039 .auto_configure_spawner()
20040 .await
20041 .unwrap()
20042 .build()
20043 .unwrap();
20044
20045 let available = agent.get_available_tool_ids().await.unwrap();
20046 assert_eq!(available.len(), 4);
20047 assert!(available.contains(&"spawn_agent".to_string()));
20048 assert!(available.contains(&"send_agent_message".to_string()));
20049 assert!(available.contains(&"list_agents".to_string()));
20050 assert!(available.contains(&"remove_agent".to_string()));
20051 }
20052
20053 #[tokio::test]
20054 async fn test_management_tools_flag_grants_core_tools_when_top_level_tools_omitted() {
20055 let mock = mock_with_response("hello");
20056 let yaml = r#"
20057name: ManagementOmittedToolsGrantAgent
20058system_prompt: "You manage agents."
20059spawner:
20060 management_tools: true
20061"#;
20062 let agent = AgentBuilder::from_yaml(yaml)
20063 .unwrap()
20064 .llm(Arc::new(mock))
20065 .auto_configure_features()
20066 .unwrap()
20067 .auto_configure_spawner()
20068 .await
20069 .unwrap()
20070 .build()
20071 .unwrap();
20072
20073 let available = agent.get_available_tool_ids().await.unwrap();
20074 assert_eq!(available.len(), 4);
20075 assert!(available.contains(&"spawn_agent".to_string()));
20076 assert!(available.contains(&"send_agent_message".to_string()));
20077 assert!(available.contains(&"list_agents".to_string()));
20078 assert!(available.contains(&"remove_agent".to_string()));
20079 }
20080
20081 #[tokio::test]
20082 async fn test_management_tools_selected_grants_only_selected_tools() {
20083 let mock = mock_with_response("hello");
20084 let yaml = r#"
20085name: ManagementSelectedGrantAgent
20086system_prompt: "You manage agents."
20087tools: []
20088spawner:
20089 management_tools:
20090 - spawn_agent
20091 - send_agent_message
20092 - list_agents
20093"#;
20094 let agent = AgentBuilder::from_yaml(yaml)
20095 .unwrap()
20096 .llm(Arc::new(mock))
20097 .auto_configure_features()
20098 .unwrap()
20099 .auto_configure_spawner()
20100 .await
20101 .unwrap()
20102 .build()
20103 .unwrap();
20104
20105 let available = agent.get_available_tool_ids().await.unwrap();
20106 assert_eq!(available.len(), 3);
20107 assert!(available.contains(&"spawn_agent".to_string()));
20108 assert!(available.contains(&"send_agent_message".to_string()));
20109 assert!(available.contains(&"list_agents".to_string()));
20110 assert!(!available.contains(&"remove_agent".to_string()));
20111 }
20112
20113 #[tokio::test]
20114 async fn test_orchestration_tools_flag_grants_tools_when_top_level_tools_empty() {
20115 let mock = mock_with_response("hello");
20116 let yaml = r#"
20117name: OrchestrationGrantAgent
20118system_prompt: "You coordinate agents."
20119llms:
20120 default:
20121 provider: openai
20122 model: gpt-4
20123 router:
20124 provider: openai
20125 model: gpt-4
20126llm:
20127 default: default
20128 router: router
20129tools: []
20130spawner:
20131 orchestration_tools: true
20132"#;
20133 let agent = AgentBuilder::from_yaml(yaml)
20134 .unwrap()
20135 .llm(Arc::new(mock))
20136 .auto_configure_features()
20137 .unwrap()
20138 .auto_configure_spawner()
20139 .await
20140 .unwrap()
20141 .build()
20142 .unwrap();
20143
20144 let available = agent.get_available_tool_ids().await.unwrap();
20145 assert_eq!(available.len(), 5);
20146 assert!(available.contains(&"route_to_agent".to_string()));
20147 assert!(available.contains(&"pipeline_process".to_string()));
20148 assert!(available.contains(&"concurrent_ask".to_string()));
20149 assert!(available.contains(&"group_discussion".to_string()));
20150 assert!(available.contains(&"handoff_conversation".to_string()));
20151 }
20152
20153 #[tokio::test]
20154 async fn test_persona_evolve_flag_grants_tool_when_top_level_tools_empty() {
20155 let mock = mock_with_response("hello");
20156 let yaml = r#"
20157name: PersonaGrantAgent
20158system_prompt: "You can evolve persona."
20159llm:
20160 provider: openai
20161 model: gpt-4
20162tools: []
20163persona:
20164 identity:
20165 name: "Guide"
20166 role: "Helper"
20167 evolution:
20168 enabled: true
20169 allow_llm_evolve: true
20170 mutable_fields:
20171 - traits.personality
20172"#;
20173 let agent = AgentBuilder::from_yaml(yaml)
20174 .unwrap()
20175 .llm(Arc::new(mock))
20176 .build()
20177 .unwrap();
20178
20179 let available = agent.get_available_tool_ids().await.unwrap();
20180 assert_eq!(available, vec!["persona_evolve".to_string()]);
20181 }
20182
20183 #[tokio::test]
20184 async fn test_persona_evolve_flag_grants_tool_when_top_level_tools_omitted() {
20185 let mock = mock_with_response("hello");
20186 let yaml = r#"
20187name: PersonaOmittedToolsGrantAgent
20188system_prompt: "You can evolve persona."
20189llm:
20190 provider: openai
20191 model: gpt-4
20192persona:
20193 identity:
20194 name: "Guide"
20195 role: "Helper"
20196 evolution:
20197 enabled: true
20198 allow_llm_evolve: true
20199 mutable_fields:
20200 - traits.personality
20201"#;
20202 let agent = AgentBuilder::from_yaml(yaml)
20203 .unwrap()
20204 .llm(Arc::new(mock))
20205 .build()
20206 .unwrap();
20207
20208 let available = agent.get_available_tool_ids().await.unwrap();
20209 assert_eq!(available, vec!["persona_evolve".to_string()]);
20210 }
20211
20212 #[tokio::test]
20213 async fn test_omitted_yaml_tools_exposes_no_tools() {
20214 let mock = mock_with_response("hello");
20215 let yaml = r#"
20216name: NoToolsAgent
20217system_prompt: "You are helpful."
20218"#;
20219 let agent = AgentBuilder::from_yaml(yaml)
20220 .unwrap()
20221 .llm(Arc::new(mock))
20222 .auto_configure_features()
20223 .unwrap()
20224 .build()
20225 .unwrap();
20226
20227 let available = agent.get_available_tool_ids().await.unwrap();
20228 assert!(available.is_empty());
20229 }
20230
20231 #[tokio::test]
20232 async fn runtime_scope_cannot_widen_omitted_or_empty_yaml_grants() {
20233 for tools in ["", "tools: []"] {
20234 let yaml = format!(
20235 r#"
20236name: RuntimeScopeNoGrantAgent
20237system_prompt: "No ordinary tools are granted."
20238{tools}
20239"#
20240 );
20241 let agent = AgentBuilder::from_yaml(&yaml)
20242 .unwrap()
20243 .llm(Arc::new(mock_with_response("done")))
20244 .auto_configure_features()
20245 .unwrap()
20246 .build()
20247 .unwrap();
20248
20249 agent
20250 .runtime_control()
20251 .set_tool_scope(vec!["calculator".to_string()]);
20252
20253 assert!(agent.get_available_tool_ids().await.unwrap().is_empty());
20254 }
20255 }
20256
20257 #[tokio::test]
20258 async fn runtime_scope_widening_attempt_keeps_only_declared_tools() {
20259 let yaml = r#"
20260name: RuntimeScopeWideningAgent
20261system_prompt: "Runtime scope cannot add authority."
20262tools: [calculator]
20263"#;
20264 let agent = AgentBuilder::from_yaml(yaml)
20265 .unwrap()
20266 .llm(Arc::new(mock_with_response("done")))
20267 .auto_configure_features()
20268 .unwrap()
20269 .build()
20270 .unwrap();
20271
20272 agent
20273 .runtime_control()
20274 .set_tool_scope(vec!["calculator".to_string(), "datetime".to_string()]);
20275
20276 assert_eq!(
20277 agent.get_available_tool_ids().await.unwrap(),
20278 vec!["calculator".to_string()]
20279 );
20280 }
20281
20282 #[tokio::test]
20283 async fn runtime_scope_is_canonical_unique_ordered_and_clear_restores_declared_grant() {
20284 let yaml = r#"
20285name: RuntimeScopeIntersectionAgent
20286system_prompt: "Use only declared tools."
20287tools: [calculator, datetime]
20288"#;
20289 let agent = AgentBuilder::from_yaml(yaml)
20290 .unwrap()
20291 .llm(Arc::new(mock_with_response("done")))
20292 .auto_configure_features()
20293 .unwrap()
20294 .build()
20295 .unwrap();
20296 let mut aliases = ai_agents_tools::ToolAliases::default();
20297 aliases
20298 .names
20299 .insert("en".to_string(), "calculate_alias".to_string());
20300 agent.tools.set_tool_aliases("calculator", aliases);
20301 let control = agent.runtime_control();
20302
20303 control.set_tool_scope(vec![
20304 "datetime".to_string(),
20305 "calculate_alias".to_string(),
20306 "calculator".to_string(),
20307 "unknown".to_string(),
20308 "datetime".to_string(),
20309 ]);
20310 assert_eq!(
20311 agent.get_available_tool_ids().await.unwrap(),
20312 vec!["calculator".to_string(), "datetime".to_string()]
20313 );
20314
20315 control.set_tool_scope(vec!["datetime".to_string()]);
20316 assert_eq!(
20317 agent.get_available_tool_ids().await.unwrap(),
20318 vec!["datetime".to_string()]
20319 );
20320
20321 control.clear_tool_scope_override();
20322 assert_eq!(
20323 agent.get_available_tool_ids().await.unwrap(),
20324 vec!["calculator".to_string(), "datetime".to_string()]
20325 );
20326 }
20327
20328 #[tokio::test]
20329 async fn runtime_scope_preserves_programmatic_registration_as_declared_grant() {
20330 let agent = AgentBuilder::new()
20331 .system_prompt("Use registered tools.")
20332 .llm(Arc::new(mock_with_response("done")))
20333 .tool(Arc::new(ContextEchoTool))
20334 .tool(Arc::new(SlowTool))
20335 .build()
20336 .unwrap();
20337
20338 agent.runtime_control().set_tool_scope(vec![
20339 "Context Echo".to_string(),
20340 "context_echo".to_string(),
20341 "unknown".to_string(),
20342 ]);
20343
20344 assert_eq!(
20345 agent.get_available_tool_ids().await.unwrap(),
20346 vec!["context_echo".to_string()]
20347 );
20348 }
20349
20350 #[tokio::test]
20351 async fn nested_state_scopes_intersect_every_ancestor_with_aliases() {
20352 let yaml = r#"
20353name: NestedStateScopeAgent
20354system_prompt: "Honor every state scope."
20355tools: [calculator, datetime, echo]
20356states:
20357 initial: root
20358 states:
20359 root:
20360 tools: [calculate_alias, datetime]
20361 initial: middle
20362 states:
20363 middle:
20364 initial: leaf
20365 states:
20366 leaf:
20367 tools: [datetime_alias, echo]
20368"#;
20369 let agent = AgentBuilder::from_yaml(yaml)
20370 .unwrap()
20371 .llm(Arc::new(mock_with_response("done")))
20372 .auto_configure_features()
20373 .unwrap()
20374 .build()
20375 .unwrap();
20376 let mut calculator_aliases = ai_agents_tools::ToolAliases::default();
20377 calculator_aliases
20378 .names
20379 .insert("en".to_string(), "calculate_alias".to_string());
20380 agent
20381 .tools
20382 .set_tool_aliases("calculator", calculator_aliases);
20383 let mut datetime_aliases = ai_agents_tools::ToolAliases::default();
20384 datetime_aliases
20385 .names
20386 .insert("en".to_string(), "datetime_alias".to_string());
20387 agent.tools.set_tool_aliases("datetime", datetime_aliases);
20388 agent.runtime_control().set_tool_scope(vec![
20389 "unknown".to_string(),
20390 "datetime_alias".to_string(),
20391 "calculate_alias".to_string(),
20392 "datetime".to_string(),
20393 ]);
20394
20395 assert_eq!(agent.current_state().as_deref(), Some("root.middle.leaf"));
20396 assert_eq!(
20397 agent.get_available_tool_ids().await.unwrap(),
20398 vec!["datetime".to_string()]
20399 );
20400 }
20401
20402 #[tokio::test]
20403 async fn ancestor_empty_state_scope_denies_omitted_descendants() {
20404 let yaml = r#"
20405name: NestedEmptyStateScopeAgent
20406system_prompt: "An empty ancestor scope denies all tools."
20407tools: [calculator]
20408states:
20409 initial: root
20410 states:
20411 root:
20412 tools: []
20413 initial: middle
20414 states:
20415 middle:
20416 initial: leaf
20417 states:
20418 leaf: {}
20419"#;
20420 let agent = AgentBuilder::from_yaml(yaml)
20421 .unwrap()
20422 .llm(Arc::new(mock_with_response("done")))
20423 .auto_configure_features()
20424 .unwrap()
20425 .build()
20426 .unwrap();
20427
20428 assert!(agent.get_available_tool_ids().await.unwrap().is_empty());
20429 }
20430
20431 #[tokio::test]
20432 async fn state_change_during_approval_invalidates_the_reviewed_authority() {
20433 let active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
20434 let max_active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
20435 let entered = Arc::new(tokio::sync::Barrier::new(2));
20436 let release = Arc::new(tokio::sync::Notify::new());
20437 let handler = Arc::new(BlockingApprovalHandler {
20438 entered: Arc::clone(&entered),
20439 release: Arc::clone(&release),
20440 result: ApprovalResult::Approved,
20441 });
20442 let yaml = r#"
20443name: ApprovalStateGenerationAgent
20444system_prompt: "State authority may change during approval."
20445tools: [locked_write]
20446states:
20447 initial: first
20448 states:
20449 first:
20450 tools: [locked_write]
20451 second:
20452 tools: [locked_write]
20453"#;
20454 let agent = Arc::new(
20455 AgentBuilder::from_yaml(yaml)
20456 .unwrap()
20457 .llm(Arc::new(mock_with_response("done")))
20458 .tool(Arc::new(LockedWriteTool {
20459 active: Arc::clone(&active),
20460 max_active: Arc::clone(&max_active),
20461 }))
20462 .tool_security(ToolSecurityEngine::new(approval_security_config(true)))
20463 .hitl_engine(HITLEngine::new(ai_agents_hitl::HITLConfig::default()))
20464 .approval_handler(handler)
20465 .build()
20466 .unwrap(),
20467 );
20468 let running = Arc::clone(&agent);
20469 let call = tokio::spawn(async move {
20470 running
20471 .invoke_tool(ToolExecutionRequest::new(
20472 "approval-state-generation",
20473 "locked_write",
20474 serde_json::json!({"path": "./state-generation.txt"}),
20475 ToolCallSource::Manual,
20476 ))
20477 .await
20478 .unwrap()
20479 });
20480
20481 entered.wait().await;
20482 agent.transition_to("second").await.unwrap();
20483 release.notify_one();
20484 let record = call.await.unwrap();
20485
20486 assert!(!record.executed);
20487 assert!(record.output.contains("Approval became stale"));
20488 assert_eq!(max_active.load(Ordering::SeqCst), 0);
20489 }
20490
20491 #[tokio::test]
20492 async fn state_change_while_waiting_for_resource_lock_fails_final_admission() {
20493 let holder_gate = PathMutationGate::new();
20494 let waiter_gate = PathMutationGate::new();
20495 let yaml = r#"
20496name: LockedStateGenerationAgent
20497system_prompt: "State authority must remain stable through admission."
20498tools: [state_lock_holder, state_lock_waiter]
20499states:
20500 initial: first
20501 states:
20502 first:
20503 tools: [state_lock_holder, state_lock_waiter]
20504 second:
20505 tools: [state_lock_holder, state_lock_waiter]
20506"#;
20507 let agent = Arc::new(
20508 AgentBuilder::from_yaml(yaml)
20509 .unwrap()
20510 .llm(Arc::new(mock_with_response("done")))
20511 .tool(Arc::new(BlockingPathMutationTool {
20512 id: "state_lock_holder",
20513 path_fields: vec![ai_agents_core::PathPolicyBinding::write("path")],
20514 gate: holder_gate.clone(),
20515 }))
20516 .tool(Arc::new(BlockingPathMutationTool {
20517 id: "state_lock_waiter",
20518 path_fields: vec![ai_agents_core::PathPolicyBinding::write("path")],
20519 gate: waiter_gate.clone(),
20520 }))
20521 .build()
20522 .unwrap(),
20523 );
20524 let holder_call = {
20525 let agent = Arc::clone(&agent);
20526 tokio::spawn(async move {
20527 agent
20528 .invoke_tool(ToolExecutionRequest::new(
20529 "state-lock-holder",
20530 "state_lock_holder",
20531 serde_json::json!({"path": "./shared-state-path.txt"}),
20532 ToolCallSource::Manual,
20533 ))
20534 .await
20535 .unwrap()
20536 })
20537 };
20538 holder_gate.wait_until_entered().await;
20539 let waiter_call = {
20540 let agent = Arc::clone(&agent);
20541 tokio::spawn(async move {
20542 agent
20543 .invoke_tool(ToolExecutionRequest::new(
20544 "state-lock-waiter",
20545 "state_lock_waiter",
20546 serde_json::json!({"path": "./shared-state-path.txt"}),
20547 ToolCallSource::Manual,
20548 ))
20549 .await
20550 .unwrap()
20551 })
20552 };
20553
20554 wait_for_resource_lock_strong_count(&agent.resource_locks, 2).await;
20555 agent.transition_to("second").await.unwrap();
20556 holder_gate.release();
20557 let holder_record = holder_call.await.unwrap();
20558 let waiter_record = waiter_call.await.unwrap();
20559
20560 assert!(holder_record.success);
20561 assert!(!waiter_record.executed);
20562 assert!(
20563 waiter_record
20564 .output
20565 .contains("state scope changed before admission")
20566 );
20567 assert!(!waiter_gate.entered.load(Ordering::SeqCst));
20568 }
20569
20570 #[tokio::test]
20571 async fn test_state_tools_cannot_widen_top_level_grant() {
20572 let mock = mock_with_response("hello");
20573 let yaml = r#"
20574name: NarrowToolsAgent
20575system_prompt: "You are helpful."
20576tools:
20577 - calculator
20578states:
20579 initial: current
20580 states:
20581 current:
20582 tools: [datetime]
20583"#;
20584 let agent = AgentBuilder::from_yaml(yaml)
20585 .unwrap()
20586 .llm(Arc::new(mock))
20587 .auto_configure_features()
20588 .unwrap()
20589 .build()
20590 .unwrap();
20591
20592 let available = agent.get_available_tool_ids().await.unwrap();
20593 assert!(available.is_empty());
20594 }
20595
20596 #[tokio::test]
20598 async fn test_integration_tool_execution() {
20599 let mock = mock_with_responses(vec![
20601 r#"I'll calculate that for you.
20603{"tool": "calculator", "arguments": {"expression": "2+2"}}"#,
20604 "The answer is 4.",
20606 ]);
20607 let observed = mock.clone();
20608 let mut tools = ai_agents_tools::ToolRegistry::new();
20609 tools
20610 .register(Arc::new(ai_agents_tools::CalculatorTool))
20611 .unwrap();
20612
20613 let agent = AgentBuilder::new()
20614 .system_prompt("You are a calculator assistant.")
20615 .llm(Arc::new(mock))
20616 .tools(tools)
20617 .build()
20618 .unwrap();
20619
20620 let response = agent.chat("What is 2+2?").await.unwrap();
20621
20622 assert_eq!(response.content, "The answer is 4.");
20623 assert_eq!(response.tool_calls.as_ref().map(Vec::len), Some(1));
20624 assert_eq!(
20625 observed.call_count(),
20626 2,
20627 "tool result must trigger a second LLM call"
20628 );
20629 let history = agent.tool_call_history();
20630 assert_eq!(history.len(), 1);
20631 assert_eq!(history[0].tool_id, "calculator");
20632 assert_eq!(
20633 history[0].result.get("result"),
20634 Some(&serde_json::json!(4.0)),
20635 "{:?}",
20636 history[0].result
20637 );
20638 }
20639
20640 #[test]
20643 fn legacy_tool_call_marker_is_plain_text() {
20644 let agent = AgentBuilder::new()
20645 .system_prompt("x")
20646 .llm(Arc::new(mock_with_response("x")))
20647 .build()
20648 .unwrap();
20649 let parsed = agent
20650 .parse_tool_calls(
20651 r#"[TOOL_CALL: {"name": "calculator", "arguments": {"expression": "2+2"}}]"#,
20652 )
20653 .unwrap();
20654 assert!(parsed.is_none());
20655 }
20656
20657 #[tokio::test]
20658 async fn test_tool_hitl_rejection_finalizes_blocking_turn() {
20659 let responses = Arc::new(std::sync::atomic::AtomicUsize::new(0));
20660 let hooks = Arc::new(ResponseCountingHooks {
20661 responses: Arc::clone(&responses),
20662 });
20663 let mock = mock_with_response(r#"{"tool":"echo","arguments":{"message":"hello"}}"#);
20664 let yaml = r#"
20665name: ToolRejectAgent
20666system_prompt: "You use tools when requested."
20667tools:
20668 - echo
20669hitl:
20670 tools:
20671 echo:
20672 require_approval: true
20673 approval_message: "Approve echo?"
20674"#;
20675 let agent = AgentBuilder::from_yaml(yaml)
20676 .unwrap()
20677 .llm(Arc::new(mock))
20678 .auto_configure_features()
20679 .unwrap()
20680 .hooks(hooks)
20681 .build()
20682 .unwrap();
20683
20684 let response = agent.chat("echo hello").await.unwrap();
20685
20686 assert!(
20687 response.content.contains("Operation cancelled"),
20688 "unexpected response: {}",
20689 response.content
20690 );
20691 assert_eq!(responses.load(Ordering::SeqCst), 1);
20692 let messages = agent.memory.get_messages(None).await.unwrap();
20693 assert_eq!(messages.len(), 3);
20694 assert_eq!(messages[0].content, "echo hello");
20695 assert!(messages[1].content.contains("\"tool\":\"echo\""));
20696 assert!(messages[2].content.contains("rejected by the approver"));
20697 }
20698
20699 #[tokio::test]
20700 async fn test_tool_hitl_rejection_finalizes_streaming_turn() {
20701 use futures::StreamExt;
20702
20703 let responses = Arc::new(std::sync::atomic::AtomicUsize::new(0));
20704 let hooks = Arc::new(ResponseCountingHooks {
20705 responses: Arc::clone(&responses),
20706 });
20707 let mock = mock_with_response(r#"{"tool":"echo","arguments":{"message":"hello"}}"#);
20708 let yaml = r#"
20709name: ToolRejectStreamingAgent
20710system_prompt: "You use tools when requested."
20711tools:
20712 - echo
20713streaming:
20714 enabled: true
20715hitl:
20716 tools:
20717 echo:
20718 require_approval: true
20719 approval_message: "Approve echo?"
20720"#;
20721 let agent = AgentBuilder::from_yaml(yaml)
20722 .unwrap()
20723 .llm(Arc::new(mock))
20724 .auto_configure_features()
20725 .unwrap()
20726 .hooks(hooks)
20727 .build()
20728 .unwrap();
20729
20730 let mut stream = agent.chat_stream("echo hello").await.unwrap();
20731 let mut terminal_error = String::new();
20732 let mut done = false;
20733 while let Some(chunk) = stream.next().await {
20734 match chunk {
20735 StreamChunk::Error { message } => terminal_error = message,
20736 StreamChunk::Done {} => {
20737 done = true;
20738 break;
20739 }
20740 _ => {}
20741 }
20742 }
20743
20744 assert!(done);
20745 assert!(
20746 terminal_error.contains("Operation cancelled"),
20747 "unexpected terminal error: {}",
20748 terminal_error
20749 );
20750 assert_eq!(responses.load(Ordering::SeqCst), 1);
20751 let messages = agent.memory.get_messages(None).await.unwrap();
20752 assert_eq!(messages.len(), 3);
20753 assert_eq!(messages[0].content, "echo hello");
20754 assert!(messages[1].content.contains("\"tool\":\"echo\""));
20755 assert!(messages[2].content.contains("rejected by the approver"));
20756 }
20757
20758 #[tokio::test]
20759 async fn tool_hitl_rejection_preserves_legacy_error_but_finalizes_event_stream() {
20760 let mock = mock_with_response(r#"{"tool":"echo","arguments":{"message":"hello"}}"#);
20761 let yaml = r#"
20762name: ToolRejectEventAgent
20763system_prompt: "You use tools when requested."
20764tools:
20765 - echo
20766streaming:
20767 enabled: true
20768hitl:
20769 tools:
20770 echo:
20771 require_approval: true
20772 approval_message: "Approve echo?"
20773"#;
20774 let agent = AgentBuilder::from_yaml(yaml)
20775 .unwrap()
20776 .llm(Arc::new(mock))
20777 .auto_configure_features()
20778 .unwrap()
20779 .build()
20780 .unwrap();
20781
20782 let mut stream = agent.chat_stream_events("echo hello").await.unwrap();
20783 let mut error_seen = false;
20784 let mut final_response = None;
20785 while let Some(event) = stream.next().await {
20786 match event {
20787 AgentStreamEvent::Chunk(StreamChunk::Error { .. }) => error_seen = true,
20788 AgentStreamEvent::Final(response) => final_response = Some(response),
20789 AgentStreamEvent::Chunk(_) => {}
20790 }
20791 }
20792
20793 assert!(!error_seen);
20794 assert!(
20795 final_response
20796 .is_some_and(|response| { response.content.contains("Operation cancelled") })
20797 );
20798 }
20799
20800 #[tokio::test]
20801 async fn test_pre_response_guard_transition_skips_old_state_llm() {
20802 let mock = mock_with_response("Billing state response");
20803 let call_counter = mock.clone();
20804 let yaml = r#"
20805name: OptimizedStateAgent
20806system_prompt: "You route before answering."
20807runtime:
20808 optimization:
20809 enabled: true
20810 pre_response_deterministic_transitions: true
20811states:
20812 initial: greeting
20813 states:
20814 greeting:
20815 prompt: "Old state prompt that should be skipped."
20816 transitions:
20817 - to: billing
20818 guard:
20819 context:
20820 topic:
20821 eq: billing
20822 timing: pre_response
20823 billing:
20824 prompt: "Answer from the billing state."
20825"#;
20826 let agent = AgentBuilder::from_yaml(yaml)
20827 .unwrap()
20828 .llm(Arc::new(mock))
20829 .build()
20830 .unwrap();
20831 agent
20832 .set_context("topic", serde_json::json!("billing"))
20833 .unwrap();
20834
20835 let response = agent.chat("I need billing help").await.unwrap();
20836
20837 assert_eq!(agent.current_state().as_deref(), Some("billing"));
20838 assert_eq!(response.content, "Billing state response");
20839 assert_eq!(call_counter.call_count(), 1);
20840 assert_eq!(agent.actor_facts().len(), 0);
20841 }
20842
20843 #[tokio::test]
20844 async fn test_set_context_supports_dotted_paths_for_pre_response_guards() {
20845 let mock = mock_with_response("Billing state response");
20846 let call_counter = mock.clone();
20847 let yaml = r#"
20848name: OptimizedStateAgent
20849system_prompt: "You route before answering."
20850runtime:
20851 optimization:
20852 enabled: true
20853 pre_response_deterministic_transitions: true
20854context:
20855 request:
20856 type: runtime
20857 default:
20858 topic: general
20859states:
20860 initial: greeting
20861 states:
20862 greeting:
20863 prompt: "Old state prompt that should be skipped."
20864 transitions:
20865 - to: billing
20866 guard:
20867 context:
20868 request.topic:
20869 eq: billing
20870 timing: pre_response
20871 billing:
20872 prompt: "Answer from the billing state."
20873"#;
20874 let agent = AgentBuilder::from_yaml(yaml)
20875 .unwrap()
20876 .llm(Arc::new(mock))
20877 .build()
20878 .unwrap();
20879 agent
20880 .set_context("request.topic", serde_json::json!("billing"))
20881 .unwrap();
20882
20883 let response = agent.chat("I need billing help").await.unwrap();
20884
20885 assert_eq!(agent.current_state().as_deref(), Some("billing"));
20886 assert_eq!(response.content, "Billing state response");
20887 assert_eq!(call_counter.call_count(), 1);
20888 assert_eq!(
20889 agent.get_context().get("request"),
20890 Some(&serde_json::json!({"topic": "billing"}))
20891 );
20892 }
20893
20894 #[tokio::test]
20895 async fn test_pre_response_rejection_does_not_commit_staged_context_or_user() {
20896 let mock = mock_with_response("billing");
20897 let yaml = r#"
20898name: OptimizedStateAgent
20899system_prompt: "You route before answering."
20900runtime:
20901 optimization:
20902 enabled: true
20903 pre_response_deterministic_transitions: true
20904hitl:
20905 states:
20906 billing:
20907 on_enter: require_approval
20908 approval_message: "Approve billing route?"
20909states:
20910 initial: greeting
20911 states:
20912 greeting:
20913 prompt: "Old state prompt."
20914 extract:
20915 - key: topic
20916 description: "Support topic"
20917 transitions:
20918 - to: billing
20919 guard:
20920 context:
20921 topic:
20922 eq: billing
20923 timing: pre_response
20924 run_extractors: true
20925 billing:
20926 prompt: "Billing state."
20927"#;
20928 let agent = AgentBuilder::from_yaml(yaml)
20929 .unwrap()
20930 .llm(Arc::new(mock))
20931 .build()
20932 .unwrap();
20933
20934 let response = agent
20935 .try_pre_response_transition("billing please")
20936 .await
20937 .unwrap();
20938
20939 assert!(response.is_none());
20940 assert_eq!(agent.current_state().as_deref(), Some("greeting"));
20941 assert!(!agent.get_context().contains_key("topic"));
20942 assert_eq!(agent.memory.get_messages(None).await.unwrap().len(), 0);
20943 }
20944
20945 #[tokio::test]
20946 async fn test_pre_response_extractor_commits_context_on_winning_path() {
20947 let mock = mock_with_responses(vec!["billing", "Billing response"]);
20948 let yaml = r#"
20949name: OptimizedStateAgent
20950system_prompt: "You route before answering."
20951runtime:
20952 optimization:
20953 enabled: true
20954 pre_response_deterministic_transitions: true
20955states:
20956 initial: greeting
20957 states:
20958 greeting:
20959 prompt: "Old state prompt."
20960 extract:
20961 - key: topic
20962 description: "Support topic"
20963 transitions:
20964 - to: billing
20965 guard:
20966 context:
20967 topic:
20968 eq: billing
20969 timing: pre_response
20970 run_extractors: true
20971 billing:
20972 prompt: "Billing state."
20973"#;
20974 let agent = AgentBuilder::from_yaml(yaml)
20975 .unwrap()
20976 .llm(Arc::new(mock))
20977 .build()
20978 .unwrap();
20979
20980 let response = agent.chat("billing please").await.unwrap();
20981
20982 assert_eq!(agent.current_state().as_deref(), Some("billing"));
20983 assert_eq!(response.content, "Billing response");
20984 assert_eq!(
20985 agent.get_context().get("topic"),
20986 Some(&serde_json::json!("billing"))
20987 );
20988 }
20989
20990 #[tokio::test]
20991 async fn test_pre_response_extractor_miss_does_not_mutate_context() {
20992 let mock = mock_with_response("__NONE__");
20993 let yaml = r#"
20994name: OptimizedStateAgent
20995system_prompt: "You route before answering."
20996runtime:
20997 optimization:
20998 enabled: true
20999 pre_response_deterministic_transitions: true
21000states:
21001 initial: greeting
21002 states:
21003 greeting:
21004 prompt: "Old state prompt."
21005 extract:
21006 - key: topic
21007 description: "Support topic"
21008 transitions:
21009 - to: billing
21010 guard:
21011 context:
21012 topic:
21013 eq: billing
21014 timing: pre_response
21015 run_extractors: true
21016 billing:
21017 prompt: "Billing state."
21018"#;
21019 let agent = AgentBuilder::from_yaml(yaml)
21020 .unwrap()
21021 .llm(Arc::new(mock))
21022 .build()
21023 .unwrap();
21024
21025 let response = agent.try_pre_response_transition("hello").await.unwrap();
21026
21027 assert!(response.is_none());
21028 assert_eq!(agent.current_state().as_deref(), Some("greeting"));
21029 assert!(!agent.get_context().contains_key("topic"));
21030 }
21031
21032 #[tokio::test]
21033 async fn test_default_guard_transition_stays_post_response() {
21034 let mock = mock_with_responses(vec!["Greeting response", "Billing response"]);
21035 let call_counter = mock.clone();
21036 let yaml = r#"
21037name: TimingAgent
21038system_prompt: "You route carefully."
21039runtime:
21040 optimization:
21041 enabled: true
21042 pre_response_deterministic_transitions: true
21043states:
21044 initial: greeting
21045 states:
21046 greeting:
21047 prompt: "Old state prompt."
21048 transitions:
21049 - to: billing
21050 guard:
21051 context:
21052 topic:
21053 eq: billing
21054 billing:
21055 prompt: "Billing state."
21056"#;
21057 let agent = AgentBuilder::from_yaml(yaml)
21058 .unwrap()
21059 .llm(Arc::new(mock))
21060 .build()
21061 .unwrap();
21062 agent
21063 .set_context("topic", serde_json::json!("billing"))
21064 .unwrap();
21065
21066 let response = agent.chat("billing please").await.unwrap();
21067
21068 assert_eq!(agent.current_state().as_deref(), Some("billing"));
21069 assert_eq!(response.content, "Billing response");
21070 assert_eq!(call_counter.call_count(), 2);
21071 }
21072
21073 #[tokio::test]
21074 async fn test_explicit_post_response_guard_transition_stays_post_response() {
21075 let mock = mock_with_responses(vec!["Greeting response", "Billing response"]);
21076 let call_counter = mock.clone();
21077 let yaml = r#"
21078name: TimingAgent
21079system_prompt: "You route carefully."
21080runtime:
21081 optimization:
21082 enabled: true
21083 pre_response_deterministic_transitions: true
21084states:
21085 initial: greeting
21086 states:
21087 greeting:
21088 prompt: "Old state prompt."
21089 transitions:
21090 - to: billing
21091 guard:
21092 context:
21093 topic:
21094 eq: billing
21095 timing: post_response
21096 billing:
21097 prompt: "Billing state."
21098"#;
21099 let agent = AgentBuilder::from_yaml(yaml)
21100 .unwrap()
21101 .llm(Arc::new(mock))
21102 .build()
21103 .unwrap();
21104 agent
21105 .set_context("topic", serde_json::json!("billing"))
21106 .unwrap();
21107
21108 let response = agent.chat("billing please").await.unwrap();
21109
21110 assert_eq!(agent.current_state().as_deref(), Some("billing"));
21111 assert_eq!(response.content, "Billing response");
21112 assert_eq!(call_counter.call_count(), 2);
21113 }
21114
21115 #[tokio::test]
21116 async fn test_pre_response_extractors_are_transition_scoped() {
21117 let mock = mock_with_responses(vec!["billing", "Billing response"]);
21118 let yaml = r#"
21119name: ScopedExtractorAgent
21120system_prompt: "You route carefully."
21121runtime:
21122 optimization:
21123 enabled: true
21124 pre_response_deterministic_transitions: true
21125states:
21126 initial: greeting
21127 states:
21128 greeting:
21129 prompt: "Old state prompt."
21130 extract:
21131 - key: topic
21132 description: "Support topic"
21133 transitions:
21134 - to: wrong
21135 guard:
21136 context:
21137 topic:
21138 eq: billing
21139 timing: pre_response
21140 - to: billing
21141 guard:
21142 context:
21143 topic:
21144 eq: billing
21145 timing: pre_response
21146 run_extractors: true
21147 wrong:
21148 prompt: "Wrong state."
21149 billing:
21150 prompt: "Billing state."
21151"#;
21152 let agent = AgentBuilder::from_yaml(yaml)
21153 .unwrap()
21154 .llm(Arc::new(mock))
21155 .build()
21156 .unwrap();
21157
21158 let response = agent.chat("billing please").await.unwrap();
21159
21160 assert_eq!(agent.current_state().as_deref(), Some("billing"));
21161 assert_eq!(response.content, "Billing response");
21162 }
21163
21164 #[tokio::test]
21165 async fn test_pre_response_resolved_intent_routes_early() {
21166 let mock = mock_with_response("Billing response");
21167 let yaml = r#"
21168name: IntentAgent
21169system_prompt: "You route carefully."
21170runtime:
21171 optimization:
21172 enabled: true
21173 pre_response_deterministic_transitions: true
21174states:
21175 initial: greeting
21176 states:
21177 greeting:
21178 prompt: "Old state prompt."
21179 transitions:
21180 - to: billing
21181 intent: billing
21182 timing: pre_response
21183 billing:
21184 prompt: "Billing state."
21185"#;
21186 let agent = AgentBuilder::from_yaml(yaml)
21187 .unwrap()
21188 .llm(Arc::new(mock))
21189 .build()
21190 .unwrap();
21191 agent
21192 .set_context("resolved_intent", serde_json::json!("billing"))
21193 .unwrap();
21194
21195 let response = agent
21196 .try_pre_response_transition("I need billing help")
21197 .await
21198 .unwrap()
21199 .unwrap();
21200
21201 assert_eq!(agent.current_state().as_deref(), Some("billing"));
21202 assert_eq!(response.content, "Billing response");
21203 }
21204
21205 #[tokio::test]
21206 async fn test_background_overflow_error_surfaces() {
21207 let mut config = RuntimeConfig::default();
21208 config.optimization.enabled = true;
21209 config.optimization.post_turn.max_background_tasks = 1;
21210 config.optimization.post_turn.on_background_overflow = BackgroundOverflowPolicy::Error;
21211 let policy = crate::optimization::MaintenanceTaskPolicy {
21212 mode: MaintenanceMode::Background,
21213 await_before_next_turn: AwaitBeforeNextTurn::Always,
21214 };
21215 let agent = AgentBuilder::new()
21216 .system_prompt("You are helpful.")
21217 .llm(Arc::new(mock_with_response("ok")))
21218 .build()
21219 .unwrap()
21220 .with_runtime_config(config);
21221 agent
21222 .background_maintenance
21223 .spawn(None, async { std::future::pending::<Result<()>>().await })
21224 .unwrap();
21225
21226 let result = agent
21227 .spawn_or_handle_background(None, async { Ok(()) }, "facts", &policy)
21228 .await;
21229
21230 assert!(result.is_err());
21231 }
21232
21233 #[tokio::test]
21234 async fn test_speculative_reasoning_low_cap_uses_serial_reasoning() {
21235 let default_mock = mock_with_response("Plain draft response");
21236 let router_mock = mock_with_response("cot");
21237 let router_counter = router_mock.clone();
21238 let yaml = r#"
21239name: ReasoningReservationAgent
21240system_prompt: "You answer plainly unless reasoning wins."
21241llm:
21242 default: default
21243 router: router
21244observability:
21245 enabled: true
21246 export:
21247 write_raw_events: true
21248reasoning:
21249 mode: auto
21250 judge_llm: router
21251runtime:
21252 optimization:
21253 enabled: true
21254 max_speculative_llm_calls_per_turn: 1
21255 speculative_reasoning_auto: true
21256 max_parallel_runtime_tasks: 2
21257"#;
21258 let agent = AgentBuilder::from_yaml(yaml)
21259 .unwrap()
21260 .llm_alias("default", Arc::new(default_mock))
21261 .llm_alias("router", Arc::new(router_mock))
21262 .build()
21263 .unwrap();
21264
21265 let response = agent.chat("hello").await.unwrap();
21266
21267 assert_eq!(response.content, "Plain draft response");
21268 assert_eq!(router_counter.call_count(), 1);
21269 let events = agent.observability().unwrap().raw_events();
21270 assert!(!events.iter().any(|event| {
21271 event.dimensions.get("commit_behavior") == Some(&"reasoning_decision".to_string())
21272 }));
21273 }
21274
21275 #[tokio::test]
21276 async fn test_forced_reasoning_skips_plain_speculative_draft() {
21277 let mock = mock_with_response("Reasoned response");
21278 let yaml = r#"
21279name: ForcedReasoningAgent
21280system_prompt: "You reason before answering."
21281observability:
21282 enabled: true
21283 export:
21284 write_raw_events: true
21285reasoning:
21286 mode: cot
21287runtime:
21288 optimization:
21289 enabled: true
21290 max_speculative_llm_calls_per_turn: 2
21291 speculative_state_transitions: true
21292 max_parallel_runtime_tasks: 2
21293states:
21294 initial: triage
21295 states:
21296 triage:
21297 prompt: "Answer from triage."
21298 transitions:
21299 - to: billing
21300 guard:
21301 context:
21302 route:
21303 eq: billing
21304 timing: parallel
21305 billing:
21306 prompt: "Billing state."
21307"#;
21308 let agent = AgentBuilder::from_yaml(yaml)
21309 .unwrap()
21310 .llm(Arc::new(mock))
21311 .build()
21312 .unwrap();
21313
21314 let response = agent.chat("hello").await.unwrap();
21315
21316 assert_eq!(response.content, "Reasoned response");
21317 let events = agent.observability().unwrap().raw_events();
21318 assert!(
21319 !events
21320 .iter()
21321 .any(|event| event.dimensions.contains_key("branch_status"))
21322 );
21323 }
21324
21325 #[tokio::test]
21326 async fn test_speculative_skill_low_cap_uses_serial_skill_route() {
21327 let default_mock = mock_with_response("Skill committed response");
21328 let router_mock = mock_with_response("helper");
21329 let router_counter = router_mock.clone();
21330 let yaml = r#"
21331name: SkillReservationAgent
21332system_prompt: "Use skills when they match."
21333llm:
21334 default: default
21335 router: router
21336observability:
21337 enabled: true
21338 export:
21339 write_raw_events: true
21340runtime:
21341 optimization:
21342 enabled: true
21343 max_speculative_llm_calls_per_turn: 1
21344 speculative_skill_routing: true
21345 max_parallel_runtime_tasks: 2
21346skills:
21347 - id: helper
21348 description: "Answer helper requests"
21349 trigger: "User asks for helper"
21350 steps:
21351 - prompt: "Answer the helper request: {{ user_input }}"
21352"#;
21353 let agent = AgentBuilder::from_yaml(yaml)
21354 .unwrap()
21355 .llm_alias("default", Arc::new(default_mock))
21356 .llm_alias("router", Arc::new(router_mock))
21357 .build()
21358 .unwrap();
21359
21360 let response = agent.chat("please use helper").await.unwrap();
21361
21362 assert_eq!(response.content, "Skill committed response");
21363 assert_eq!(router_counter.call_count(), 1);
21364 let events = agent.observability().unwrap().raw_events();
21365 assert!(
21366 !events
21367 .iter()
21368 .any(|event| event.dimensions.contains_key("branch_status"))
21369 );
21370 }
21371
21372 #[tokio::test]
21373 async fn test_parallel_transition_low_cap_allows_deterministic_route() {
21374 let mock = mock_with_response("unused");
21375 let call_counter = mock.clone();
21376 let yaml = r#"
21377name: ParallelTransitionLowCapAgent
21378system_prompt: "Route before stale responses when safe."
21379runtime:
21380 optimization:
21381 enabled: true
21382 max_speculative_llm_calls_per_turn: 1
21383 speculative_state_transitions: true
21384 max_parallel_runtime_tasks: 2
21385states:
21386 initial: triage
21387 states:
21388 triage:
21389 prompt: "Triage state."
21390 transitions:
21391 - to: billing
21392 guard:
21393 context:
21394 route:
21395 eq: billing
21396 timing: parallel
21397 billing:
21398 prompt: "Billing state."
21399"#;
21400 let agent = AgentBuilder::from_yaml(yaml)
21401 .unwrap()
21402 .llm(Arc::new(mock))
21403 .build()
21404 .unwrap();
21405 agent
21406 .set_context("route", serde_json::json!("billing"))
21407 .unwrap();
21408 agent.update_active_turn_context("billing help", HashMap::new());
21409 assert!(
21410 agent.reserve_active_speculative_llm_call(
21411 RuntimeOptimizationKind::ParallelStateTransition
21412 )
21413 );
21414
21415 let selection = agent
21416 .select_parallel_transition_candidate("billing help")
21417 .await
21418 .unwrap();
21419 agent.end_root_turn();
21420
21421 match selection {
21422 ParallelTransitionSelection::Candidate(candidate) => {
21423 assert_eq!(candidate.target(), "billing");
21424 }
21425 ParallelTransitionSelection::NoMatch => panic!("deterministic route did not match"),
21426 ParallelTransitionSelection::ReservationExhausted => {
21427 panic!("deterministic route consumed LLM budget")
21428 }
21429 }
21430 assert_eq!(call_counter.call_count(), 0);
21431 }
21432
21433 #[tokio::test]
21434 async fn speculative_transition_drops_loser_before_state_actions() {
21435 let lock = Arc::new(tokio::sync::Mutex::new(()));
21436 let first_started = Arc::new(tokio::sync::Notify::new());
21437 let first_dropped = Arc::new(AtomicBool::new(false));
21438 let committed_after_drop = Arc::new(AtomicBool::new(false));
21439 let default = Arc::new(FirstCallLockingProvider {
21440 lock,
21441 first_started: Arc::clone(&first_started),
21442 first_dropped: Arc::clone(&first_dropped),
21443 committed_after_drop: Arc::clone(&committed_after_drop),
21444 calls: AtomicU64::new(0),
21445 });
21446 let router = Arc::new(RoutingAfterProviderStart {
21447 provider_started: first_started,
21448 });
21449 let yaml = r#"
21450name: SpeculativeCancellationAgent
21451system_prompt: "Route before committed work."
21452llm:
21453 default: default
21454 router: router
21455runtime:
21456 optimization:
21457 enabled: true
21458 max_speculative_llm_calls_per_turn: 2
21459 speculative_state_transitions: true
21460 max_parallel_runtime_tasks: 2
21461states:
21462 initial: triage
21463 states:
21464 triage:
21465 prompt: "Triage state."
21466 transitions:
21467 - to: technical
21468 when: "The request needs technical support"
21469 timing: parallel
21470 technical:
21471 prompt: "Technical state."
21472 on_enter:
21473 - prompt: "Prepare technical context."
21474 llm: default
21475 store_as: preparation
21476"#;
21477 let agent = AgentBuilder::from_yaml(yaml)
21478 .unwrap()
21479 .llm_alias("default", default)
21480 .llm_alias("router", router)
21481 .build()
21482 .unwrap();
21483
21484 let response = tokio::time::timeout(
21485 std::time::Duration::from_secs(2),
21486 agent.chat("I cannot log in because of AUTH-17."),
21487 )
21488 .await
21489 .expect("committed work must not wait on the losing provider future")
21490 .unwrap();
21491
21492 assert_eq!(response.content, "Committed technical response.");
21493 assert_eq!(agent.current_state().as_deref(), Some("technical"));
21494 assert!(first_dropped.load(Ordering::SeqCst));
21495 assert!(committed_after_drop.load(Ordering::SeqCst));
21496 }
21497
21498 #[tokio::test]
21499 async fn buffered_transition_drops_stale_stream_before_redispatch() {
21500 use futures::StreamExt;
21501
21502 let lock = Arc::new(tokio::sync::Mutex::new(()));
21503 let stream_started = Arc::new(tokio::sync::Notify::new());
21504 let stream_dropped = Arc::new(AtomicBool::new(false));
21505 let committed_after_drop = Arc::new(AtomicBool::new(false));
21506 let default = Arc::new(BufferedLockingProvider {
21507 lock,
21508 stream_started: Arc::clone(&stream_started),
21509 stream_dropped: Arc::clone(&stream_dropped),
21510 committed_after_drop: Arc::clone(&committed_after_drop),
21511 });
21512 let router = Arc::new(RoutingAfterProviderStart {
21513 provider_started: stream_started,
21514 });
21515 let yaml = r#"
21516name: BufferedCancellationAgent
21517system_prompt: "Hide stale streamed output."
21518llm:
21519 default: default
21520 router: router
21521streaming:
21522 enabled: true
21523 buffer_size: 8
21524runtime:
21525 optimization:
21526 enabled: true
21527 max_speculative_llm_calls_per_turn: 2
21528 speculative_state_transitions: true
21529 streaming_policy: buffer_until_routing_done
21530 max_parallel_runtime_tasks: 2
21531states:
21532 initial: triage
21533 states:
21534 triage:
21535 prompt: "Triage state."
21536 transitions:
21537 - to: technical
21538 when: "The request needs technical support"
21539 timing: parallel
21540 technical:
21541 prompt: "Technical state."
21542"#;
21543 let agent = AgentBuilder::from_yaml(yaml)
21544 .unwrap()
21545 .llm_alias("default", default)
21546 .llm_alias("router", router)
21547 .build()
21548 .unwrap();
21549
21550 let content = tokio::time::timeout(std::time::Duration::from_secs(2), async {
21551 let mut stream = agent
21552 .chat_stream("AUTH-17 needs technical help.")
21553 .await
21554 .unwrap();
21555 let mut content = String::new();
21556 while let Some(chunk) = stream.next().await {
21557 match chunk {
21558 StreamChunk::Content { text } => content.push_str(&text),
21559 StreamChunk::Done {} => break,
21560 StreamChunk::Error { message } => panic!("unexpected stream error: {message}"),
21561 _ => {}
21562 }
21563 }
21564 content
21565 })
21566 .await
21567 .expect("redispatch must not wait on the stale streaming future");
21568
21569 assert_eq!(content, "Committed technical response.");
21570 assert_eq!(agent.current_state().as_deref(), Some("technical"));
21571 assert!(stream_dropped.load(Ordering::SeqCst));
21572 assert!(committed_after_drop.load(Ordering::SeqCst));
21573 }
21574
21575 #[tokio::test]
21576 async fn buffered_transition_drops_established_stream_before_redispatch() {
21577 use futures::StreamExt;
21578
21579 let stream_started = Arc::new(tokio::sync::Notify::new());
21580 let stream_dropped = Arc::new(AtomicBool::new(false));
21581 let stream_dropped_notify = Arc::new(tokio::sync::Notify::new());
21582 let committed_after_drop = Arc::new(AtomicBool::new(false));
21583 let default = Arc::new(EstablishedStreamProvider {
21584 stream_started: Arc::clone(&stream_started),
21585 stream_dropped: Arc::clone(&stream_dropped),
21586 stream_dropped_notify,
21587 committed_after_drop: Arc::clone(&committed_after_drop),
21588 });
21589 let router = Arc::new(RoutingAfterProviderStart {
21590 provider_started: stream_started,
21591 });
21592 let yaml = r#"
21593name: EstablishedStreamCancellationAgent
21594system_prompt: "Hide stale streamed output."
21595llm:
21596 default: default
21597 router: router
21598streaming:
21599 enabled: true
21600 buffer_size: 8
21601runtime:
21602 optimization:
21603 enabled: true
21604 max_speculative_llm_calls_per_turn: 2
21605 speculative_state_transitions: true
21606 streaming_policy: buffer_until_routing_done
21607 max_parallel_runtime_tasks: 2
21608states:
21609 initial: triage
21610 states:
21611 triage:
21612 prompt: "Triage state."
21613 transitions:
21614 - to: technical
21615 when: "The request needs technical support"
21616 timing: parallel
21617 technical:
21618 prompt: "Technical state."
21619"#;
21620 let agent = AgentBuilder::from_yaml(yaml)
21621 .unwrap()
21622 .llm_alias("default", default)
21623 .llm_alias("router", router)
21624 .build()
21625 .unwrap();
21626
21627 let content = tokio::time::timeout(std::time::Duration::from_secs(2), async {
21628 let mut stream = agent
21629 .chat_stream("AUTH-17 needs technical help.")
21630 .await
21631 .unwrap();
21632 let mut content = String::new();
21633 while let Some(chunk) = stream.next().await {
21634 match chunk {
21635 StreamChunk::Content { text } => content.push_str(&text),
21636 StreamChunk::Done {} => break,
21637 StreamChunk::Error { message } => panic!("unexpected stream error: {message}"),
21638 _ => {}
21639 }
21640 }
21641 content
21642 })
21643 .await
21644 .expect("redispatch must wait for the established stale stream to be dropped");
21645
21646 assert_eq!(content, "Committed technical response.");
21647 assert_eq!(agent.current_state().as_deref(), Some("technical"));
21648 assert!(stream_dropped.load(Ordering::SeqCst));
21649 assert!(committed_after_drop.load(Ordering::SeqCst));
21650 }
21651
21652 #[tokio::test]
21653 async fn test_buffered_streaming_transition_reservation_falls_back() {
21654 use futures::StreamExt;
21655
21656 let mock = mock_with_responses(vec![
21657 "Serial streaming response",
21658 "Serial streaming response",
21659 ]);
21660 let router_mock = mock_with_response("1");
21661 let router_counter = router_mock.clone();
21662 let yaml = r#"
21663name: BufferedReservationFallbackAgent
21664system_prompt: "Stream normally if speculative routing cannot be evaluated."
21665llm:
21666 default: default
21667 router: router
21668observability:
21669 enabled: true
21670 export:
21671 write_raw_events: true
21672streaming:
21673 enabled: true
21674 buffer_size: 8
21675runtime:
21676 optimization:
21677 enabled: true
21678 max_speculative_llm_calls_per_turn: 1
21679 speculative_state_transitions: true
21680 streaming_policy: buffer_until_routing_done
21681 max_parallel_runtime_tasks: 2
21682states:
21683 initial: triage
21684 states:
21685 triage:
21686 prompt: "Triage state."
21687 transitions:
21688 - to: billing
21689 guard:
21690 context:
21691 route:
21692 eq: billing
21693 when: "User asks about billing"
21694 timing: parallel
21695 billing:
21696 prompt: "Billing state."
21697"#;
21698 let agent = AgentBuilder::from_yaml(yaml)
21699 .unwrap()
21700 .llm_alias("default", Arc::new(mock))
21701 .llm_alias("router", Arc::new(router_mock))
21702 .build()
21703 .unwrap();
21704
21705 let mut stream = agent.chat_stream("hello").await.unwrap();
21706 let mut content = String::new();
21707 let mut error = None;
21708 while let Some(chunk) = stream.next().await {
21709 match chunk {
21710 StreamChunk::Content { text } => content.push_str(&text),
21711 StreamChunk::Error { message } => error = Some(message),
21712 StreamChunk::Done {} => break,
21713 _ => {}
21714 }
21715 }
21716
21717 assert_eq!(error, None);
21718 assert_eq!(content, "Serial streaming response");
21719 assert_eq!(router_counter.call_count(), 0);
21720 let events = agent.observability().unwrap().raw_events();
21721 assert!(events.iter().any(|event| {
21722 event.dimensions.get("branch_status") == Some(&"cancelled".to_string())
21723 && event.dimensions.get("commit_behavior")
21724 == Some(&"transition_decision".to_string())
21725 }));
21726 }
21727
21728 #[tokio::test]
21729 async fn test_blocking_error_cleanup_resets_root_turn_for_next_chat() {
21730 let mut mock = mock_with_response("Recovered response");
21731 mock.set_error("boom");
21732 let mut handle = mock.clone();
21733 let agent = AgentBuilder::new()
21734 .system_prompt("You are helpful.")
21735 .llm(Arc::new(mock))
21736 .build()
21737 .unwrap();
21738
21739 assert!(agent.chat("first").await.is_err());
21740 handle.clear_error();
21741 let response = agent.chat("second").await.unwrap();
21742
21743 assert_eq!(response.content, "Recovered response");
21744 let messages = agent.memory.get_messages(None).await.unwrap();
21745 let user_count = messages
21746 .iter()
21747 .filter(|message| message.role == ai_agents_core::Role::User)
21748 .count();
21749 assert_eq!(user_count, 2);
21750 }
21751
21752 #[tokio::test]
21753 async fn test_streaming_error_cleanup_resets_root_turn_for_next_chat() {
21754 use futures::StreamExt;
21755
21756 let mut mock = mock_with_response("Recovered response");
21757 mock.set_error("stream boom");
21758 let mut handle = mock.clone();
21759 let agent = AgentBuilder::new()
21760 .system_prompt("You are helpful.")
21761 .llm(Arc::new(mock))
21762 .build()
21763 .unwrap();
21764
21765 let mut stream = agent.chat_stream("first").await.unwrap();
21766 let mut saw_error = false;
21767 while let Some(chunk) = stream.next().await {
21768 if matches!(chunk, StreamChunk::Error { .. }) {
21769 saw_error = true;
21770 }
21771 }
21772 assert!(saw_error);
21773
21774 handle.clear_error();
21775 let response = agent.chat("second").await.unwrap();
21776
21777 assert_eq!(response.content, "Recovered response");
21778 let messages = agent.memory.get_messages(None).await.unwrap();
21779 let user_count = messages
21780 .iter()
21781 .filter(|message| message.role == ai_agents_core::Role::User)
21782 .count();
21783 assert_eq!(user_count, 2);
21784 }
21785
21786 #[tokio::test]
21787 async fn test_buffered_streaming_route_miss_releases_buffer_limit() {
21788 use futures::StreamExt;
21789
21790 let mut mock = mock_with_response("one two three");
21791 mock.set_latency(10);
21792 let yaml = r#"
21793name: BufferedMissAgent
21794system_prompt: "You stream safely."
21795llm:
21796 default: default
21797streaming:
21798 enabled: true
21799 buffer_size: 1
21800runtime:
21801 optimization:
21802 enabled: true
21803 max_speculative_llm_calls_per_turn: 2
21804 speculative_state_transitions: true
21805 streaming_policy: buffer_until_routing_done
21806 max_parallel_runtime_tasks: 2
21807states:
21808 initial: triage
21809 states:
21810 triage:
21811 prompt: "Answer from triage."
21812 transitions:
21813 - to: billing
21814 guard:
21815 context:
21816 route:
21817 eq: billing
21818 timing: parallel
21819 billing:
21820 prompt: "Billing state."
21821"#;
21822 let agent = AgentBuilder::from_yaml(yaml)
21823 .unwrap()
21824 .llm_alias("default", Arc::new(mock))
21825 .build()
21826 .unwrap();
21827
21828 let mut stream = agent.chat_stream("hello").await.unwrap();
21829 let mut content = String::new();
21830 let mut error = None;
21831 while let Some(chunk) = stream.next().await {
21832 match chunk {
21833 StreamChunk::Content { text } => content.push_str(&text),
21834 StreamChunk::Error { message } => error = Some(message),
21835 StreamChunk::Done {} => break,
21836 _ => {}
21837 }
21838 }
21839
21840 assert_eq!(error, None);
21841 assert_eq!(content, "one two three");
21842 }
21843
21844 #[tokio::test]
21845 async fn test_buffered_streaming_main_failure_finalizes_branch() {
21846 use futures::StreamExt;
21847
21848 let mock = mock_with_response("one two");
21849 let mut router_mock = mock_with_response("0");
21850 router_mock.set_latency(50);
21851 let yaml = r#"
21852name: BufferedFailureAgent
21853system_prompt: "You stream safely."
21854llm:
21855 default: default
21856 router: router
21857observability:
21858 enabled: true
21859 export:
21860 write_raw_events: true
21861streaming:
21862 enabled: true
21863 buffer_size: 1
21864runtime:
21865 optimization:
21866 enabled: true
21867 max_speculative_llm_calls_per_turn: 2
21868 speculative_state_transitions: true
21869 streaming_policy: buffer_until_routing_done
21870 max_parallel_runtime_tasks: 2
21871states:
21872 initial: triage
21873 states:
21874 triage:
21875 prompt: "Ask for the category."
21876 transitions:
21877 - to: billing
21878 when: "User asks about billing"
21879 timing: parallel
21880 billing:
21881 prompt: "Billing state."
21882"#;
21883 let agent = AgentBuilder::from_yaml(yaml)
21884 .unwrap()
21885 .llm_alias("default", Arc::new(mock))
21886 .llm_alias("router", Arc::new(router_mock))
21887 .build()
21888 .unwrap();
21889
21890 let mut stream = agent.chat_stream("hello").await.unwrap();
21891 let mut error = String::new();
21892 while let Some(chunk) = stream.next().await {
21893 if let StreamChunk::Error { message } = chunk {
21894 error = message;
21895 }
21896 }
21897
21898 assert!(
21899 error.contains("stream buffer filled"),
21900 "unexpected stream error: {}",
21901 error
21902 );
21903 let events = agent.observability().unwrap().raw_events();
21904 assert!(events.iter().any(|event| {
21905 event.dimensions.get("branch_status") == Some(&"failed".to_string())
21906 && event.dimensions.get("commit_behavior") == Some(&"final_response".to_string())
21907 && event.dimensions.get("optimization")
21908 == Some(&"buffered_streaming_routing".to_string())
21909 }));
21910 }
21911
21912 #[tokio::test]
21913 async fn test_streaming_preflight_does_not_emit_old_state_content() {
21914 use futures::StreamExt;
21915
21916 let mock = mock_with_response("Billing streamed response");
21917 let yaml = r#"
21918name: StreamingOptimizedAgent
21919system_prompt: "You route before streaming."
21920runtime:
21921 optimization:
21922 enabled: true
21923 pre_response_deterministic_transitions: true
21924streaming:
21925 enabled: true
21926states:
21927 initial: greeting
21928 states:
21929 greeting:
21930 prompt: "OLD_STATE_SENTINEL"
21931 transitions:
21932 - to: billing
21933 guard:
21934 context:
21935 topic:
21936 eq: billing
21937 timing: pre_response
21938 billing:
21939 prompt: "Billing state."
21940"#;
21941 let agent = AgentBuilder::from_yaml(yaml)
21942 .unwrap()
21943 .llm(Arc::new(mock))
21944 .build()
21945 .unwrap();
21946 agent
21947 .set_context("topic", serde_json::json!("billing"))
21948 .unwrap();
21949
21950 let mut stream = agent.chat_stream("billing please").await.unwrap();
21951 let mut content = String::new();
21952 while let Some(chunk) = stream.next().await {
21953 match chunk {
21954 StreamChunk::Content { text } => content.push_str(&text),
21955 StreamChunk::Error { message } => panic!("stream error: {}", message),
21956 StreamChunk::Done {} => break,
21957 _ => {}
21958 }
21959 }
21960
21961 assert_eq!(agent.current_state().as_deref(), Some("billing"));
21962 assert!(content.contains("Billing streamed response"));
21963 assert!(!content.contains("OLD_STATE_SENTINEL"));
21964 }
21965
21966 #[tokio::test]
21968 async fn test_integration_state_machine_basic() {
21969 let yaml = r#"
21970name: StateAgent
21971system_prompt: "You are a support agent."
21972states:
21973 initial: greeting
21974 states:
21975 greeting:
21976 prompt: "Welcome the user warmly."
21977 transitions:
21978 - to: support
21979 when: "User needs help"
21980 auto: true
21981 support:
21982 prompt: "Help solve the user's problem."
21983"#;
21984 let mock = mock_with_responses(vec![
21985 "Welcome! How can I help?", "1", "I'll help you with that.", ]);
21989 let builder = AgentBuilder::from_yaml(yaml).unwrap();
21990 let agent = builder.llm(Arc::new(mock)).build().unwrap();
21991
21992 assert_eq!(agent.current_state(), Some("greeting".to_string()));
21993 let _ = agent.chat("I need help").await.unwrap();
21994 }
21997
21998 #[tokio::test]
22000 async fn test_integration_state_on_enter_set_context() {
22001 let yaml = r#"
22002name: ActionAgent
22003system_prompt: "You are helpful."
22004states:
22005 initial: step1
22006 states:
22007 step1:
22008 prompt: "Step 1"
22009 on_exit:
22010 - set_context:
22011 step1_exited: true
22012 transitions:
22013 - to: step2
22014 when: "always"
22015 auto: true
22016 step2:
22017 prompt: "Step 2"
22018 on_enter:
22019 - set_context:
22020 step2_entered: true
22021"#;
22022 let mock = mock_with_responses(vec![
22024 "Processing step 1.",
22025 "0", ]);
22027 let builder = AgentBuilder::from_yaml(yaml).unwrap();
22028 let agent = builder.llm(Arc::new(mock)).build().unwrap();
22029
22030 assert_eq!(agent.current_state(), Some("step1".to_string()));
22031
22032 agent.transition_to("step2").await.unwrap();
22034
22035 assert_eq!(agent.current_state(), Some("step2".to_string()));
22036
22037 let ctx = agent.get_context();
22039 assert_eq!(ctx.get("step1_exited"), Some(&serde_json::json!(true)));
22040 assert_eq!(ctx.get("step2_entered"), Some(&serde_json::json!(true)));
22041 }
22042
22043 #[tokio::test]
22044 async fn state_action_tool_preserves_source_in_stored_record() {
22045 let yaml = r#"
22046name: StateActionToolAgent
22047system_prompt: "You are helpful."
22048tools:
22049 - context_echo
22050states:
22051 initial: idle
22052 states:
22053 idle:
22054 prompt: "Idle"
22055 active:
22056 prompt: "Active"
22057 on_enter:
22058 - set_context:
22059 action_started: true
22060 - tool: context_echo
22061 args: {}
22062"#;
22063 let agent = AgentBuilder::from_yaml(yaml)
22064 .unwrap()
22065 .llm(Arc::new(mock_with_response("unused")))
22066 .tool(Arc::new(ContextEchoTool))
22067 .build()
22068 .unwrap();
22069
22070 agent.transition_to("active").await.unwrap();
22071
22072 let record: ToolExecutionRecord = serde_json::from_value(
22073 agent
22074 .get_context()
22075 .get("last_tool_record")
22076 .cloned()
22077 .expect("successful state action must store its execution record"),
22078 )
22079 .unwrap();
22080 assert!(record.executed);
22081 assert!(record.success);
22082 assert_eq!(record.canonical_id, "context_echo");
22083 assert!(matches!(
22084 &record.source,
22085 ToolCallSource::StateAction {
22086 state: Some(state),
22087 action_index: 1,
22088 } if state == "active"
22089 ));
22090 }
22091
22092 #[tokio::test]
22093 async fn test_ordinary_transition_uses_on_enter_then_on_reenter() {
22094 let yaml = r#"
22095name: OrdinaryLifecycleAgent
22096system_prompt: "You are helpful."
22097states:
22098 initial: intake
22099 regenerate_on_transition: false
22100 states:
22101 intake:
22102 prompt: "Intake"
22103 transitions:
22104 - to: drafting
22105 guard:
22106 context:
22107 route:
22108 eq: drafting
22109 drafting:
22110 prompt: "Drafting"
22111 on_enter:
22112 - set_context:
22113 draft_version: 1
22114 on_reenter:
22115 - set_context:
22116 draft_version: 2
22117 transitions:
22118 - to: review
22119 guard:
22120 context:
22121 route:
22122 eq: review
22123 review:
22124 prompt: "Review"
22125 on_enter:
22126 - set_context:
22127 review_entry: first
22128 transitions:
22129 - to: drafting
22130 guard:
22131 context:
22132 route:
22133 eq: drafting
22134"#;
22135 let agent = AgentBuilder::from_yaml(yaml)
22136 .unwrap()
22137 .llm(Arc::new(mock_with_responses(vec![
22138 "Intake response",
22139 "Draft response",
22140 "Review response",
22141 ])))
22142 .build()
22143 .unwrap();
22144
22145 agent
22146 .set_context("route", serde_json::json!("drafting"))
22147 .unwrap();
22148 agent.chat("Start a draft").await.unwrap();
22149 assert_eq!(agent.current_state().as_deref(), Some("drafting"));
22150 assert_eq!(
22151 agent.get_context().get("draft_version"),
22152 Some(&serde_json::json!(1))
22153 );
22154
22155 agent
22156 .set_context("route", serde_json::json!("review"))
22157 .unwrap();
22158 agent.chat("Review this").await.unwrap();
22159 assert_eq!(agent.current_state().as_deref(), Some("review"));
22160 assert_eq!(
22161 agent.get_context().get("review_entry"),
22162 Some(&serde_json::json!("first"))
22163 );
22164
22165 agent
22166 .set_context("route", serde_json::json!("drafting"))
22167 .unwrap();
22168 agent.chat("Revise this").await.unwrap();
22169 assert_eq!(agent.current_state().as_deref(), Some("drafting"));
22170 assert_eq!(
22171 agent.get_context().get("draft_version"),
22172 Some(&serde_json::json!(2))
22173 );
22174 }
22175
22176 #[tokio::test]
22177 async fn test_manual_transition_uses_on_enter_then_on_reenter() {
22178 let yaml = r#"
22179name: ManualLifecycleAgent
22180system_prompt: "You are helpful."
22181states:
22182 initial: intake
22183 states:
22184 intake:
22185 prompt: "Intake"
22186 drafting:
22187 prompt: "Drafting"
22188 on_enter:
22189 - set_context:
22190 draft_version: 1
22191 on_reenter:
22192 - set_context:
22193 draft_version: 2
22194 review:
22195 prompt: "Review"
22196"#;
22197 let agent = AgentBuilder::from_yaml(yaml)
22198 .unwrap()
22199 .llm(Arc::new(mock_with_response("unused")))
22200 .build()
22201 .unwrap();
22202
22203 assert!(!agent.get_context().contains_key("draft_version"));
22204 agent.transition_to("drafting").await.unwrap();
22205 assert_eq!(agent.current_state().as_deref(), Some("drafting"));
22206 assert_eq!(
22207 agent.get_context().get("draft_version"),
22208 Some(&serde_json::json!(1))
22209 );
22210
22211 agent.transition_to("review").await.unwrap();
22212 agent.transition_to("drafting").await.unwrap();
22213 assert_eq!(agent.current_state().as_deref(), Some("drafting"));
22214 assert_eq!(
22215 agent.get_context().get("draft_version"),
22216 Some(&serde_json::json!(2))
22217 );
22218 }
22219
22220 #[tokio::test]
22221 async fn test_timeout_transition_uses_on_enter_then_on_reenter() {
22222 let yaml = r#"
22223name: TimeoutLifecycleAgent
22224system_prompt: "You are helpful."
22225states:
22226 initial: intake
22227 regenerate_on_transition: false
22228 states:
22229 intake:
22230 prompt: "Intake"
22231 max_turns: 1
22232 timeout_to: drafting
22233 drafting:
22234 prompt: "Drafting"
22235 max_turns: 1
22236 timeout_to: review
22237 on_enter:
22238 - set_context:
22239 draft_version: 1
22240 on_reenter:
22241 - set_context:
22242 draft_version: 2
22243 review:
22244 prompt: "Review"
22245 max_turns: 1
22246 timeout_to: drafting
22247 on_enter:
22248 - set_context:
22249 review_entry: first
22250"#;
22251 let agent = AgentBuilder::from_yaml(yaml)
22252 .unwrap()
22253 .llm(Arc::new(mock_with_responses(vec![
22254 "Intake",
22255 "First draft",
22256 "Review",
22257 "Revised draft",
22258 ])))
22259 .build()
22260 .unwrap();
22261
22262 agent.chat("First turn").await.unwrap();
22263 assert_eq!(agent.current_state().as_deref(), Some("intake"));
22264 assert!(!agent.get_context().contains_key("draft_version"));
22265
22266 agent.chat("Second turn").await.unwrap();
22267 assert_eq!(agent.current_state().as_deref(), Some("drafting"));
22268 assert_eq!(
22269 agent.get_context().get("draft_version"),
22270 Some(&serde_json::json!(1))
22271 );
22272
22273 agent.chat("Third turn").await.unwrap();
22274 assert_eq!(agent.current_state().as_deref(), Some("review"));
22275 assert_eq!(
22276 agent.get_context().get("review_entry"),
22277 Some(&serde_json::json!("first"))
22278 );
22279
22280 agent.chat("Fourth turn").await.unwrap();
22281 assert_eq!(agent.current_state().as_deref(), Some("drafting"));
22282 assert_eq!(
22283 agent.get_context().get("draft_version"),
22284 Some(&serde_json::json!(2))
22285 );
22286 }
22287
22288 #[tokio::test]
22290 async fn test_integration_process_normalize() {
22291 let yaml = r#"
22292name: ProcessAgent
22293system_prompt: "You are helpful."
22294process:
22295 input:
22296 - type: normalize
22297 config:
22298 trim: true
22299 collapse_whitespace: true
22300"#;
22301 let mock = mock_with_response("Got your message.");
22302 let builder = AgentBuilder::from_yaml(yaml).unwrap();
22303 let agent = builder.llm(Arc::new(mock.clone())).build().unwrap();
22304
22305 let _ = agent.chat(" hello world ").await.unwrap();
22306
22307 let history = mock.call_history();
22309 assert!(!history.is_empty());
22310 let last_call = history.last().unwrap();
22312 let user_msg = last_call
22313 .messages
22314 .iter()
22315 .find(|m| m.role == ai_agents_core::Role::User)
22316 .unwrap();
22317 assert_eq!(user_msg.content, "hello world");
22318 }
22319
22320 #[tokio::test]
22324 async fn test_integration_memory_compression() {
22325 let yaml = r#"
22326name: MemoryAgent
22327system_prompt: "You are helpful."
22328memory:
22329 type: compacting
22330 max_messages: 100
22331 compress_threshold: 5
22332 max_recent_messages: 3
22333 summarize_batch_size: 2
22334"#;
22335 let responses: Vec<&str> = (0..8).map(|_| "Response from assistant.").collect();
22337 let mock = mock_with_responses(responses);
22338 let builder = AgentBuilder::from_yaml(yaml).unwrap();
22339 let agent = builder.llm(Arc::new(mock)).build().unwrap();
22340
22341 for i in 0..6 {
22343 let _ = agent.chat(&format!("Message {}", i)).await.unwrap();
22344 }
22345
22346 let messages = agent.memory.get_messages(None).await.unwrap();
22349 assert!(messages.len() <= 12); }
22353
22354 #[tokio::test]
22356 async fn test_integration_multi_llm_registry() {
22357 let mut mock_default = MockLLMProvider::new("default");
22358 mock_default.set_response("Default LLM response.");
22359 let mut mock_router = MockLLMProvider::new("router");
22360 mock_router.set_response("Router response.");
22361
22362 let agent = AgentBuilder::new()
22363 .system_prompt("You are helpful.")
22364 .llm_alias("default", Arc::new(mock_default))
22365 .llm_alias("router", Arc::new(mock_router))
22366 .build()
22367 .unwrap();
22368
22369 let response = agent.chat("Hello").await.unwrap();
22370 assert_eq!(response.content, "Default LLM response.");
22371 }
22372
22373 #[tokio::test]
22375 async fn test_integration_agent_reset() {
22376 let mock = mock_with_responses(vec!["Hello!", "Hello again!"]);
22377 let agent = AgentBuilder::new()
22378 .system_prompt("You are helpful.")
22379 .llm(Arc::new(mock))
22380 .build()
22381 .unwrap();
22382
22383 let _ = agent.chat("Hi").await.unwrap();
22384 let messages = agent.memory.get_messages(None).await.unwrap();
22385 assert_eq!(messages.len(), 2); agent.reset().await.unwrap();
22388 let messages = agent.memory.get_messages(None).await.unwrap();
22389 assert_eq!(messages.len(), 0);
22390 }
22391
22392 #[tokio::test]
22394 async fn test_integration_process_validate_reject() {
22395 use ai_agents_process::{ProcessConfig, ProcessProcessor};
22396
22397 let validate_config = ai_agents_process::ValidateStage {
22398 id: Some("length_check".to_string()),
22399 condition: None,
22400 config: ai_agents_process::ValidateConfig {
22401 rules: vec![ai_agents_process::ValidationRule::MinLength {
22402 min_length: 10,
22403 on_fail: ai_agents_process::ValidationAction {
22404 action: ai_agents_process::ValidationActionType::Reject,
22405 message: None,
22406 },
22407 }],
22408 ..Default::default()
22409 },
22410 };
22411 let process_config = ProcessConfig {
22412 input: vec![ai_agents_process::ProcessStage::Validate(validate_config)],
22413 ..Default::default()
22414 };
22415 let processor = ProcessProcessor::new(process_config);
22416
22417 let mock = mock_with_response("Should not reach here.");
22418 let agent = AgentBuilder::new()
22419 .system_prompt("You are helpful.")
22420 .llm(Arc::new(mock))
22421 .process_processor(processor)
22422 .build()
22423 .unwrap();
22424
22425 let response = agent.chat("Hi").await.unwrap();
22426 assert!(
22428 response.content.contains("rejected")
22429 || response.content.contains("Input rejected")
22430 || response.content.contains("too short")
22431 || response.content.contains("Too short")
22432 || response.content.len() < 50, "Expected rejection response, got: {}",
22434 response.content
22435 );
22436 }
22437
22438 #[tokio::test]
22440 async fn test_llm_fallback_on_failure() {
22441 use ai_agents_recovery::{ErrorRecoveryConfig, LLMFailureAction, LLMRecoveryConfig};
22442
22443 let mut primary = MockLLMProvider::new("primary");
22444 primary.set_error("Primary LLM is unavailable");
22445
22446 let mut fallback = MockLLMProvider::new("fallback");
22447 fallback.set_response("Fallback response works!");
22448
22449 let agent = AgentBuilder::new()
22450 .system_prompt("You are helpful.")
22451 .llm_alias("default", Arc::new(primary))
22452 .llm_alias("backup", Arc::new(fallback))
22453 .recovery_manager(RecoveryManager::new(ErrorRecoveryConfig {
22454 llm: LLMRecoveryConfig {
22455 on_failure: LLMFailureAction::FallbackLlm {
22456 fallback_llm: "backup".to_string(),
22457 },
22458 ..Default::default()
22459 },
22460 ..Default::default()
22461 }))
22462 .build()
22463 .unwrap();
22464
22465 let response = agent.chat("Hello").await.unwrap();
22466 assert!(
22467 response.content.contains("Fallback response"),
22468 "Expected fallback response, got: {}",
22469 response.content
22470 );
22471 }
22472
22473 #[tokio::test]
22475 async fn test_llm_fallback_response_static_message() {
22476 use ai_agents_recovery::{ErrorRecoveryConfig, LLMFailureAction, LLMRecoveryConfig};
22477
22478 let mut primary = MockLLMProvider::new("primary");
22479 primary.set_error("Primary LLM is unavailable");
22480
22481 let agent = AgentBuilder::new()
22482 .system_prompt("You are helpful.")
22483 .llm(Arc::new(primary))
22484 .recovery_manager(RecoveryManager::new(ErrorRecoveryConfig {
22485 llm: LLMRecoveryConfig {
22486 on_failure: LLMFailureAction::FallbackResponse {
22487 message: "I am temporarily unavailable. Please try again later."
22488 .to_string(),
22489 },
22490 ..Default::default()
22491 },
22492 ..Default::default()
22493 }))
22494 .build()
22495 .unwrap();
22496
22497 let response = agent.chat("Hello").await.unwrap();
22498 assert!(
22499 response.content.contains("temporarily unavailable"),
22500 "Expected static fallback message, got: {}",
22501 response.content
22502 );
22503 }
22504
22505 #[tokio::test]
22508 async fn test_tool_failure_skip() {
22509 use ai_agents_recovery::{
22510 ErrorRecoveryConfig, ToolFailureAction, ToolRecoveryConfig, ToolRetryConfig,
22511 };
22512
22513 let mock = mock_with_responses(vec![
22514 r#"{"tool": "calculator", "arguments": {"expression": "not a number +"}}"#,
22515 "The calculation was skipped, but I can still help you.",
22516 ]);
22517 let observed = mock.clone();
22518 let mut tools = ai_agents_tools::ToolRegistry::new();
22519 tools
22520 .register(Arc::new(ai_agents_tools::CalculatorTool))
22521 .unwrap();
22522
22523 let agent = AgentBuilder::new()
22524 .system_prompt("You are helpful.")
22525 .llm(Arc::new(mock))
22526 .tools(tools)
22527 .recovery_manager(RecoveryManager::new(ErrorRecoveryConfig {
22528 tools: ToolRecoveryConfig {
22529 default: ToolRetryConfig {
22530 max_retries: 0,
22531 timeout_ms: None,
22532 on_failure: ToolFailureAction::Skip,
22533 },
22534 ..Default::default()
22535 },
22536 ..Default::default()
22537 }))
22538 .build()
22539 .unwrap();
22540
22541 let response = agent.chat("Compute this").await.unwrap();
22542
22543 assert_eq!(
22544 response.content,
22545 "The calculation was skipped, but I can still help you."
22546 );
22547 assert_eq!(observed.call_count(), 2);
22548 let history = agent.tool_call_history();
22550 assert_eq!(history.len(), 1);
22551 assert_eq!(history[0].tool_id, "calculator");
22552 assert_eq!(
22553 history[0].result.get("skipped"),
22554 Some(&serde_json::json!(true)),
22555 "{:?}",
22556 history[0].result
22557 );
22558 }
22559
22560 #[tokio::test]
22562 async fn test_unregistered_tool_call_records_unavailable_and_continues() {
22563 let mock = mock_with_responses(vec![
22564 r#"{"tool": "nonexistent_tool", "arguments": {}}"#,
22565 "The tool was unavailable, but I can still help you.",
22566 ]);
22567 let observed = mock.clone();
22568
22569 let agent = AgentBuilder::new()
22570 .system_prompt("You are helpful.")
22571 .llm(Arc::new(mock))
22572 .build()
22573 .unwrap();
22574
22575 let response = agent.chat("Use the nonexistent tool").await.unwrap();
22576
22577 assert_eq!(
22578 response.content,
22579 "The tool was unavailable, but I can still help you."
22580 );
22581 assert_eq!(observed.call_count(), 2);
22582 let history = agent.tool_call_history();
22583 assert_eq!(history.len(), 1);
22584 assert_eq!(history[0].tool_id, "nonexistent_tool");
22585 assert_eq!(
22586 history[0].result.pointer("/error/kind"),
22587 Some(&serde_json::json!("tool_unavailable")),
22588 "{:?}",
22589 history[0].result
22590 );
22591 }
22592
22593 fn fallback_llm_recovery(fallback_llm: &str) -> RecoveryManager {
22598 use ai_agents_recovery::{ErrorRecoveryConfig, LLMFailureAction, LLMRecoveryConfig};
22599 RecoveryManager::new(ErrorRecoveryConfig {
22600 llm: LLMRecoveryConfig {
22601 on_failure: LLMFailureAction::FallbackLlm {
22602 fallback_llm: fallback_llm.to_string(),
22603 },
22604 ..Default::default()
22605 },
22606 ..Default::default()
22607 })
22608 }
22609
22610 #[tokio::test]
22611 async fn test_stream_llm_fallback_on_open_failure() {
22612 let mut primary = MockLLMProvider::new("primary");
22613 primary.set_error("Primary LLM is unavailable");
22614 let mut fallback = MockLLMProvider::new("fallback");
22615 fallback.set_response("Fallback response works!");
22616
22617 let agent = AgentBuilder::new()
22618 .system_prompt("You are helpful.")
22619 .llm_alias("default", Arc::new(primary))
22620 .llm_alias("backup", Arc::new(fallback))
22621 .recovery_manager(fallback_llm_recovery("backup"))
22622 .build()
22623 .unwrap();
22624
22625 let (content, chunks, final_response) = collect_stream_events(&agent, "Hello").await;
22626 assert!(
22627 !chunks.iter().any(StreamChunk::is_error),
22628 "fallback must not surface as a stream error: {chunks:?}"
22629 );
22630 let final_response = final_response.expect("Final must be emitted after fallback");
22631 assert!(content.contains("Fallback response"));
22632 assert!(final_response.content.contains("Fallback response"));
22633 }
22634
22635 #[tokio::test]
22636 async fn test_stream_llm_fallback_response_static_message() {
22637 use ai_agents_recovery::{ErrorRecoveryConfig, LLMFailureAction, LLMRecoveryConfig};
22638
22639 let mut primary = MockLLMProvider::new("primary");
22640 primary.set_error("Primary LLM is unavailable");
22641
22642 let agent = AgentBuilder::new()
22643 .system_prompt("You are helpful.")
22644 .llm(Arc::new(primary))
22645 .recovery_manager(RecoveryManager::new(ErrorRecoveryConfig {
22646 llm: LLMRecoveryConfig {
22647 on_failure: LLMFailureAction::FallbackResponse {
22648 message: "Service is temporarily unavailable.".to_string(),
22649 },
22650 ..Default::default()
22651 },
22652 ..Default::default()
22653 }))
22654 .build()
22655 .unwrap();
22656
22657 let (content, chunks, final_response) = collect_stream_events(&agent, "Hello").await;
22658 assert!(!chunks.iter().any(StreamChunk::is_error));
22659 let content_chunks = chunks.iter().filter(|c| c.is_content()).count();
22660 assert_eq!(content_chunks, 1, "static fallback is one content chunk");
22661 assert_eq!(content, "Service is temporarily unavailable.");
22662 assert_eq!(
22663 final_response.expect("Final").content,
22664 "Service is temporarily unavailable."
22665 );
22666 }
22667
22668 struct FailOnceStreamProvider {
22670 remaining_failures: Arc<std::sync::atomic::AtomicUsize>,
22671 open_attempts: Arc<std::sync::atomic::AtomicUsize>,
22672 }
22673
22674 #[async_trait]
22675 impl LLMProvider for FailOnceStreamProvider {
22676 async fn complete(
22677 &self,
22678 _messages: &[ChatMessage],
22679 _config: Option<&LLMConfig>,
22680 ) -> std::result::Result<LLMResponse, LLMError> {
22681 Ok(LLMResponse::new("blocking path", FinishReason::Stop))
22682 }
22683
22684 async fn complete_stream(
22685 &self,
22686 _messages: &[ChatMessage],
22687 _config: Option<&LLMConfig>,
22688 ) -> std::result::Result<
22689 Box<dyn Stream<Item = std::result::Result<LLMChunk, LLMError>> + Unpin + Send>,
22690 LLMError,
22691 > {
22692 self.open_attempts.fetch_add(1, Ordering::SeqCst);
22693 if self
22694 .remaining_failures
22695 .fetch_update(Ordering::SeqCst, Ordering::SeqCst, |n| n.checked_sub(1))
22696 .is_ok()
22697 {
22698 return Err(LLMError::Network("connection reset".to_string()));
22699 }
22700 Ok(Box::new(futures::stream::iter(vec![Ok(LLMChunk::new(
22701 "Recovered after retry",
22702 true,
22703 ))])))
22704 }
22705
22706 fn provider_name(&self) -> &str {
22707 "fail-once-stream"
22708 }
22709
22710 fn supports(&self, feature: LLMFeature) -> bool {
22711 matches!(feature, LLMFeature::Streaming)
22712 }
22713 }
22714
22715 #[tokio::test]
22716 async fn test_stream_llm_retry_then_success() {
22717 use ai_agents_recovery::{BackoffConfig, ErrorRecoveryConfig, RetryConfig};
22718
22719 let open_attempts = Arc::new(std::sync::atomic::AtomicUsize::new(0));
22720 let provider = FailOnceStreamProvider {
22721 remaining_failures: Arc::new(std::sync::atomic::AtomicUsize::new(1)),
22722 open_attempts: Arc::clone(&open_attempts),
22723 };
22724
22725 let agent = AgentBuilder::new()
22726 .system_prompt("You are helpful.")
22727 .llm(Arc::new(provider))
22728 .recovery_manager(RecoveryManager::new(ErrorRecoveryConfig {
22729 default: RetryConfig {
22730 max_retries: 1,
22731 backoff: BackoffConfig {
22732 initial_ms: 1,
22733 max_ms: 1,
22734 ..Default::default()
22735 },
22736 ..Default::default()
22737 },
22738 ..Default::default()
22739 }))
22740 .build()
22741 .unwrap();
22742
22743 let (content, chunks, final_response) = collect_stream_events(&agent, "Hello").await;
22744 assert!(!chunks.iter().any(StreamChunk::is_error), "{chunks:?}");
22745 assert_eq!(open_attempts.load(Ordering::SeqCst), 2);
22746 assert_eq!(content, "Recovered after retry");
22747 assert_eq!(
22748 final_response.expect("Final").content,
22749 "Recovered after retry"
22750 );
22751 }
22752
22753 #[tokio::test]
22754 async fn test_stream_llm_error_action_error_emits_terminal_error() {
22755 let mut primary = MockLLMProvider::new("primary");
22756 primary.set_error("Primary LLM is unavailable");
22757
22758 let agent = AgentBuilder::new()
22759 .system_prompt("You are helpful.")
22760 .llm(Arc::new(primary))
22761 .build()
22762 .unwrap();
22763
22764 let (_, chunks, final_response) = collect_stream_events(&agent, "Hello").await;
22765 assert!(
22766 final_response.is_none(),
22767 "default Error action must not produce Final"
22768 );
22769 assert!(
22770 chunks.iter().any(StreamChunk::is_error),
22771 "default Error action must surface a stream error"
22772 );
22773 }
22774
22775 struct MidStreamFailureProvider;
22777
22778 #[async_trait]
22779 impl LLMProvider for MidStreamFailureProvider {
22780 async fn complete(
22781 &self,
22782 _messages: &[ChatMessage],
22783 _config: Option<&LLMConfig>,
22784 ) -> std::result::Result<LLMResponse, LLMError> {
22785 Ok(LLMResponse::new("blocking path", FinishReason::Stop))
22786 }
22787
22788 async fn complete_stream(
22789 &self,
22790 _messages: &[ChatMessage],
22791 _config: Option<&LLMConfig>,
22792 ) -> std::result::Result<
22793 Box<dyn Stream<Item = std::result::Result<LLMChunk, LLMError>> + Unpin + Send>,
22794 LLMError,
22795 > {
22796 Ok(Box::new(futures::stream::iter(vec![
22797 Ok(LLMChunk::new("Partial ", false)),
22798 Err(LLMError::Network("connection dropped".to_string())),
22799 ])))
22800 }
22801
22802 fn provider_name(&self) -> &str {
22803 "mid-stream-failure"
22804 }
22805
22806 fn supports(&self, feature: LLMFeature) -> bool {
22807 matches!(feature, LLMFeature::Streaming)
22808 }
22809 }
22810
22811 #[tokio::test]
22812 async fn test_stream_mid_stream_failure_is_terminal() {
22813 let mut fallback = MockLLMProvider::new("fallback");
22814 fallback.set_response("Fallback must not run");
22815 let fallback_calls = fallback.clone();
22816
22817 let agent = AgentBuilder::new()
22818 .system_prompt("You are helpful.")
22819 .llm_alias("default", Arc::new(MidStreamFailureProvider))
22820 .llm_alias("backup", Arc::new(fallback))
22821 .recovery_manager(fallback_llm_recovery("backup"))
22822 .build()
22823 .unwrap();
22824
22825 let (content, chunks, final_response) = collect_stream_events(&agent, "Hello").await;
22826 assert_eq!(content, "Partial ");
22827 assert!(chunks.iter().any(StreamChunk::is_error));
22828 assert!(final_response.is_none());
22829 assert_eq!(
22830 fallback_calls.call_count(),
22831 0,
22832 "fallback must not run after a visible delta"
22833 );
22834 }
22835
22836 #[tokio::test]
22837 async fn test_buffered_streaming_draft_uses_fallback_llm() {
22838 use futures::StreamExt;
22839
22840 let mut primary = MockLLMProvider::new("primary");
22841 primary.set_error("Primary LLM is unavailable");
22842 let fallback = mock_with_response("fallback one two");
22843 let yaml = r#"
22844name: BufferedFallbackAgent
22845system_prompt: "You stream safely."
22846llm:
22847 default: default
22848streaming:
22849 enabled: true
22850 buffer_size: 8
22851runtime:
22852 optimization:
22853 enabled: true
22854 max_speculative_llm_calls_per_turn: 2
22855 speculative_state_transitions: true
22856 streaming_policy: buffer_until_routing_done
22857 max_parallel_runtime_tasks: 2
22858states:
22859 initial: triage
22860 states:
22861 triage:
22862 prompt: "Answer from triage."
22863 transitions:
22864 - to: billing
22865 guard:
22866 context:
22867 route:
22868 eq: billing
22869 timing: parallel
22870 billing:
22871 prompt: "Billing state."
22872"#;
22873 let agent = AgentBuilder::from_yaml(yaml)
22874 .unwrap()
22875 .llm_alias("default", Arc::new(primary))
22876 .llm_alias("backup", Arc::new(fallback))
22877 .recovery_manager(fallback_llm_recovery("backup"))
22878 .build()
22879 .unwrap();
22880
22881 let mut stream = agent.chat_stream("hello").await.unwrap();
22882 let mut content = String::new();
22883 let mut error = None;
22884 while let Some(chunk) = stream.next().await {
22885 match chunk {
22886 StreamChunk::Content { text } => content.push_str(&text),
22887 StreamChunk::Error { message } => error = Some(message),
22888 StreamChunk::Done {} => break,
22889 _ => {}
22890 }
22891 }
22892
22893 assert_eq!(error, None);
22894 assert_eq!(content, "fallback one two");
22895 }
22896
22897 #[tokio::test]
22898 async fn parity_llm_fallback_llm() {
22899 let build = || {
22900 let mut primary = MockLLMProvider::new("primary");
22901 primary.set_error("Primary LLM is unavailable");
22902 let mut fallback = MockLLMProvider::new("fallback");
22903 fallback.set_response("Fallback response works!");
22904 AgentBuilder::new()
22905 .system_prompt("You are helpful.")
22906 .llm_alias("default", Arc::new(primary))
22907 .llm_alias("backup", Arc::new(fallback))
22908 .recovery_manager(fallback_llm_recovery("backup"))
22909 .build()
22910 .unwrap()
22911 };
22912 assert_blocking_streaming_parity(build, "Hello").await;
22913 }
22914
22915 #[tokio::test]
22916 async fn parity_llm_fallback_response() {
22917 use ai_agents_recovery::{ErrorRecoveryConfig, LLMFailureAction, LLMRecoveryConfig};
22918 let build = || {
22919 let mut primary = MockLLMProvider::new("primary");
22920 primary.set_error("Primary LLM is unavailable");
22921 AgentBuilder::new()
22922 .system_prompt("You are helpful.")
22923 .llm(Arc::new(primary))
22924 .recovery_manager(RecoveryManager::new(ErrorRecoveryConfig {
22925 llm: LLMRecoveryConfig {
22926 on_failure: LLMFailureAction::FallbackResponse {
22927 message: "Service is temporarily unavailable.".to_string(),
22928 },
22929 ..Default::default()
22930 },
22931 ..Default::default()
22932 }))
22933 .build()
22934 .unwrap()
22935 };
22936 assert_blocking_streaming_parity(build, "Hello").await;
22937 }
22938
22939 #[tokio::test]
22940 async fn parity_basic_chat() {
22941 let build = || {
22942 AgentBuilder::new()
22943 .system_prompt("You are helpful.")
22944 .llm(Arc::new(mock_with_response("Plain answer")))
22945 .build()
22946 .unwrap()
22947 };
22948 assert_blocking_streaming_parity(build, "Hello").await;
22949 }
22950
22951 fn skills_with_parallel_transition_yaml(extra_optimization: &str, streaming: &str) -> String {
22957 format!(
22958 r#"
22959name: SkillsBesideTransitionAgent
22960system_prompt: "Use skills when they match."
22961llm:
22962 default: default
22963 router: router
22964observability:
22965 enabled: true
22966 export:
22967 write_raw_events: true
22968{streaming}
22969runtime:
22970 optimization:
22971 enabled: true
22972 speculative_state_transitions: true
22973{extra_optimization}
22974states:
22975 initial: triage
22976 states:
22977 triage:
22978 prompt: "Triage state."
22979 transitions:
22980 - to: billing
22981 guard:
22982 context:
22983 route:
22984 eq: billing
22985 timing: parallel
22986 billing:
22987 prompt: "Billing state."
22988skills:
22989 - id: helper
22990 description: "Answer helper requests"
22991 trigger: "User asks for helper"
22992 steps:
22993 - prompt: "Answer the helper request: {{{{ user_input }}}}"
22994 llm: skill
22995"#
22996 )
22997 }
22998
22999 struct RoleMocks {
23003 main: MockLLMProvider,
23004 router: MockLLMProvider,
23005 skill: MockLLMProvider,
23006 }
23007
23008 fn role_mocks(main: MockLLMProvider, router: MockLLMProvider) -> RoleMocks {
23009 RoleMocks {
23010 main,
23011 router,
23012 skill: mock_with_response("Skill step response"),
23013 }
23014 }
23015
23016 fn build_skills_beside_transition_agent(yaml: &str, mocks: RoleMocks) -> RuntimeAgent {
23017 AgentBuilder::from_yaml(yaml)
23018 .unwrap()
23019 .llm_alias("default", Arc::new(mocks.main))
23020 .llm_alias("router", Arc::new(mocks.router))
23021 .llm_alias("skill", Arc::new(mocks.skill))
23022 .build()
23023 .unwrap()
23024 }
23025
23026 fn branch_events_with_commit_behavior(agent: &RuntimeAgent, behavior: &str) -> usize {
23027 agent
23028 .observability()
23029 .unwrap()
23030 .raw_events()
23031 .iter()
23032 .filter(|event| event.dimensions.get("commit_behavior") == Some(&behavior.to_string()))
23033 .count()
23034 }
23035
23036 #[tokio::test]
23037 async fn test_speculative_transition_with_skills_and_no_skill_branch_routes_skill_serially() {
23038 let default_mock = mock_with_response("Draft response");
23039 let router_mock = mock_with_response("helper");
23040 let router_counter = router_mock.clone();
23041 let yaml = skills_with_parallel_transition_yaml(
23042 " max_speculative_llm_calls_per_turn: 2\n max_parallel_runtime_tasks: 2",
23043 "",
23044 );
23045 let agent =
23046 build_skills_beside_transition_agent(&yaml, role_mocks(default_mock, router_mock));
23047
23048 let response = agent.chat("please use helper").await.unwrap();
23049
23050 assert_eq!(
23051 response.metadata.as_ref().and_then(|m| m.get("skill_id")),
23052 Some(&serde_json::json!("helper")),
23053 "skill must route even without a skill branch: {response:?}"
23054 );
23055 assert_eq!(router_counter.call_count(), 1);
23056 assert!(branch_events_with_commit_behavior(&agent, "transition_decision") > 0);
23058 assert_eq!(
23059 branch_events_with_commit_behavior(&agent, "skill_selection"),
23060 0
23061 );
23062 }
23063
23064 #[tokio::test]
23065 async fn test_speculative_transition_with_skills_no_match_commits_draft() {
23066 let default_mock = mock_with_response("Draft response");
23067 let router_mock = mock_with_response("none");
23068 let router_counter = router_mock.clone();
23069 let yaml = skills_with_parallel_transition_yaml(
23070 " max_speculative_llm_calls_per_turn: 2\n max_parallel_runtime_tasks: 2",
23071 "",
23072 );
23073 let agent =
23074 build_skills_beside_transition_agent(&yaml, role_mocks(default_mock, router_mock));
23075
23076 let response = agent.chat("just chat").await.unwrap();
23077
23078 assert_eq!(response.content, "Draft response");
23079 assert!(
23080 response
23081 .metadata
23082 .as_ref()
23083 .is_none_or(|m| !m.contains_key("skill_id"))
23084 );
23085 assert_eq!(router_counter.call_count(), 1);
23086 assert!(branch_events_with_commit_behavior(&agent, "final_response") > 0);
23087 }
23088
23089 #[tokio::test]
23090 async fn test_speculative_transition_win_skips_serial_skill_selection() {
23091 let default_mock = mock_with_response("Billing answer");
23092 let router_mock = mock_with_response("none");
23093 let router_counter = router_mock.clone();
23094 let yaml = skills_with_parallel_transition_yaml(
23095 " max_speculative_llm_calls_per_turn: 2\n max_parallel_runtime_tasks: 2",
23096 "",
23097 );
23098 let agent =
23099 build_skills_beside_transition_agent(&yaml, role_mocks(default_mock, router_mock));
23100 agent
23101 .set_context("route", serde_json::json!("billing"))
23102 .unwrap();
23103
23104 let response = agent.chat("billing please").await.unwrap();
23105
23106 assert_eq!(agent.current_state().as_deref(), Some("billing"));
23107 assert_eq!(response.content, "Billing answer");
23108 assert_eq!(router_counter.call_count(), 1);
23110 }
23111
23112 #[tokio::test]
23113 async fn test_speculative_skill_capacity_exhausted_still_routes_skill_serially() {
23114 let default_mock = mock_with_response("Draft response");
23115 let router_mock = mock_with_response("helper");
23116 let router_counter = router_mock.clone();
23117 let yaml = skills_with_parallel_transition_yaml(
23119 " speculative_skill_routing: true\n max_speculative_llm_calls_per_turn: 2\n max_parallel_runtime_tasks: 2",
23120 "",
23121 );
23122 let agent =
23123 build_skills_beside_transition_agent(&yaml, role_mocks(default_mock, router_mock));
23124
23125 let response = agent.chat("please use helper").await.unwrap();
23126
23127 assert_eq!(
23128 response.metadata.as_ref().and_then(|m| m.get("skill_id")),
23129 Some(&serde_json::json!("helper"))
23130 );
23131 assert_eq!(router_counter.call_count(), 1);
23132 assert_eq!(
23133 branch_events_with_commit_behavior(&agent, "skill_selection"),
23134 0
23135 );
23136 }
23137
23138 #[tokio::test]
23139 async fn test_speculative_transition_and_skill_both_enabled_unchanged() {
23140 let default_mock = mock_with_response("Draft response");
23141 let router_mock = mock_with_response("helper");
23142 let router_counter = router_mock.clone();
23143 let yaml = skills_with_parallel_transition_yaml(
23144 " speculative_skill_routing: true\n max_speculative_llm_calls_per_turn: 3\n max_parallel_runtime_tasks: 3",
23145 "",
23146 );
23147 let agent =
23148 build_skills_beside_transition_agent(&yaml, role_mocks(default_mock, router_mock));
23149
23150 let response = agent.chat("please use helper").await.unwrap();
23151
23152 assert_eq!(
23153 response.metadata.as_ref().and_then(|m| m.get("skill_id")),
23154 Some(&serde_json::json!("helper"))
23155 );
23156 assert_eq!(router_counter.call_count(), 1);
23157 assert!(branch_events_with_commit_behavior(&agent, "skill_selection") > 0);
23159 }
23160
23161 const BUFFERED_STREAMING_YAML_FRAGMENT: &str = "streaming:\n enabled: true\n buffer_size: 16";
23162 const BUFFERED_OPTIMIZATION_FRAGMENT: &str = " max_speculative_llm_calls_per_turn: 2\n streaming_policy: buffer_until_routing_done\n max_parallel_runtime_tasks: 2";
23163
23164 #[tokio::test]
23165 async fn test_buffered_streaming_skill_wins_after_transition_miss() {
23166 let mut default_mock = mock_with_response("draft one two");
23167 default_mock.set_latency(10);
23168 let router_mock = mock_with_response("helper");
23169 let router_counter = router_mock.clone();
23170 let yaml = skills_with_parallel_transition_yaml(
23171 BUFFERED_OPTIMIZATION_FRAGMENT,
23172 BUFFERED_STREAMING_YAML_FRAGMENT,
23173 );
23174 let agent =
23175 build_skills_beside_transition_agent(&yaml, role_mocks(default_mock, router_mock));
23176
23177 let (content, chunks, final_response) =
23178 collect_stream_events(&agent, "please use helper").await;
23179
23180 assert!(!chunks.iter().any(StreamChunk::is_error), "{chunks:?}");
23181 assert!(
23182 !content.contains("draft"),
23183 "buffered draft must be discarded when a skill wins: {content:?}"
23184 );
23185 let final_response = final_response.expect("Final");
23186 assert_eq!(
23187 final_response
23188 .metadata
23189 .as_ref()
23190 .and_then(|m| m.get("skill_id")),
23191 Some(&serde_json::json!("helper"))
23192 );
23193 assert_eq!(content, final_response.content);
23194 assert_eq!(router_counter.call_count(), 1);
23195 }
23196
23197 #[tokio::test]
23198 async fn test_buffered_streaming_skill_miss_releases_buffer_and_commits_draft() {
23199 let mut default_mock = mock_with_response("draft one two");
23200 default_mock.set_latency(10);
23201 let router_mock = mock_with_response("none");
23202 let router_counter = router_mock.clone();
23203 let yaml = skills_with_parallel_transition_yaml(
23204 BUFFERED_OPTIMIZATION_FRAGMENT,
23205 BUFFERED_STREAMING_YAML_FRAGMENT,
23206 );
23207 let agent =
23208 build_skills_beside_transition_agent(&yaml, role_mocks(default_mock, router_mock));
23209
23210 let (content, chunks, final_response) = collect_stream_events(&agent, "just chat").await;
23211
23212 assert!(!chunks.iter().any(StreamChunk::is_error), "{chunks:?}");
23213 assert_eq!(content, "draft one two");
23214 assert_eq!(final_response.expect("Final").content, "draft one two");
23215 assert_eq!(router_counter.call_count(), 1);
23216 }
23217
23218 #[tokio::test]
23219 async fn test_buffered_streaming_transition_win_skips_skill_selection() {
23220 let default_mock = mock_with_response("Billing answer");
23221 let router_mock = mock_with_response("none");
23222 let router_counter = router_mock.clone();
23223 let yaml = skills_with_parallel_transition_yaml(
23224 BUFFERED_OPTIMIZATION_FRAGMENT,
23225 BUFFERED_STREAMING_YAML_FRAGMENT,
23226 );
23227 let agent =
23228 build_skills_beside_transition_agent(&yaml, role_mocks(default_mock, router_mock));
23229 agent
23230 .set_context("route", serde_json::json!("billing"))
23231 .unwrap();
23232
23233 let (content, chunks, final_response) =
23234 collect_stream_events(&agent, "billing please").await;
23235
23236 assert!(!chunks.iter().any(StreamChunk::is_error), "{chunks:?}");
23237 assert_eq!(agent.current_state().as_deref(), Some("billing"));
23238 assert_eq!(content, "Billing answer");
23239 assert_eq!(final_response.expect("Final").content, "Billing answer");
23240 assert_eq!(router_counter.call_count(), 1);
23242 }
23243
23244 #[tokio::test]
23245 async fn parity_buffered_policy_with_skills() {
23246 let yaml = skills_with_parallel_transition_yaml(
23247 BUFFERED_OPTIMIZATION_FRAGMENT,
23248 BUFFERED_STREAMING_YAML_FRAGMENT,
23249 );
23250 let build = || {
23251 build_skills_beside_transition_agent(
23252 &yaml,
23253 role_mocks(
23254 mock_with_response("draft one two"),
23255 mock_with_response("helper"),
23256 ),
23257 )
23258 };
23259 let (blocking, _, _) = assert_blocking_streaming_parity(build, "please use helper").await;
23260 assert_eq!(
23261 blocking.metadata.as_ref().and_then(|m| m.get("skill_id")),
23262 Some(&serde_json::json!("helper"))
23263 );
23264 }
23265
23266 #[tokio::test]
23267 async fn parity_buffered_policy_with_cot() {
23268 let yaml = format!(
23269 r#"
23270name: BufferedCotAgent
23271system_prompt: "Think first."
23272llm:
23273 default: default
23274streaming:
23275 enabled: true
23276 buffer_size: 16
23277reasoning:
23278 mode: cot
23279runtime:
23280 optimization:
23281 enabled: true
23282 speculative_state_transitions: true
23283{BUFFERED_OPTIMIZATION_FRAGMENT}
23284states:
23285 initial: triage
23286 states:
23287 triage:
23288 prompt: "Triage state."
23289 transitions:
23290 - to: billing
23291 guard:
23292 context:
23293 route:
23294 eq: billing
23295 timing: parallel
23296 billing:
23297 prompt: "Billing state."
23298"#
23299 );
23300 let build = || {
23301 AgentBuilder::from_yaml(&yaml)
23302 .unwrap()
23303 .llm_alias(
23304 "default",
23305 Arc::new(mock_with_response(
23306 "<thinking>step by step</thinking>Reasoned answer",
23307 )),
23308 )
23309 .build()
23310 .unwrap()
23311 };
23312 let (blocking, streamed, _) = assert_blocking_streaming_parity(build, "hello").await;
23313 assert_eq!(blocking.content, "Reasoned answer");
23314 let mode = streamed
23315 .metadata
23316 .as_ref()
23317 .and_then(|m| m.get("reasoning"))
23318 .and_then(|r| r.get("mode_used"))
23319 .cloned();
23320 assert_eq!(
23322 mode,
23323 Some(serde_json::to_value(ReasoningMode::CoT).unwrap())
23324 );
23325 }
23326
23327 fn post_response_transition_yaml(states_extra: &str, billing_extra: &str) -> String {
23334 format!(
23335 r#"
23336name: PostResponseTransitionAgent
23337system_prompt: "You are helpful."
23338streaming:
23339 enabled: true
23340states:
23341 initial: intake
23342{states_extra}
23343 states:
23344 intake:
23345 prompt: "Intake"
23346 transitions:
23347 - to: billing
23348 guard:
23349 context:
23350 route:
23351 eq: billing
23352 billing:
23353 prompt: "Billing"
23354{billing_extra}
23355"#
23356 )
23357 }
23358
23359 fn build_post_response_transition_agent(yaml: &str, mock: MockLLMProvider) -> RuntimeAgent {
23360 let agent = AgentBuilder::from_yaml(yaml)
23361 .unwrap()
23362 .llm(Arc::new(mock))
23363 .build()
23364 .unwrap();
23365 agent
23366 .set_context("route", serde_json::json!("billing"))
23367 .unwrap();
23368 agent
23369 }
23370
23371 fn count_occurrences(haystack: &str, needle: &str) -> usize {
23372 haystack.matches(needle).count()
23373 }
23374
23375 #[tokio::test]
23376 async fn test_stream_transition_without_regeneration_emits_content_once() {
23377 let yaml = post_response_transition_yaml(" regenerate_on_transition: false", "");
23378 let agent =
23379 build_post_response_transition_agent(&yaml, mock_with_response("Intake answer"));
23380
23381 let (content, chunks, final_response) = collect_stream_events(&agent, "hello").await;
23382
23383 assert!(!chunks.iter().any(StreamChunk::is_error), "{chunks:?}");
23384 assert_eq!(agent.current_state().as_deref(), Some("billing"));
23385 assert_eq!(
23386 count_occurrences(&content, "Intake answer"),
23387 1,
23388 "committed content must not be emitted twice: {content:?}"
23389 );
23390 assert_eq!(final_response.expect("Final").content, content);
23391 assert!(
23392 chunks
23393 .iter()
23394 .any(|c| matches!(c, StreamChunk::StateTransition { .. }))
23395 );
23396 }
23397
23398 #[tokio::test]
23399 async fn test_stream_transition_without_regeneration_buffered_emits_content_once() {
23400 let yaml = post_response_transition_yaml(" regenerate_on_transition: false", "");
23401 let mut mock = mock_with_response("Intake answer");
23402 mock.set_tool_choice(Some(ToolChoice::Auto));
23404 let agent = build_post_response_transition_agent(&yaml, mock);
23405
23406 let (content, chunks, final_response) = collect_stream_events(&agent, "hello").await;
23407
23408 assert!(!chunks.iter().any(StreamChunk::is_error), "{chunks:?}");
23409 assert_eq!(agent.current_state().as_deref(), Some("billing"));
23410 assert_eq!(
23411 count_occurrences(&content, "Intake answer"),
23412 1,
23413 "{content:?}"
23414 );
23415 assert_eq!(final_response.expect("Final").content, content);
23416 }
23417
23418 #[tokio::test]
23419 async fn test_stream_state_regenerate_on_enter_false_emits_content_once() {
23420 let yaml = post_response_transition_yaml("", " regenerate_on_enter: false");
23421 let agent =
23422 build_post_response_transition_agent(&yaml, mock_with_response("Intake answer"));
23423
23424 let (content, chunks, final_response) = collect_stream_events(&agent, "hello").await;
23425
23426 assert!(!chunks.iter().any(StreamChunk::is_error), "{chunks:?}");
23427 assert_eq!(agent.current_state().as_deref(), Some("billing"));
23428 assert_eq!(
23429 count_occurrences(&content, "Intake answer"),
23430 1,
23431 "{content:?}"
23432 );
23433 assert_eq!(final_response.expect("Final").content, content);
23434 }
23435
23436 #[tokio::test]
23437 async fn test_stream_transition_with_regeneration_emits_replacement() {
23438 let yaml = post_response_transition_yaml("", "");
23439 let agent = build_post_response_transition_agent(
23440 &yaml,
23441 mock_with_responses(vec!["Intake answer", "Billing answer"]),
23442 );
23443
23444 let (content, chunks, final_response) = collect_stream_events(&agent, "hello").await;
23445
23446 assert!(!chunks.iter().any(StreamChunk::is_error), "{chunks:?}");
23447 assert_eq!(agent.current_state().as_deref(), Some("billing"));
23448 assert_eq!(count_occurrences(&content, "Intake answer"), 1);
23450 assert_eq!(count_occurrences(&content, "Billing answer"), 1);
23451 assert_eq!(final_response.expect("Final").content, "Billing answer");
23452 }
23453
23454 #[tokio::test]
23455 async fn test_blocking_transition_without_regeneration_unchanged() {
23456 let yaml = post_response_transition_yaml(" regenerate_on_transition: false", "");
23457 let agent =
23458 build_post_response_transition_agent(&yaml, mock_with_response("Intake answer"));
23459
23460 let response = agent.chat("hello").await.unwrap();
23461
23462 assert_eq!(response.content, "Intake answer");
23463 assert_eq!(agent.current_state().as_deref(), Some("billing"));
23464 }
23465
23466 #[tokio::test]
23467 async fn parity_transition_regenerate_off() {
23468 let yaml = post_response_transition_yaml(" regenerate_on_transition: false", "");
23469 let build =
23470 || build_post_response_transition_agent(&yaml, mock_with_response("Intake answer"));
23471 assert_blocking_streaming_parity(build, "hello").await;
23472 }
23473
23474 #[tokio::test]
23475 async fn parity_transition_regenerate_on() {
23476 let yaml = post_response_transition_yaml("", "");
23477 let build = || {
23478 build_post_response_transition_agent(
23479 &yaml,
23480 mock_with_responses(vec!["Intake answer", "Billing answer"]),
23481 )
23482 };
23483 let (blocking, _, _) = assert_blocking_streaming_parity(build, "hello").await;
23484 assert_eq!(blocking.content, "Billing answer");
23485 }
23486
23487 fn rejecting_process_processor() -> ProcessProcessor {
23492 use ai_agents_process::ProcessConfig;
23493 let validate_config = ai_agents_process::ValidateStage {
23494 id: Some("length_check".to_string()),
23495 condition: None,
23496 config: ai_agents_process::ValidateConfig {
23497 rules: vec![ai_agents_process::ValidationRule::MinLength {
23498 min_length: 10,
23499 on_fail: ai_agents_process::ValidationAction {
23500 action: ai_agents_process::ValidationActionType::Reject,
23501 message: None,
23502 },
23503 }],
23504 ..Default::default()
23505 },
23506 };
23507 ProcessProcessor::new(ProcessConfig {
23508 input: vec![ai_agents_process::ProcessStage::Validate(validate_config)],
23509 ..Default::default()
23510 })
23511 }
23512
23513 fn looks_like_rejection(content: &str) -> bool {
23515 content.contains("rejected")
23516 || content.contains("Input rejected")
23517 || content.contains("too short")
23518 || content.contains("Too short")
23519 || content.len() < 50
23520 }
23521
23522 #[tokio::test]
23523 async fn test_stream_input_rejection_is_final_response() {
23524 let responses = Arc::new(std::sync::atomic::AtomicUsize::new(0));
23525 let hooks = Arc::new(ResponseCountingHooks {
23526 responses: Arc::clone(&responses),
23527 });
23528 let mock = mock_with_response("Should not reach here.");
23529 let llm_calls = mock.clone();
23530 let agent = AgentBuilder::new()
23531 .system_prompt("You are helpful.")
23532 .llm(Arc::new(mock))
23533 .process_processor(rejecting_process_processor())
23534 .hooks(hooks.clone())
23535 .build()
23536 .unwrap();
23537
23538 let (content, chunks, final_response) = collect_stream_events(&agent, "Hi").await;
23539
23540 assert!(
23541 !chunks.iter().any(StreamChunk::is_error),
23542 "rejection is a response, not a stream error: {chunks:?}"
23543 );
23544 let final_response = final_response.expect("rejection must finalize as Final");
23545 assert!(
23546 looks_like_rejection(&final_response.content),
23547 "Expected rejection response, got: {}",
23548 final_response.content
23549 );
23550 assert_eq!(content, final_response.content);
23551 assert_eq!(
23552 llm_calls.call_count(),
23553 0,
23554 "rejected input must not reach the LLM"
23555 );
23556 assert_eq!(responses.load(Ordering::SeqCst), 1, "on_response must fire");
23557 }
23558
23559 #[tokio::test]
23560 async fn parity_input_rejection() {
23561 let build = || {
23562 AgentBuilder::new()
23563 .system_prompt("You are helpful.")
23564 .llm(Arc::new(mock_with_response("Should not reach here.")))
23565 .process_processor(rejecting_process_processor())
23566 .build()
23567 .unwrap()
23568 };
23569 let (blocking, _, _) = assert_blocking_streaming_parity(build, "Hi").await;
23570 assert!(
23571 looks_like_rejection(&blocking.content),
23572 "{}",
23573 blocking.content
23574 );
23575 }
23576
23577 fn pre_response_transition_yaml(streaming_policy: &str) -> String {
23578 format!(
23579 r#"
23580name: StreamingPreflightAgent
23581system_prompt: "You route before streaming."
23582runtime:
23583 optimization:
23584 enabled: true
23585 pre_response_deterministic_transitions: true
23586 streaming_policy: {streaming_policy}
23587streaming:
23588 enabled: true
23589 buffer_size: 16
23590states:
23591 initial: greeting
23592 states:
23593 greeting:
23594 prompt: "OLD_STATE_SENTINEL"
23595 transitions:
23596 - to: billing
23597 guard:
23598 context:
23599 topic:
23600 eq: billing
23601 timing: pre_response
23602 billing:
23603 prompt: "Billing state."
23604"#
23605 )
23606 }
23607
23608 fn build_pre_response_transition_agent(yaml: &str) -> RuntimeAgent {
23609 let agent = AgentBuilder::from_yaml(yaml)
23610 .unwrap()
23611 .llm(Arc::new(mock_with_response("Billing streamed response")))
23612 .build()
23613 .unwrap();
23614 agent
23615 .set_context("topic", serde_json::json!("billing"))
23616 .unwrap();
23617 agent
23618 }
23619
23620 #[tokio::test]
23621 async fn test_stream_buffered_policy_runs_pre_response_deterministic_transition() {
23622 let yaml = pre_response_transition_yaml("buffer_until_routing_done");
23623 let agent = build_pre_response_transition_agent(&yaml);
23624
23625 let (content, chunks, final_response) =
23626 collect_stream_events(&agent, "billing please").await;
23627
23628 assert!(!chunks.iter().any(StreamChunk::is_error), "{chunks:?}");
23629 assert_eq!(agent.current_state().as_deref(), Some("billing"));
23630 assert!(content.contains("Billing streamed response"));
23631 assert!(!content.contains("OLD_STATE_SENTINEL"));
23632 assert_eq!(final_response.expect("Final").content, content);
23633 }
23634
23635 #[tokio::test]
23636 async fn test_stream_disabled_policy_skips_preflight() {
23637 let yaml = pre_response_transition_yaml("disabled");
23638 let agent = build_pre_response_transition_agent(&yaml);
23639
23640 let (_, chunks, final_response) = collect_stream_events(&agent, "billing please").await;
23641
23642 assert!(!chunks.iter().any(StreamChunk::is_error), "{chunks:?}");
23643 assert!(final_response.is_some());
23644 assert_eq!(agent.current_state().as_deref(), Some("greeting"));
23648 }
23649
23650 #[tokio::test]
23651 async fn parity_pre_response_transition_buffered_policy() {
23652 let yaml = pre_response_transition_yaml("buffer_until_routing_done");
23653 let build = || build_pre_response_transition_agent(&yaml);
23654 assert_blocking_streaming_parity(build, "billing please").await;
23655 }
23656
23657 fn calculator_agent_with(mock: MockLLMProvider) -> RuntimeAgent {
23662 let mut tools = ai_agents_tools::ToolRegistry::new();
23663 tools
23664 .register(Arc::new(ai_agents_tools::CalculatorTool))
23665 .unwrap();
23666 AgentBuilder::new()
23667 .system_prompt("You are a calculator assistant.")
23668 .llm(Arc::new(mock))
23669 .tools(tools)
23670 .build()
23671 .unwrap()
23672 }
23673
23674 #[tokio::test]
23675 async fn test_stream_tool_start_events_precede_results_for_batch() {
23676 let mock = mock_with_responses(vec![
23677 r#"[{"tool": "calculator", "arguments": {"expression": "1+1"}}, {"tool": "calculator", "arguments": {"expression": "2+2"}}]"#,
23678 "Both answers are ready.",
23679 ]);
23680 let agent = calculator_agent_with(mock);
23681
23682 let (_, chunks, final_response) = collect_stream_events(&agent, "compute both").await;
23683
23684 assert!(!chunks.iter().any(StreamChunk::is_error), "{chunks:?}");
23685 let final_response = final_response.expect("Final");
23686 assert_eq!(final_response.tool_calls.as_ref().map(Vec::len), Some(2));
23687
23688 let tool_events: Vec<&StreamChunk> = chunks
23689 .iter()
23690 .filter(|c| {
23691 matches!(
23692 c,
23693 StreamChunk::ToolCallStart { .. }
23694 | StreamChunk::ToolResult { .. }
23695 | StreamChunk::ToolCallEnd { .. }
23696 )
23697 })
23698 .collect();
23699 assert_eq!(tool_events.len(), 6, "{tool_events:?}");
23700 assert!(matches!(tool_events[0], StreamChunk::ToolCallStart { .. }));
23702 assert!(matches!(tool_events[1], StreamChunk::ToolCallStart { .. }));
23703 assert!(matches!(
23704 tool_events[2],
23705 StreamChunk::ToolResult { success: true, .. }
23706 ));
23707 assert!(matches!(tool_events[3], StreamChunk::ToolCallEnd { .. }));
23708 assert!(matches!(
23709 tool_events[4],
23710 StreamChunk::ToolResult { success: true, .. }
23711 ));
23712 assert!(matches!(tool_events[5], StreamChunk::ToolCallEnd { .. }));
23713 }
23714
23715 #[tokio::test]
23716 async fn test_stream_clarification_final_carries_options_and_detection() {
23717 let responses = || {
23718 vec![
23719 r#"{"is_ambiguous":true,"confidence":0.2,"ambiguity_type":"missing_target","reasoning":"target missing","what_is_unclear":["target"],"detected_language":"en"}"#,
23720 r#"{"question":"What should I send?","options":["report","invoice"]}"#,
23721 ]
23722 };
23723 let (blocking_agent, _) = state_disambiguation_agent(responses(), true, None, true);
23724 let (streaming_agent, _) = state_disambiguation_agent(responses(), true, None, true);
23725
23726 let blocking = blocking_agent.chat("Send it").await.unwrap();
23727 let (_, chunks, streamed) = collect_stream_events(&streaming_agent, "Send it").await;
23728 assert!(!chunks.iter().any(StreamChunk::is_error), "{chunks:?}");
23729 let streamed = streamed.expect("clarification must finalize as Final");
23730
23731 assert_eq!(streamed.content, "What should I send?");
23732 let streamed_meta = streamed
23733 .metadata
23734 .as_ref()
23735 .and_then(|m| m.get("disambiguation"))
23736 .cloned()
23737 .expect("disambiguation metadata");
23738 for key in ["status", "options", "clarifying", "detection"] {
23739 assert!(
23740 streamed_meta.get(key).is_some(),
23741 "missing {key}: {streamed_meta}"
23742 );
23743 }
23744 assert_eq!(
23745 streamed_meta.get("detection").and_then(|d| d.get("type")),
23746 Some(&serde_json::json!("missing_target"))
23747 );
23748 assert_eq!(
23749 blocking
23750 .metadata
23751 .as_ref()
23752 .and_then(|m| m.get("disambiguation")),
23753 Some(&streamed_meta),
23754 "blocking and streaming clarification metadata must be identical"
23755 );
23756 }
23757
23758 struct FailingMemory {
23760 messages: parking_lot::RwLock<Vec<ChatMessage>>,
23761 fail_on_add: usize,
23762 adds: std::sync::atomic::AtomicUsize,
23763 }
23764
23765 #[async_trait]
23766 impl ai_agents_core::Memory for FailingMemory {
23767 async fn add_message(&self, message: ChatMessage) -> Result<()> {
23768 let n = self.adds.fetch_add(1, Ordering::SeqCst) + 1;
23769 if n == self.fail_on_add {
23770 return Err(AgentError::Other(format!(
23771 "simulated memory failure on add #{n}"
23772 )));
23773 }
23774 self.messages.write().push(message);
23775 Ok(())
23776 }
23777
23778 async fn get_messages(&self, limit: Option<usize>) -> Result<Vec<ChatMessage>> {
23779 let messages = self.messages.read();
23780 Ok(match limit {
23781 Some(n) if n < messages.len() => messages[messages.len() - n..].to_vec(),
23782 _ => messages.clone(),
23783 })
23784 }
23785
23786 async fn clear(&self) -> Result<()> {
23787 self.messages.write().clear();
23788 Ok(())
23789 }
23790
23791 fn len(&self) -> usize {
23792 self.messages.read().len()
23793 }
23794
23795 async fn restore(&self, snapshot: ai_agents_core::MemorySnapshot) -> Result<()> {
23796 *self.messages.write() = snapshot.messages;
23797 Ok(())
23798 }
23799 }
23800
23801 impl ai_agents_memory::Memory for FailingMemory {}
23802
23803 #[tokio::test]
23804 async fn test_stream_memory_write_failure_surfaces_as_error() {
23805 let yaml = r#"
23808name: TransitionOnToolCallAgent
23809system_prompt: "You are helpful."
23810streaming:
23811 enabled: true
23812states:
23813 initial: intake
23814 states:
23815 intake:
23816 prompt: "Intake"
23817 transitions:
23818 - to: billing
23819 guard:
23820 context:
23821 route:
23822 eq: billing
23823 billing:
23824 prompt: "Billing"
23825"#;
23826 let build = |fail_on_add: usize| {
23827 let mut tools = ai_agents_tools::ToolRegistry::new();
23828 tools
23829 .register(Arc::new(ai_agents_tools::CalculatorTool))
23830 .unwrap();
23831 let agent = AgentBuilder::from_yaml(yaml)
23832 .unwrap()
23833 .llm(Arc::new(mock_with_responses(vec![
23834 r#"{"tool": "calculator", "arguments": {"expression": "1+1"}}"#,
23835 "Billing answer",
23836 ])))
23837 .tools(tools)
23838 .memory(Arc::new(FailingMemory {
23839 messages: parking_lot::RwLock::new(Vec::new()),
23840 fail_on_add,
23841 adds: std::sync::atomic::AtomicUsize::new(0),
23842 }))
23843 .build()
23844 .unwrap();
23845 agent
23846 .set_context("route", serde_json::json!("billing"))
23847 .unwrap();
23848 agent
23849 };
23850
23851 let blocking = build(2).chat("compute").await;
23852 assert!(
23853 blocking.is_err(),
23854 "blocking must surface the memory failure"
23855 );
23856
23857 let (_, chunks, final_response) = collect_stream_events(&build(2), "compute").await;
23858 assert!(
23859 final_response.is_none(),
23860 "streaming must not finalize after a memory failure"
23861 );
23862 assert!(
23863 chunks.iter().any(|c| matches!(c, StreamChunk::Error { message } if message.contains("simulated memory failure"))),
23864 "streaming must surface the memory failure: {chunks:?}"
23865 );
23866
23867 assert!(build(usize::MAX).chat("compute").await.is_ok());
23869 }
23870
23871 #[tokio::test]
23872 async fn parity_tool_execution() {
23873 let build = || {
23874 calculator_agent_with(mock_with_responses(vec![
23875 r#"{"tool": "calculator", "arguments": {"expression": "2+2"}}"#,
23876 "The answer is 4.",
23877 ]))
23878 };
23879 let (blocking, _, chunks) = assert_blocking_streaming_parity(build, "What is 2+2?").await;
23880 assert_eq!(blocking.content, "The answer is 4.");
23881 assert!(
23882 chunks
23883 .iter()
23884 .any(|c| matches!(c, StreamChunk::ToolResult { .. }))
23885 );
23886 }
23887
23888 #[tokio::test]
23889 async fn parity_disambiguation_clarification() {
23890 let build = || {
23891 state_disambiguation_agent(
23892 vec![
23893 r#"{"is_ambiguous":true,"confidence":0.2,"ambiguity_type":"missing_target","reasoning":"target missing","what_is_unclear":["target"],"detected_language":"en"}"#,
23894 r#"{"question":"What should I send?","options":null}"#,
23895 ],
23896 true,
23897 None,
23898 true,
23899 )
23900 .0
23901 };
23902 let (blocking, _, _) = assert_blocking_streaming_parity(build, "Send it").await;
23903 assert_eq!(blocking.content, "What should I send?");
23904 }
23905
23906 #[tokio::test]
23907 async fn parity_reflection_enabled() {
23908 let yaml = r#"
23909name: ReflectionAgent
23910system_prompt: "You are careful."
23911reflection:
23912 enabled: true
23913 criteria:
23914 - "Is the answer helpful?"
23915"#;
23916 let build = || {
23917 AgentBuilder::from_yaml(yaml)
23918 .unwrap()
23919 .llm(Arc::new(mock_with_responses(vec![
23920 "Main answer",
23921 "OVERALL: PASS\nCONFIDENCE: 0.9",
23922 ])))
23923 .build()
23924 .unwrap()
23925 };
23926 let (blocking, streamed, _) = assert_blocking_streaming_parity(build, "hello").await;
23927 assert_eq!(blocking.content, "Main answer");
23928 assert!(metadata_keys(&streamed).contains("reflection"));
23929 }
23930
23931 #[tokio::test]
23932 async fn parity_cot_hidden_thinking() {
23933 let yaml = r#"
23934name: CotHiddenAgent
23935system_prompt: "Think first."
23936reasoning:
23937 mode: cot
23938 output: hidden
23939"#;
23940 let build = || {
23941 AgentBuilder::from_yaml(yaml)
23942 .unwrap()
23943 .llm(Arc::new(mock_with_response(
23944 "<thinking>step by step</thinking>Visible answer",
23945 )))
23946 .build()
23947 .unwrap()
23948 };
23949 let (blocking, streamed, chunks) = assert_blocking_streaming_parity(build, "hello").await;
23950 assert_eq!(blocking.content, "Visible answer");
23951 assert_eq!(content_chunks(&chunks).concat(), streamed.content);
23953 }
23954
23955 fn content_chunks(chunks: &[StreamChunk]) -> Vec<String> {
23960 chunks
23961 .iter()
23962 .filter_map(|c| match c {
23963 StreamChunk::Content { text } => Some(text.clone()),
23964 _ => None,
23965 })
23966 .collect()
23967 }
23968
23969 fn reflection_auto_agent(main: MockLLMProvider, judge: MockLLMProvider) -> RuntimeAgent {
23970 let yaml = r#"
23971name: ReflectionAutoAgent
23972system_prompt: "You are careful."
23973llm:
23974 default: default
23975 router: router
23976reflection:
23977 enabled: auto
23978 evaluator_llm: router
23979 criteria:
23980 - "Is the answer helpful?"
23981"#;
23982 AgentBuilder::from_yaml(yaml)
23983 .unwrap()
23984 .llm_alias("default", Arc::new(main))
23985 .llm_alias("router", Arc::new(judge))
23986 .build()
23987 .unwrap()
23988 }
23989
23990 #[tokio::test]
23991 async fn test_stream_reflection_auto_buffers_and_calls_judge_once_per_iteration() {
23992 let judge = mock_with_responses(vec!["YES", "OVERALL: PASS\nCONFIDENCE: 0.9"]);
23993 let judge_calls = judge.clone();
23994 let agent = reflection_auto_agent(mock_with_response("Main answer one two"), judge);
23995
23996 let (content, chunks, final_response) = collect_stream_events(&agent, "hello").await;
23997
23998 assert!(!chunks.iter().any(StreamChunk::is_error), "{chunks:?}");
23999 assert_eq!(
24000 content_chunks(&chunks).len(),
24001 1,
24002 "auto reflection must buffer the main response: {chunks:?}"
24003 );
24004 assert_eq!(content, "Main answer one two");
24005 assert_eq!(final_response.expect("Final").content, content);
24006 assert_eq!(judge_calls.call_count(), 2);
24008 }
24009
24010 #[tokio::test]
24011 async fn test_stream_reflection_auto_rewrite_is_streamed() {
24012 let judge = mock_with_responses(vec![
24013 "YES",
24014 "OVERALL: FAIL\nCONFIDENCE: 0.1",
24015 "OVERALL: PASS\nCONFIDENCE: 0.9",
24016 ]);
24017 let agent = reflection_auto_agent(
24018 mock_with_responses(vec!["First attempt", "Improved answer"]),
24019 judge,
24020 );
24021
24022 let (content, chunks, final_response) = collect_stream_events(&agent, "hello").await;
24023
24024 assert!(!chunks.iter().any(StreamChunk::is_error), "{chunks:?}");
24025 assert_eq!(content_chunks(&chunks).len(), 1, "{chunks:?}");
24026 assert_eq!(
24027 content, "Improved answer",
24028 "the rewritten answer is what streams"
24029 );
24030 assert_eq!(final_response.expect("Final").content, "Improved answer");
24031 }
24032
24033 fn reasoning_agent(mode: &str, output: &str) -> RuntimeAgent {
24034 let yaml = format!(
24035 r#"
24036name: ReasoningStreamAgent
24037system_prompt: "Think first."
24038reasoning:
24039 mode: {mode}
24040 output: {output}
24041"#
24042 );
24043 AgentBuilder::from_yaml(&yaml)
24044 .unwrap()
24045 .llm(Arc::new(mock_with_response(
24046 "<thinking>step by step</thinking>Visible answer",
24047 )))
24048 .build()
24049 .unwrap()
24050 }
24051
24052 #[tokio::test]
24053 async fn test_stream_cot_hidden_emits_no_thinking_tags() {
24054 let agent = reasoning_agent("cot", "hidden");
24055 let (content, chunks, final_response) = collect_stream_events(&agent, "hello").await;
24056 assert!(!chunks.iter().any(StreamChunk::is_error), "{chunks:?}");
24057 assert_eq!(content_chunks(&chunks).len(), 1, "{chunks:?}");
24058 assert!(!content.contains("<thinking>"), "{content:?}");
24059 assert_eq!(content, "Visible answer");
24060 assert_eq!(final_response.expect("Final").content, content);
24061 }
24062
24063 #[tokio::test]
24064 async fn test_stream_cot_visible_matches_final_format() {
24065 let agent = reasoning_agent("cot", "visible");
24066 let (content, chunks, final_response) = collect_stream_events(&agent, "hello").await;
24067 assert!(!chunks.iter().any(StreamChunk::is_error), "{chunks:?}");
24068 assert!(content.starts_with("Thinking:"), "{content:?}");
24069 assert!(content.contains("Answer:\nVisible answer"), "{content:?}");
24070 assert_eq!(final_response.expect("Final").content, content);
24071 }
24072
24073 #[tokio::test]
24074 async fn test_stream_react_buffers() {
24075 let agent = reasoning_agent("react", "hidden");
24076 let (content, chunks, final_response) = collect_stream_events(&agent, "hello").await;
24077 assert!(!chunks.iter().any(StreamChunk::is_error), "{chunks:?}");
24078 assert_eq!(content_chunks(&chunks).len(), 1, "{chunks:?}");
24079 assert_eq!(content, "Visible answer");
24080 assert_eq!(final_response.expect("Final").content, content);
24081 }
24082
24083 #[tokio::test]
24084 async fn test_stream_plain_mode_still_streams_deltas() {
24085 let agent = AgentBuilder::new()
24086 .system_prompt("You are helpful.")
24087 .llm(Arc::new(mock_with_response("one two three")))
24088 .build()
24089 .unwrap();
24090 let (content, chunks, final_response) = collect_stream_events(&agent, "hello").await;
24091 assert!(
24092 content_chunks(&chunks).len() >= 2,
24093 "plain turns must keep token-level streaming: {chunks:?}"
24094 );
24095 assert_eq!(content, "one two three");
24096 assert_eq!(final_response.expect("Final").content, content);
24097 }
24098
24099 struct ActorProbeHooks {
24101 seen: parking_lot::Mutex<Option<crate::TurnActorContext>>,
24102 }
24103
24104 #[async_trait]
24105 impl AgentHooks for ActorProbeHooks {
24106 async fn on_message_received(&self, _input: &str) {
24107 *self.seen.lock() = current_turn_actor_context();
24108 }
24109 }
24110
24111 fn actor_probe_agent(hooks: Arc<ActorProbeHooks>) -> RuntimeAgent {
24112 let yaml = r#"
24113name: ActorStreamAgent
24114system_prompt: "You are helpful."
24115observability:
24116 enabled: true
24117 export:
24118 write_raw_events: true
24119"#;
24120 AgentBuilder::from_yaml(yaml)
24121 .unwrap()
24122 .llm(Arc::new(mock_with_response("Hello actor")))
24123 .hooks(hooks)
24124 .build()
24125 .unwrap()
24126 }
24127
24128 async fn collect_actor_stream_final(
24129 agent: &RuntimeAgent,
24130 input: &str,
24131 actor_context: crate::TurnActorContext,
24132 ) -> AgentResponse {
24133 use futures::StreamExt;
24134 let mut events = agent
24135 .chat_stream_events_with_actor_context(input, actor_context)
24136 .await
24137 .expect("stream opens");
24138 let mut final_response = None;
24139 while let Some(event) = events.next().await {
24140 match event {
24141 AgentStreamEvent::Final(response) => final_response = Some(response),
24142 AgentStreamEvent::Chunk(StreamChunk::Error { message }) => {
24143 panic!("unexpected stream error: {message}")
24144 }
24145 AgentStreamEvent::Chunk(_) => {}
24146 }
24147 }
24148 final_response.expect("Final")
24149 }
24150
24151 #[tokio::test]
24152 async fn test_stream_events_with_actor_context_scopes_actor_for_turn() {
24153 let hooks = Arc::new(ActorProbeHooks {
24154 seen: parking_lot::Mutex::new(None),
24155 });
24156 let agent = actor_probe_agent(Arc::clone(&hooks));
24157 let actor_context = crate::TurnActorContext::new().with_origin_actor("customer_42");
24158
24159 let final_response = collect_actor_stream_final(&agent, "hi", actor_context).await;
24160
24161 assert_eq!(final_response.content, "Hello actor");
24162 assert_eq!(
24163 hooks
24164 .seen
24165 .lock()
24166 .as_ref()
24167 .and_then(|context| context.effective_actor_id().map(str::to_string)),
24168 Some("customer_42".to_string()),
24169 "the actor context must be visible inside the streaming turn"
24170 );
24171 assert!(
24172 agent.actor_id().is_none(),
24173 "a turn-scoped actor must not mutate the global actor ID"
24174 );
24175 let events = agent.observability().unwrap().raw_events();
24176 assert!(
24177 events
24178 .iter()
24179 .any(|event| event.dimensions.get("actor") == Some(&"customer_42".to_string())),
24180 "observation events must carry the actor dimension"
24181 );
24182 }
24183
24184 #[tokio::test]
24185 async fn test_stream_events_with_actor_context_matches_blocking_actor_context() {
24186 let actor_context = crate::TurnActorContext::new()
24187 .with_origin_actor("customer_42")
24188 .with_sender_agent("coordinator");
24189
24190 let blocking_hooks = Arc::new(ActorProbeHooks {
24191 seen: parking_lot::Mutex::new(None),
24192 });
24193 let blocking_agent = actor_probe_agent(Arc::clone(&blocking_hooks));
24194 let blocking = blocking_agent
24195 .chat_with_actor_context("hi", actor_context.clone())
24196 .await
24197 .unwrap();
24198
24199 let streaming_hooks = Arc::new(ActorProbeHooks {
24200 seen: parking_lot::Mutex::new(None),
24201 });
24202 let streaming_agent = actor_probe_agent(Arc::clone(&streaming_hooks));
24203 let streamed =
24204 collect_actor_stream_final(&streaming_agent, "hi", actor_context.clone()).await;
24205
24206 assert_eq!(blocking.content, streamed.content);
24207 assert_eq!(metadata_keys(&blocking), metadata_keys(&streamed));
24208 assert_eq!(
24209 *blocking_hooks.seen.lock(),
24210 *streaming_hooks.seen.lock(),
24211 "both entry points must expose the same turn actor context"
24212 );
24213 assert_eq!(*streaming_hooks.seen.lock(), Some(actor_context));
24214 }
24215
24216 #[tokio::test]
24217 async fn test_stream_events_with_actor_context_releases_root_turn_on_drop() {
24218 use futures::StreamExt;
24219 let agent = AgentBuilder::new()
24220 .system_prompt("You are helpful.")
24221 .llm(Arc::new(mock_with_response("one two three")))
24222 .build()
24223 .unwrap();
24224 {
24225 let mut events = agent
24226 .chat_stream_events_with_actor_context(
24227 "hi",
24228 crate::TurnActorContext::new().with_origin_actor("customer_42"),
24229 )
24230 .await
24231 .unwrap();
24232 let _first = events.next().await;
24234 }
24235 let next = tokio::time::timeout(Duration::from_secs(5), agent.chat("next")).await;
24236 assert!(
24237 matches!(next, Ok(Ok(_))),
24238 "the root turn must be released when the actor stream is dropped: {next:?}"
24239 );
24240 }
24241
24242 #[tokio::test]
24243 async fn parity_skill_route() {
24244 let yaml = skills_with_parallel_transition_yaml(
24245 " max_speculative_llm_calls_per_turn: 2\n max_parallel_runtime_tasks: 2",
24246 "",
24247 );
24248 let build = || {
24249 build_skills_beside_transition_agent(
24250 &yaml,
24251 role_mocks(
24252 mock_with_response("Draft response"),
24253 mock_with_response("helper"),
24254 ),
24255 )
24256 };
24257 assert_blocking_streaming_parity(build, "please use helper").await;
24258 }
24259
24260 fn required_context_agent(mock: MockLLMProvider, default: bool) -> RuntimeAgent {
24262 let default_yaml = if default {
24263 " default:\n brief: fallback\n"
24264 } else {
24265 ""
24266 };
24267 let yaml = format!(
24268 "name: RequiredContextAgent\nsystem_prompt: 'Voice: {{{{ context.voice.brief }}}}'\ncontext:\n voice:\n type: runtime\n required: true\n{default_yaml}"
24269 );
24270 AgentBuilder::from_yaml(&yaml)
24271 .unwrap()
24272 .llm(Arc::new(mock))
24273 .build()
24274 .unwrap()
24275 }
24276
24277 struct FailOnceContextProvider {
24278 attempts: std::sync::atomic::AtomicUsize,
24279 }
24280
24281 #[async_trait]
24282 impl ContextProvider for FailOnceContextProvider {
24283 async fn get(&self, _key: &str, _current_context: &Value) -> Result<Value> {
24284 if self.attempts.fetch_add(1, Ordering::SeqCst) == 0 {
24285 return Err(AgentError::Other("context initialization failed".into()));
24286 }
24287 Ok(serde_json::json!({"brief": "ready"}))
24288 }
24289 }
24290
24291 #[tokio::test]
24292 async fn test_context_initialization_retries_after_failure() {
24293 let mock = mock_with_response("Voice response");
24294 let calls = mock.clone();
24295 let yaml = "name: CallbackAgent\nsystem_prompt: 'Voice: {{ context.voice.brief }}'\ncontext:\n voice:\n type: callback\n name: flaky\n";
24296 let agent = AgentBuilder::from_yaml(yaml)
24297 .unwrap()
24298 .llm(Arc::new(mock))
24299 .build()
24300 .unwrap();
24301 let provider = Arc::new(FailOnceContextProvider {
24302 attempts: std::sync::atomic::AtomicUsize::new(0),
24303 });
24304 agent.register_context_provider("flaky", provider.clone());
24305
24306 assert!(agent.chat("first").await.is_err());
24307 assert_eq!(calls.call_count(), 0);
24308 assert_eq!(
24309 agent.chat("second").await.unwrap().content,
24310 "Voice response"
24311 );
24312 assert_eq!(provider.attempts.load(Ordering::SeqCst), 2);
24313 assert_eq!(calls.call_count(), 1);
24314 }
24315
24316 #[tokio::test]
24317 async fn test_required_context_blocks_chat_until_supplied_and_after_removal() {
24318 let mock = mock_with_response("Voice response");
24319 let calls = mock.clone();
24320 let agent = required_context_agent(mock, false);
24321
24322 let error = agent.chat("first").await.unwrap_err();
24323 assert!(
24324 error
24325 .to_string()
24326 .contains("Required context 'voice' not provided")
24327 );
24328 assert_eq!(calls.call_count(), 0);
24329 assert!(agent.memory.get_messages(None).await.unwrap().is_empty());
24330
24331 agent
24332 .set_context("voice.brief", serde_json::json!("ready"))
24333 .unwrap();
24334 assert_eq!(
24335 agent.chat("second").await.unwrap().content,
24336 "Voice response"
24337 );
24338 assert_eq!(calls.call_count(), 1);
24339
24340 agent.remove_context("voice");
24341 let error = agent.chat("third").await.unwrap_err();
24342 assert!(
24343 error
24344 .to_string()
24345 .contains("Required context 'voice' not provided")
24346 );
24347 assert_eq!(calls.call_count(), 1);
24348 }
24349
24350 #[tokio::test]
24351 async fn test_required_context_default_satisfies_presence_check() {
24352 let mock = mock_with_response("Fallback response");
24353 let calls = mock.clone();
24354 let agent = required_context_agent(mock, true);
24355
24356 assert_eq!(
24357 agent.chat("hello").await.unwrap().content,
24358 "Fallback response"
24359 );
24360 assert_eq!(
24361 agent.context_manager().get_path("voice.brief"),
24362 Some(serde_json::json!("fallback"))
24363 );
24364 assert_eq!(calls.call_count(), 1);
24365 }
24366
24367 #[tokio::test]
24368 async fn test_required_context_blocks_legacy_stream_before_model_call() {
24369 use futures::StreamExt;
24370
24371 let mock = mock_with_response("Voice response");
24372 let calls = mock.clone();
24373 let agent = required_context_agent(mock, false);
24374 let mut stream = agent.chat_stream("first").await.unwrap();
24375 assert!(
24376 matches!(stream.next().await, Some(StreamChunk::Error { message }) if message.contains("Required context 'voice' not provided"))
24377 );
24378 assert!(stream.next().await.is_none());
24379 drop(stream);
24380 assert_eq!(calls.call_count(), 0);
24381 assert!(agent.memory.get_messages(None).await.unwrap().is_empty());
24382
24383 agent
24384 .set_context("voice.brief", serde_json::json!("ready"))
24385 .unwrap();
24386 assert_eq!(
24387 agent.chat("second").await.unwrap().content,
24388 "Voice response"
24389 );
24390 }
24391
24392 #[tokio::test]
24393 async fn test_required_context_blocks_event_streams_without_final() {
24394 use futures::StreamExt;
24395
24396 for actor_scoped in [false, true] {
24397 let mock = mock_with_response("Voice response");
24398 let calls = mock.clone();
24399 let agent = required_context_agent(mock, false);
24400 let mut events = if actor_scoped {
24401 agent
24402 .chat_stream_events_with_actor_context(
24403 "first",
24404 crate::TurnActorContext::new().with_origin_actor("caller"),
24405 )
24406 .await
24407 .unwrap()
24408 } else {
24409 agent.chat_stream_events("first").await.unwrap()
24410 };
24411 assert!(
24412 matches!(events.next().await, Some(AgentStreamEvent::Chunk(StreamChunk::Error { message })) if message.contains("Required context 'voice' not provided"))
24413 );
24414 assert!(events.next().await.is_none());
24415 drop(events);
24416 assert_eq!(calls.call_count(), 0);
24417 assert!(agent.memory.get_messages(None).await.unwrap().is_empty());
24418
24419 agent
24420 .set_context("voice.brief", serde_json::json!("ready"))
24421 .unwrap();
24422 assert_eq!(
24423 agent.chat("second").await.unwrap().content,
24424 "Voice response"
24425 );
24426 }
24427 }
24428}