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 run_loop(&self, input: &str) -> Result<AgentResponse> {
9432 self.init_storage().await?;
9436 self.begin_root_turn();
9437 let _root_cleanup = RootTurnCleanup::new(self);
9438 info!(input_len = input.len(), "Starting chat");
9439
9440 self.hooks.on_message_received(input).await;
9441
9442 if !self.context_initialized.swap(true, Ordering::SeqCst) {
9446 self.context_manager.initialize().await?;
9447 debug!("Context manager initialized (defaults, env, builtins)");
9448 }
9449
9450 self.check_turn_timeout().await?;
9451 self.context_manager.refresh_per_turn().await?;
9452
9453 self.clear_disambiguation_context();
9456
9457 let input_to_run = match self.resolve_disambiguation(input).await? {
9460 DisambiguationDispatch::Terminal(response) => return Ok(response),
9461 DisambiguationDispatch::RecheckSkill {
9462 skill_id,
9463 enriched_input,
9464 disambiguation_epoch,
9465 state_generation,
9466 } => {
9467 return self
9468 .recheck_skill_disambiguation(
9469 &skill_id,
9470 &enriched_input,
9471 disambiguation_epoch,
9472 state_generation,
9473 )
9474 .await;
9475 }
9476 DisambiguationDispatch::Proceed(input) => input,
9477 };
9478
9479 self.run_loop_internal(&input_to_run).await
9480 }
9481
9482 async fn generate_localized_apology(&self, instruction: &str, reason: &str) -> Result<String> {
9484 let llm = self.llm_registry.router().map_err(|e| {
9485 AgentError::LLM(format!(
9486 "Router LLM not available for localized response: {}",
9487 e
9488 ))
9489 })?;
9490
9491 let recent: Vec<String> = self
9492 .memory
9493 .get_messages(Some(3))
9494 .await?
9495 .iter()
9496 .map(|m| m.content.clone())
9497 .collect();
9498
9499 let context_hint = if recent.is_empty() {
9500 String::new()
9501 } else {
9502 format!(
9503 "\nRecent conversation (detect the user's language from this):\n{}\n",
9504 recent.join("\n")
9505 )
9506 };
9507
9508 let prompt = format!(
9509 "{}\nReason: {}\n{}Respond in the same language as the user. Output ONLY the message, nothing else.",
9510 instruction, reason, context_hint
9511 );
9512
9513 let messages = vec![ChatMessage::user(&prompt)];
9514 let response = self
9515 .observe_purpose(
9516 ObservationPurpose::DisambiguationClarification,
9517 llm.complete(&messages, None),
9518 )
9519 .await
9520 .map_err(|e| AgentError::LLM(format!("Localized response generation failed: {}", e)))?;
9521
9522 Ok(response.content.trim().to_string())
9523 }
9524
9525 fn render_action_args(&self, args: &Value) -> Value {
9529 let context = self.build_context_with_overlays();
9530 match args {
9531 Value::Object(map) => {
9532 let mut rendered = serde_json::Map::new();
9533 for (k, v) in map {
9534 match v {
9535 Value::String(s) if s.contains("{{") => {
9536 match self.template_renderer.render(s, &context) {
9537 Ok(rendered_str) => {
9538 rendered.insert(k.clone(), Value::String(rendered_str));
9539 }
9540 Err(_) => {
9541 rendered.insert(k.clone(), v.clone());
9542 }
9543 }
9544 }
9545 _ => {
9546 rendered.insert(k.clone(), v.clone());
9547 }
9548 }
9549 }
9550 Value::Object(rendered)
9551 }
9552 _ => args.clone(),
9553 }
9554 }
9555
9556 fn clear_disambiguation_context(&self) {
9558 let _ = self
9559 .context_manager
9560 .set("resolved_intent", serde_json::Value::Null);
9561
9562 let all = self.context_manager.get_all();
9563 for key in all.keys() {
9564 if key.starts_with("disambiguation.") {
9565 let _ = self.context_manager.set(key, serde_json::Value::Null);
9566 }
9567 }
9568 }
9569
9570 async fn recheck_skill_disambiguation(
9576 &self,
9577 skill_id: &str,
9578 enriched_input: &str,
9579 expected_disambiguation_epoch: u64,
9580 expected_state_generation: Option<u64>,
9581 ) -> Result<AgentResponse> {
9582 let skill = self
9583 .skill_router
9584 .as_ref()
9585 .and_then(|r| r.get_skill(skill_id).cloned());
9586
9587 if let Some(ref skill) = skill
9589 && let Some(ref skill_disambig) = skill.disambiguation
9590 && skill_disambig.enabled.unwrap_or(false)
9591 && let Some(ref disambiguator) = self.disambiguation_manager
9592 {
9593 let context = self.build_disambiguation_context().await?;
9594 let state_override = self
9595 .state_machine
9596 .as_ref()
9597 .and_then(|sm| sm.current_definition())
9598 .and_then(|def| def.disambiguation.clone());
9599
9600 let disambiguation_result = self
9601 .observe_purpose(
9602 ObservationPurpose::DisambiguationDetection,
9603 disambiguator.process_input_with_override(
9604 enriched_input,
9605 &context,
9606 state_override.as_ref(),
9607 Some(skill_disambig),
9608 ),
9609 )
9610 .await?;
9611 let current_state_generation = self
9612 .state_machine
9613 .as_ref()
9614 .map(|state_machine| state_machine.generation());
9615 if current_state_generation != expected_state_generation
9616 || self.disambiguation_epoch.load(Ordering::SeqCst) != expected_disambiguation_epoch
9617 {
9618 disambiguator.clear_pending().await;
9619 *self.pending_skill_id.write() = None;
9620 return Err(AgentError::Other(
9621 "State or reset ownership changed during skill disambiguation recheck"
9622 .to_string(),
9623 ));
9624 }
9625 match disambiguation_result {
9626 DisambiguationResult::Clear => {
9627 debug!(skill_id = %skill_id, "Skill re-check: all fields present");
9628 }
9629 DisambiguationResult::NeedsClarification {
9630 question,
9631 detection,
9632 } => {
9633 let admission = self
9634 .admit_disambiguation_redispatch(
9635 expected_disambiguation_epoch,
9636 expected_state_generation,
9637 )
9638 .await?;
9639 let awaiting_confirmation = disambiguator.has_pending_confirmation().await;
9640 info!(
9641 skill_id = %skill_id,
9642 ambiguity_type = ?detection.ambiguity_type,
9643 what_is_unclear = ?detection.what_is_unclear,
9644 "Skill re-check: still missing fields, asking again"
9645 );
9646 self.memory
9650 .add_message(ChatMessage::user(enriched_input))
9651 .await?;
9652 self.memory
9653 .add_message(ChatMessage::assistant(&question.question))
9654 .await?;
9655
9656 let response = AgentResponse::new(&question.question).with_metadata(
9657 "disambiguation",
9658 serde_json::json!({
9659 "status": if awaiting_confirmation { "awaiting_confirmation" } else { "awaiting_clarification" },
9660 "skill_id": skill_id,
9661 "options": question.options,
9662 "clarifying": question.clarifying,
9663 "detection": {
9664 "type": detection.ambiguity_type,
9665 "confidence": detection.confidence,
9666 "what_is_unclear": detection.what_is_unclear,
9667 }
9668 }),
9669 );
9670 drop(admission);
9671 self.finish_turn_if_root(&response).await?;
9672 return Ok(response);
9673 }
9674 DisambiguationResult::Clarified {
9675 enriched_input: re_enriched,
9676 ..
9677 } => {
9678 debug!(skill_id = %skill_id, "Skill re-check: clarified immediately, executing");
9679 let admission = self
9680 .admit_disambiguation_redispatch(
9681 expected_disambiguation_epoch,
9682 expected_state_generation,
9683 )
9684 .await?;
9685 *self.pending_skill_id.write() = None;
9686 drop(admission);
9687 let skill_response = self.execute_skill_by_id(skill_id, &re_enriched).await?;
9688 self.memory
9689 .add_message(ChatMessage::user(&re_enriched))
9690 .await?;
9691 return self
9692 .handle_skill_response(
9693 &re_enriched,
9694 skill_id,
9695 skill_response,
9696 &HashMap::new(),
9697 )
9698 .await;
9699 }
9700 DisambiguationResult::ProceedWithBestGuess {
9701 enriched_input: re_enriched,
9702 } => {
9703 debug!(skill_id = %skill_id, "Skill re-check: proceeding with best guess");
9704 let admission = self
9705 .admit_disambiguation_redispatch(
9706 expected_disambiguation_epoch,
9707 expected_state_generation,
9708 )
9709 .await?;
9710 *self.pending_skill_id.write() = None;
9711 drop(admission);
9712 let skill_response = self.execute_skill_by_id(skill_id, &re_enriched).await?;
9713 self.memory
9714 .add_message(ChatMessage::user(&re_enriched))
9715 .await?;
9716 return self
9717 .handle_skill_response(
9718 &re_enriched,
9719 skill_id,
9720 skill_response,
9721 &HashMap::new(),
9722 )
9723 .await;
9724 }
9725 DisambiguationResult::GiveUp { reason } => {
9726 *self.pending_skill_id.write() = None;
9727 let apology = self
9728 .generate_localized_apology(
9729 "Generate a brief, polite apology saying you couldn't understand the request. Be concise.",
9730 &reason,
9731 )
9732 .await
9733 .unwrap_or_else(|_| {
9734 format!("I'm sorry, I couldn't understand your request: {}", reason)
9735 });
9736 let response = AgentResponse::new(&apology);
9737 self.finish_turn_if_root(&response).await?;
9738 return Ok(response);
9739 }
9740 DisambiguationResult::Escalate { reason } => {
9741 *self.pending_skill_id.write() = None;
9742 let apology = self
9743 .generate_localized_apology(
9744 "Explain briefly that you're transferring the user to a human agent for help.",
9745 &reason,
9746 )
9747 .await
9748 .unwrap_or_else(|_| {
9749 format!("I need human assistance to help with your request: {}", reason)
9750 });
9751 let response = AgentResponse::new(&apology);
9752 self.finish_turn_if_root(&response).await?;
9753 return Ok(response);
9754 }
9755 DisambiguationResult::Abandoned { new_input } => {
9756 *self.pending_skill_id.write() = None;
9759 debug!(skill_id = %skill_id, "Skill re-check: abandoned by user");
9760 if let Some(fresh) = new_input {
9761 return self.run_loop_internal(&fresh).await;
9762 }
9763 let ack = self
9764 .generate_localized_apology(
9765 "The user changed their mind about their previous request. \
9766 Generate a brief, friendly acknowledgment (e.g. 'OK, no problem. What else can I help with?'). \
9767 Do NOT apologize excessively. Be concise.",
9768 "User abandoned clarification",
9769 )
9770 .await
9771 .unwrap_or_else(|_| {
9772 "OK, no problem. What else can I help with?".to_string()
9773 });
9774 self.memory
9775 .add_message(ChatMessage::assistant(&ack))
9776 .await?;
9777 let response = AgentResponse::new(&ack);
9778 self.finish_turn_if_root(&response).await?;
9779 return Ok(response);
9780 }
9781 }
9782 }
9783
9784 let admission = self
9786 .admit_disambiguation_redispatch(
9787 expected_disambiguation_epoch,
9788 expected_state_generation,
9789 )
9790 .await?;
9791 *self.pending_skill_id.write() = None;
9792 drop(admission);
9793 let skill_response = self.execute_skill_by_id(skill_id, enriched_input).await?;
9794 self.memory
9795 .add_message(ChatMessage::user(enriched_input))
9796 .await?;
9797 self.handle_skill_response(enriched_input, skill_id, skill_response, &HashMap::new())
9798 .await
9799 }
9800
9801 async fn handle_skill_response(
9804 &self,
9805 processed_input: &str,
9806 skill_id: &str,
9807 skill_response: String,
9808 input_context: &HashMap<String, Value>,
9809 ) -> Result<AgentResponse> {
9810 let output_data = self.process_output(&skill_response, input_context).await?;
9811 let final_response = output_data.content;
9812
9813 self.memory
9814 .add_message(ChatMessage::assistant(&final_response))
9815 .await?;
9816
9817 self.check_memory_compression().await?;
9818
9819 self.increment_turn();
9820 self.evaluate_transitions(processed_input, &final_response)
9821 .await?;
9822
9823 let response = AgentResponse::new(final_response)
9824 .with_metadata("skill_id", serde_json::json!(skill_id));
9825 self.finish_turn_if_root(&response).await?;
9826 Ok(response)
9827 }
9828
9829 async fn handle_plan_and_execute(
9832 &self,
9833 processed_input: &str,
9834 input_context: &HashMap<String, Value>,
9835 auto_detected: bool,
9836 ) -> Result<AgentResponse> {
9837 let effective = self.get_effective_reasoning_config();
9838 let plan_reflection = effective
9839 .get_planning()
9840 .map(|c| c.reflection.clone())
9841 .unwrap_or_default();
9842
9843 let max_attempts = if plan_reflection.enabled {
9844 1 + plan_reflection.max_replans
9845 } else {
9846 1
9847 };
9848
9849 let mut plan = self.generate_plan(processed_input).await?;
9850 info!(
9851 plan_id = %plan.id,
9852 steps = plan.steps.len(),
9853 "Plan generated"
9854 );
9855
9856 let mut plan_result = String::new();
9857
9858 for attempt in 0..max_attempts {
9859 *self.current_plan.write() = Some(plan.clone());
9860 plan_result = self.execute_plan(&mut plan).await?;
9861
9862 info!(
9863 plan_status = ?plan.status,
9864 completed_steps = plan.completed_steps().count(),
9865 attempt = attempt + 1,
9866 "Plan execution completed"
9867 );
9868
9869 if !plan_reflection.enabled {
9870 break;
9871 }
9872
9873 let has_failures = plan.steps.iter().any(|s| s.status.is_failed());
9874 if !has_failures {
9875 break;
9876 }
9877
9878 if attempt + 1 >= max_attempts {
9879 break;
9880 }
9881
9882 match plan_reflection.on_step_failure {
9883 StepFailureAction::Replan => {
9884 info!(attempt = attempt + 1, "Plan had failures, replanning");
9885 plan = self.generate_plan(processed_input).await?;
9886 }
9887 StepFailureAction::Abort => {
9888 warn!("Plan step failed, aborting");
9889 break;
9890 }
9891 StepFailureAction::Skip | StepFailureAction::Continue => {
9892 break;
9893 }
9894 }
9895 }
9896
9897 *self.current_plan.write() = Some(plan);
9898
9899 let output_data = self.process_output(&plan_result, input_context).await?;
9900 let final_content = output_data.content;
9901
9902 self.memory
9903 .add_message(ChatMessage::assistant(&final_content))
9904 .await?;
9905
9906 self.check_memory_compression().await?;
9907 self.increment_turn();
9908 self.evaluate_transitions(processed_input, &final_content)
9909 .await?;
9910
9911 let reasoning_metadata =
9912 ReasoningMetadata::new(ReasoningMode::PlanAndExecute).with_auto_detected(auto_detected);
9913
9914 let response = AgentResponse::new(&final_content).with_metadata(
9915 "reasoning",
9916 serde_json::to_value(&reasoning_metadata).unwrap_or_default(),
9917 );
9918
9919 self.finish_turn_if_root(&response).await?;
9920 Ok(response)
9921 }
9922
9923 fn inject_reasoning_prompt(
9925 &self,
9926 messages: &mut [ChatMessage],
9927 reasoning_mode: &ReasoningMode,
9928 is_first_iteration: bool,
9929 ) {
9930 if !is_first_iteration {
9931 return;
9932 }
9933 match reasoning_mode {
9934 ReasoningMode::CoT => {
9935 if let Some(msg) = messages.first_mut()
9936 && matches!(msg.role, ai_agents_core::Role::System)
9937 {
9938 msg.content = self.build_cot_system_prompt(&msg.content);
9939 debug!("Applied Chain-of-Thought system prompt");
9940 }
9941 }
9942 ReasoningMode::React => {
9943 if let Some(msg) = messages.first_mut()
9944 && matches!(msg.role, ai_agents_core::Role::System)
9945 {
9946 msg.content = self.build_react_system_prompt(&msg.content);
9947 debug!("Applied ReAct system prompt");
9948 }
9949 }
9950 _ => {}
9951 }
9952 }
9953
9954 async fn generate_main_response_draft(
9959 &self,
9960 processed_input: &str,
9961 reasoning_mode: &ReasoningMode,
9962 ) -> Result<MainResponseDraft> {
9963 let llm = self.get_state_llm()?;
9964 let protocol = self.main_tool_protocol(llm.as_ref(), true).await?;
9965 let mut messages = self
9966 .build_messages_internal(false, Some(processed_input), protocol.choice.is_none())
9967 .await?;
9968 self.inject_reasoning_prompt(&mut messages, reasoning_mode, true);
9969 let response = self
9970 .complete_main_llm_with_recovery(llm, &messages, &protocol)
9971 .await?;
9972 let content = response.content.trim().to_string();
9973 let (thinking, answer) = self.extract_thinking(&content);
9974 if let Some(calls) = self.parse_main_tool_calls(&content, &protocol)? {
9975 return Ok(MainResponseDraft::ToolCalls {
9976 raw_content: content,
9977 calls,
9978 thinking,
9979 });
9980 }
9981 Ok(MainResponseDraft::Text {
9982 raw_content: answer,
9983 thinking,
9984 })
9985 }
9986
9987 async fn commit_main_response_draft(
9992 &self,
9993 processed_input: &str,
9994 input_context: &HashMap<String, Value>,
9995 draft: MainResponseDraft,
9996 reasoning_mode: ReasoningMode,
9997 auto_detected: bool,
9998 ) -> Result<AgentResponse> {
9999 self.commit_root_user_message(processed_input).await?;
10000 match draft {
10001 MainResponseDraft::Text {
10002 raw_content,
10003 thinking,
10004 } => {
10005 self.finish_text_response_from_model(CommittedTextResponse {
10006 processed_input,
10007 input_context,
10008 answer: raw_content,
10009 reasoning_mode,
10010 auto_detected,
10011 iterations: 1,
10012 thinking_content: thinking,
10013 all_tool_calls: Vec::new(),
10014 })
10015 .await
10016 }
10017 MainResponseDraft::ToolCalls {
10018 raw_content,
10019 calls,
10020 thinking: _,
10021 } => {
10022 let mut all_tool_calls = Vec::new();
10023 match self
10024 .handle_tool_calls(
10025 processed_input,
10026 &raw_content,
10027 calls,
10028 &mut all_tool_calls,
10029 None,
10030 )
10031 .await?
10032 {
10033 ToolCallOutcome::Rejected(response) => {
10034 self.finish_turn_if_root(&response).await?;
10035 Ok(response)
10036 }
10037 ToolCallOutcome::Continue | ToolCallOutcome::TransitionFired => {
10038 self.continue_after_committed_tool_draft(processed_input)
10039 .await
10040 }
10041 }
10042 }
10043 }
10044 }
10045
10046 async fn continue_after_committed_tool_draft(
10051 &self,
10052 processed_input: &str,
10053 ) -> Result<AgentResponse> {
10054 *self.redispatch_depth.write() += 1;
10055 if let Some(context) = self.active_turn_context.write().as_mut() {
10056 context.enter_redispatch();
10057 }
10058 let result = Box::pin(self.run_loop_internal(processed_input)).await;
10059 *self.redispatch_depth.write() -= 1;
10060 if let Some(context) = self.active_turn_context.write().as_mut() {
10061 context.exit_redispatch();
10062 }
10063 let response = result?;
10064 self.finish_turn_if_root(&response).await?;
10065 Ok(response)
10066 }
10067
10068 async fn finish_text_response_from_model(
10073 &self,
10074 response: CommittedTextResponse<'_>,
10075 ) -> Result<AgentResponse> {
10076 let CommittedTextResponse {
10077 processed_input,
10078 input_context,
10079 answer,
10080 reasoning_mode,
10081 auto_detected,
10082 iterations,
10083 thinking_content,
10084 all_tool_calls,
10085 } = response;
10086 let output_data = self.process_output(&answer, input_context).await?;
10087 let mut final_content = if output_data.metadata.rejected {
10088 output_data
10089 .metadata
10090 .rejection_reason
10091 .unwrap_or_else(|| answer.to_string())
10092 } else {
10093 output_data.content
10094 };
10095 let llm = self.get_state_llm()?;
10096 let reflection_metadata;
10097 (final_content, reflection_metadata) = self
10098 .run_reflection(&*llm, processed_input, final_content)
10099 .await?;
10100 final_content =
10101 self.format_response_with_thinking(thinking_content.as_deref(), &final_content);
10102 let final_content = {
10103 let result = self
10104 .post_loop_processing(processed_input, final_content)
10105 .await?;
10106 self.apply_post_loop_result(processed_input, result)
10107 .await?
10108 .content
10109 };
10110 let response = self.build_agent_response(AgentResponseParts {
10111 content: final_content,
10112 all_tool_calls,
10113 reasoning_mode,
10114 auto_detected,
10115 iterations,
10116 thinking: thinking_content,
10117 reflection_metadata,
10118 });
10119 self.finish_turn_if_root(&response).await?;
10120 Ok(response)
10121 }
10122
10123 async fn run_committed_response_loop_with_reasoning(
10128 &self,
10129 processed_input: &str,
10130 input_context: &HashMap<String, Value>,
10131 reasoning_mode: ReasoningMode,
10132 auto_detected: bool,
10133 ) -> Result<AgentResponse> {
10134 self.commit_root_user_message(processed_input).await?;
10135 let llm = self.get_state_llm()?;
10136 let mut iterations = 0u32;
10137 let mut all_tool_calls = Vec::new();
10138 let mut thinking_content = None;
10139 loop {
10140 let effective_max = if reasoning_mode != ReasoningMode::None {
10141 let rc = self.get_effective_reasoning_config();
10142 self.max_iterations.min(rc.max_iterations)
10143 } else {
10144 self.max_iterations
10145 };
10146 if iterations >= effective_max {
10147 return Err(AgentError::Other(format!(
10148 "Max iterations ({}) exceeded",
10149 effective_max
10150 )));
10151 }
10152 iterations += 1;
10153 *self.iteration_count.write() = iterations;
10154 let protocol = self.main_tool_protocol(llm.as_ref(), false).await?;
10155 let mut messages = self
10156 .build_messages_internal(true, None, protocol.choice.is_none())
10157 .await?;
10158 self.inject_reasoning_prompt(&mut messages, &reasoning_mode, iterations == 1);
10159 self.hooks.on_llm_start(&messages).await;
10160 let llm_start = Instant::now();
10161 let response = self
10162 .complete_main_llm_with_recovery(Arc::clone(&llm), &messages, &protocol)
10163 .await?;
10164 let llm_duration_ms = llm_start.elapsed().as_millis() as u64;
10165 self.hooks.on_llm_complete(&response, llm_duration_ms).await;
10166 let content = response.content.trim();
10167 if let Some(tool_calls) = self.parse_main_tool_calls(content, &protocol)? {
10168 match self
10169 .handle_tool_calls(
10170 processed_input,
10171 content,
10172 tool_calls,
10173 &mut all_tool_calls,
10174 None,
10175 )
10176 .await?
10177 {
10178 ToolCallOutcome::Continue | ToolCallOutcome::TransitionFired => continue,
10179 ToolCallOutcome::Rejected(resp) => {
10180 self.finish_turn_if_root(&resp).await?;
10181 return Ok(resp);
10182 }
10183 }
10184 }
10185 let (extracted_thinking, answer) = self.extract_thinking(content);
10186 if extracted_thinking.is_some() {
10187 thinking_content = extracted_thinking;
10188 }
10189 return self
10190 .finish_text_response_from_model(CommittedTextResponse {
10191 processed_input,
10192 input_context,
10193 answer,
10194 reasoning_mode,
10195 auto_detected,
10196 iterations,
10197 thinking_content,
10198 all_tool_calls,
10199 })
10200 .await;
10201 }
10202 }
10203
10204 async fn handle_tool_calls(
10210 &self,
10211 processed_input: &str,
10212 content: &str,
10213 tool_calls: Vec<ToolCall>,
10214 all_tool_calls: &mut Vec<ToolCall>,
10215 mut events: Option<&mut Vec<StreamChunk>>,
10216 ) -> Result<ToolCallOutcome> {
10217 let include_tool_events = self.streaming.include_tool_events;
10218 let transition_content = native_readable_projection(content)
10222 .map_err(|error| AgentError::LLM(error.to_string()))?;
10223 let transition_fired = self
10224 .evaluate_transitions(processed_input, &transition_content)
10225 .await?;
10226 if transition_fired {
10227 self.memory
10228 .add_message(ChatMessage::assistant(
10229 "(Transitioned to new state — tool call handled by workflow)",
10230 ))
10231 .await?;
10232 if let Some(events) = events.as_deref_mut()
10233 && self.streaming.include_state_events
10234 && let Some(state) = self.current_state()
10235 {
10236 events.push(StreamChunk::state_transition(None, state));
10237 }
10238 return Ok(ToolCallOutcome::TransitionFired);
10239 }
10240
10241 self.memory
10243 .add_message(ChatMessage::assistant(content))
10244 .await?;
10245 self.remember_committed_native_exchange(content).await?;
10246 let native_tool_call = Self::is_native_tool_call_content(content)?;
10247
10248 if let Some(events) = events.as_deref_mut()
10249 && include_tool_events
10250 {
10251 for tool_call in &tool_calls {
10252 events.push(StreamChunk::tool_start(&tool_call.id, &tool_call.name));
10253 }
10254 }
10255 let results = self.execute_tools_parallel(&tool_calls).await;
10256 let mut rejection = None;
10257
10258 for ((_id, result), tool_call) in results.into_iter().zip(tool_calls.iter()) {
10259 match result {
10260 Ok(output) => {
10261 if let Some(events) = events.as_deref_mut()
10262 && include_tool_events
10263 {
10264 events.push(StreamChunk::tool_result(
10265 &tool_call.id,
10266 &tool_call.name,
10267 &output,
10268 true,
10269 ));
10270 }
10271 self.memory
10272 .add_message(Self::tool_result_message(
10273 tool_call,
10274 &output,
10275 native_tool_call,
10276 )?)
10277 .await?;
10278 }
10279 Err(e) => {
10280 if matches!(e, AgentError::HITLRejected(_)) {
10281 if !native_tool_call {
10282 self.memory
10283 .add_message(ChatMessage::assistant(format!(
10284 "The operation was rejected by the approver: {e}"
10285 )))
10286 .await?;
10287 return Ok(ToolCallOutcome::Rejected(AgentResponse {
10288 content: format!("Operation cancelled: {e}"),
10289 metadata: None,
10290 tool_calls: Some(all_tool_calls.clone()),
10291 }));
10292 }
10293 if rejection.is_none() {
10294 rejection = Some(e.to_string());
10295 }
10296 }
10297 if let Some(events) = events.as_deref_mut()
10298 && include_tool_events
10299 {
10300 events.push(StreamChunk::tool_result(
10301 &tool_call.id,
10302 &tool_call.name,
10303 e.to_string(),
10304 false,
10305 ));
10306 }
10307 self.memory
10308 .add_message(Self::tool_result_message(
10309 tool_call,
10310 &format!("Error: {}", e),
10311 native_tool_call,
10312 )?)
10313 .await?;
10314 }
10315 }
10316 all_tool_calls.push(tool_call.clone());
10317 if let Some(events) = events.as_deref_mut()
10318 && include_tool_events
10319 {
10320 events.push(StreamChunk::tool_end(&tool_call.id));
10321 }
10322 }
10323 if let Some(rejection) = rejection {
10324 self.memory
10325 .add_message(ChatMessage::assistant(format!(
10326 "The operation was rejected by the approver: {rejection}"
10327 )))
10328 .await?;
10329 return Ok(ToolCallOutcome::Rejected(AgentResponse {
10330 content: format!("Operation cancelled: {rejection}"),
10331 metadata: None,
10332 tool_calls: Some(all_tool_calls.clone()),
10333 }));
10334 }
10335 Ok(ToolCallOutcome::Continue)
10336 }
10337
10338 async fn run_reflection(
10340 &self,
10341 llm: &dyn LLMProvider,
10342 processed_input: &str,
10343 mut content: String,
10344 ) -> Result<(String, Option<ReflectionMetadata>)> {
10345 let should_reflect = self.should_reflect(processed_input, &content).await?;
10346 if !should_reflect {
10347 return Ok((content, None));
10348 }
10349
10350 info!("Starting response reflection evaluation");
10351 let mut attempts = 0u32;
10352 let max_retries = self.reflection_config.max_retries;
10353 let mut history: Vec<ReflectionAttempt> = Vec::new();
10354
10355 loop {
10356 let evaluation = self.evaluate_response(processed_input, &content).await?;
10357
10358 if evaluation.passed || attempts >= max_retries {
10359 info!(
10360 passed = evaluation.passed,
10361 confidence = evaluation.confidence,
10362 attempts = attempts + 1,
10363 "Reflection evaluation complete"
10364 );
10365 let reflection_metadata = Some(
10366 ReflectionMetadata::new(evaluation)
10367 .with_attempts(attempts + 1)
10368 .with_history(history),
10369 );
10370 return Ok((content, reflection_metadata));
10371 }
10372
10373 debug!(
10374 attempt = attempts + 1,
10375 failed_criteria = evaluation.failed_criteria().count(),
10376 "Response did not meet criteria, retrying"
10377 );
10378
10379 history.push(
10380 ReflectionAttempt::new(&content, evaluation.clone())
10381 .with_feedback("Response did not meet quality criteria"),
10382 );
10383
10384 let feedback: Vec<String> = evaluation
10385 .failed_criteria()
10386 .map(|c| format!("- {}", c.criterion))
10387 .collect();
10388
10389 let retry_prompt = format!(
10390 "Your previous response did not meet these criteria:\n{}\n\nPlease provide an improved response.",
10391 feedback.join("\n")
10392 );
10393
10394 self.memory
10395 .add_message(ChatMessage::user(&retry_prompt))
10396 .await?;
10397
10398 let retry_messages = self.build_messages().await?;
10399 let retry_response = self
10400 .observe_purpose(
10401 ObservationPurpose::ReflectionEvaluation,
10402 llm.complete(&retry_messages, None),
10403 )
10404 .await
10405 .map_err(|e| AgentError::LLM(e.to_string()))?;
10406
10407 content = retry_response.content.trim().to_string();
10408 attempts += 1;
10409 }
10410 }
10411
10412 async fn post_loop_processing(
10415 &self,
10416 processed_input: &str,
10417 content: String,
10418 ) -> Result<PostLoopResult> {
10419 self.increment_turn();
10424
10425 self.run_context_extractors(processed_input).await;
10427
10428 let transitioned = self.evaluate_transitions(processed_input, &content).await?;
10429
10430 if !transitioned {
10431 self.memory
10432 .add_message(ChatMessage::assistant(&content))
10433 .await?;
10434 self.check_memory_compression().await?;
10435 return Ok(PostLoopResult::NoTransition(content));
10436 }
10437
10438 if !self.should_regenerate_after_transition() {
10440 self.memory
10441 .add_message(ChatMessage::assistant(&content))
10442 .await?;
10443 self.check_memory_compression().await?;
10444 return Ok(PostLoopResult::Transitioned {
10445 content,
10446 regenerated: false,
10447 });
10448 }
10449
10450 if self.needs_redispatch_for_new_state() {
10454 info!("Post-transition NeedsRedispatch: new state requires full dispatch");
10455 return Ok(PostLoopResult::NeedsRedispatch);
10458 }
10459
10460 self.memory
10463 .add_message(ChatMessage::assistant(&content))
10464 .await?;
10465 self.check_memory_compression().await?;
10466
10467 let new_llm = self.get_state_llm()?;
10473 let mut final_content;
10474
10475 for post_iter in 0..self.max_iterations {
10476 let protocol = self.main_tool_protocol(new_llm.as_ref(), false).await?;
10477 let new_messages = self
10478 .build_messages_internal(true, None, protocol.choice.is_none())
10479 .await?;
10480 if post_iter == 0
10481 && let Some(system_msg) = new_messages.first()
10482 && system_msg.role == ai_agents_core::Role::System
10483 {
10484 debug!(
10485 prompt_preview =
10486 &system_msg.content[system_msg.content.len().saturating_sub(200)..],
10487 "Post-transition system prompt (last 200 chars)"
10488 );
10489 }
10490
10491 let new_response = self
10492 .complete_main_llm_with_recovery(Arc::clone(&new_llm), &new_messages, &protocol)
10493 .await?;
10494 final_content = new_response.content.trim().to_string();
10495
10496 if let Some(tool_calls) = self.parse_main_tool_calls(&final_content, &protocol)? {
10499 let native_tool_call = Self::is_native_tool_call_content(&final_content)?;
10500 debug!(
10501 post_iter = post_iter,
10502 tools = tool_calls.len(),
10503 "Post-transition tool call detected, executing"
10504 );
10505
10506 self.memory
10507 .add_message(ChatMessage::assistant(&final_content))
10508 .await?;
10509 self.remember_committed_native_exchange(&final_content)
10510 .await?;
10511
10512 let results = self.execute_tools_parallel(&tool_calls).await;
10513 let mut rejection = None;
10514 for ((_id, result), tool_call) in results.into_iter().zip(tool_calls.iter()) {
10515 match result {
10516 Ok(output) => {
10517 self.memory
10518 .add_message(Self::tool_result_message(
10519 tool_call,
10520 &output,
10521 native_tool_call,
10522 )?)
10523 .await?;
10524 }
10525 Err(e) => {
10526 if native_tool_call
10527 && rejection.is_none()
10528 && matches!(e, AgentError::HITLRejected(_))
10529 {
10530 rejection = Some(e.to_string());
10531 }
10532 self.memory
10533 .add_message(Self::tool_result_message(
10534 tool_call,
10535 &format!("Error: {}", e),
10536 native_tool_call,
10537 )?)
10538 .await?;
10539 }
10540 }
10541 }
10542 if let Some(rejection) = rejection {
10543 self.memory
10544 .add_message(ChatMessage::assistant(format!(
10545 "The operation was rejected by the approver: {rejection}"
10546 )))
10547 .await?;
10548 return Err(AgentError::HITLRejected(rejection));
10549 }
10550 continue;
10552 }
10553
10554 self.memory
10556 .add_message(ChatMessage::assistant(&final_content))
10557 .await?;
10558 return Ok(PostLoopResult::Transitioned {
10559 content: final_content,
10560 regenerated: true,
10561 });
10562 }
10563
10564 final_content = "Post-transition processing completed.".to_string();
10566 self.memory
10567 .add_message(ChatMessage::assistant(&final_content))
10568 .await?;
10569
10570 Ok(PostLoopResult::Transitioned {
10571 content: final_content,
10572 regenerated: true,
10573 })
10574 }
10575
10576 fn should_regenerate_after_transition(&self) -> bool {
10579 if let Some(ref sm) = self.state_machine {
10580 if !sm.config().regenerate_on_transition {
10582 return false;
10583 }
10584 if let Some(def) = sm.current_definition()
10586 && let Some(regen) = def.regenerate_on_enter
10587 {
10588 return regen;
10589 }
10590 }
10591 true
10592 }
10593
10594 fn needs_redispatch_for_new_state(&self) -> bool {
10597 if let Some(ref sm) = self.state_machine
10598 && let Some(def) = sm.current_definition()
10599 {
10600 if def.concurrent.is_some()
10601 || def.group_chat.is_some()
10602 || def.pipeline.is_some()
10603 || def.handoff.is_some()
10604 || def.delegate.is_some()
10605 {
10606 return true;
10607 }
10608 let effective = self.get_effective_reasoning_config();
10610 if !matches!(effective.mode, ReasoningMode::None) {
10611 return true;
10612 }
10613 }
10614 false
10615 }
10616
10617 async fn apply_post_loop_result(
10623 &self,
10624 processed_input: &str,
10625 result: PostLoopResult,
10626 ) -> Result<AppliedPostLoop> {
10627 match result {
10628 PostLoopResult::NoTransition(content) => Ok(AppliedPostLoop {
10629 content,
10630 transitioned: false,
10631 regenerated: false,
10632 }),
10633 PostLoopResult::Transitioned {
10634 content,
10635 regenerated,
10636 } => Ok(AppliedPostLoop {
10637 content,
10638 transitioned: true,
10639 regenerated,
10640 }),
10641 PostLoopResult::NeedsRedispatch => {
10642 const MAX_REDISPATCH_DEPTH: u32 = 3;
10643 let current_depth = *self.redispatch_depth.read();
10644 if current_depth >= MAX_REDISPATCH_DEPTH {
10645 warn!(
10646 depth = current_depth,
10647 "Post-transition re-dispatch depth limit reached, returning empty response"
10648 );
10649 let content = String::new();
10650 self.memory
10651 .add_message(ChatMessage::assistant(&content))
10652 .await?;
10653 return Ok(AppliedPostLoop {
10655 content,
10656 transitioned: true,
10657 regenerated: false,
10658 });
10659 }
10660 *self.redispatch_depth.write() += 1;
10661 if let Some(context) = self.active_turn_context.write().as_mut() {
10662 context.enter_redispatch();
10663 }
10664 info!(
10665 depth = current_depth + 1,
10666 "Re-dispatching for new state after transition"
10667 );
10668 let resp = Box::pin(self.run_loop_internal(processed_input)).await;
10669 *self.redispatch_depth.write() -= 1;
10670 if let Some(context) = self.active_turn_context.write().as_mut() {
10671 context.exit_redispatch();
10672 }
10673 resp.map(|r| AppliedPostLoop {
10674 content: r.content,
10675 transitioned: true,
10676 regenerated: true,
10677 })
10678 }
10679 }
10680 }
10681
10682 fn build_agent_response(&self, parts: AgentResponseParts) -> AgentResponse {
10684 let AgentResponseParts {
10685 content,
10686 all_tool_calls,
10687 reasoning_mode,
10688 auto_detected,
10689 iterations,
10690 thinking,
10691 reflection_metadata,
10692 } = parts;
10693 let reasoning_metadata = ReasoningMetadata::new(reasoning_mode.clone())
10694 .with_thinking(thinking.clone().unwrap_or_default())
10695 .with_iterations(iterations)
10696 .with_auto_detected(auto_detected);
10697
10698 let mut response = AgentResponse::new(&content);
10699 if !all_tool_calls.is_empty() {
10700 response = response.with_tool_calls(all_tool_calls);
10701 }
10702
10703 if let Some(state) = self.current_state() {
10704 response = response.with_metadata("current_state", serde_json::json!(state));
10705 }
10706
10707 response = response.with_metadata(
10708 "reasoning",
10709 serde_json::to_value(&reasoning_metadata).unwrap_or_default(),
10710 );
10711
10712 if let Some(ref refl_meta) = reflection_metadata {
10713 response = response.with_metadata(
10714 "reflection",
10715 serde_json::to_value(refl_meta).unwrap_or_default(),
10716 );
10717 }
10718
10719 response
10720 }
10721
10722 async fn handle_delegated_state(
10724 &self,
10725 input: &str,
10726 delegate_id: &str,
10727 state_def: &ai_agents_state::StateDefinition,
10728 ) -> Result<AgentResponse> {
10729 use std::time::Instant;
10730
10731 let registry = self.spawner_registry.as_ref().ok_or_else(|| {
10732 AgentError::Config(format!(
10733 "State delegates to '{}' but no agent registry is configured. \
10734 Add a spawner section with auto_spawn to your YAML.",
10735 delegate_id
10736 ))
10737 })?;
10738
10739 let state_name = self
10740 .state_machine
10741 .as_ref()
10742 .map(|sm| sm.current())
10743 .unwrap_or_else(|| "unknown".to_string());
10744
10745 self.hooks.on_delegate_start(delegate_id, &state_name).await;
10746 let start = Instant::now();
10747
10748 let delegate = registry.get(delegate_id).ok_or_else(|| {
10749 AgentError::Other(format!(
10750 "State '{}' delegates to '{}' but no agent with that ID exists in the registry.",
10751 state_name, delegate_id
10752 ))
10753 })?;
10754
10755 let context_mode = state_def.delegate_context.clone().unwrap_or_default();
10757 let effective_input = self
10758 .observe_purpose(
10759 ObservationPurpose::OrchestrationRouting,
10760 crate::orchestration::context::prepare_delegate_input(
10761 input,
10762 &context_mode,
10763 &*self.memory,
10764 self.llm_registry.get("router").ok().as_deref(),
10765 ),
10766 )
10767 .await?;
10768
10769 let response = delegate
10770 .chat_with_actor_context(&effective_input, self.outbound_actor_context())
10771 .await?;
10772
10773 let duration_ms = start.elapsed().as_millis() as u64;
10774 self.hooks
10775 .on_delegate_complete(delegate_id, &state_name, duration_ms)
10776 .await;
10777
10778 let ctx_key = format!("delegation.{}.last_response", delegate_id);
10780 let _ = self.context_manager.set(
10781 &ctx_key,
10782 serde_json::Value::String(response.content.clone()),
10783 );
10784
10785 let _ = self.context_manager.set(
10787 "orchestration",
10788 serde_json::json!({
10789 "type": "delegate",
10790 "agent": delegate_id,
10791 "state": state_name,
10792 "response": response.content,
10793 "duration_ms": duration_ms,
10794 }),
10795 );
10796
10797 self.commit_root_user_message(input).await?;
10798
10799 let post_result = self
10802 .post_loop_processing(
10803 input,
10804 format!("[Delegated to {}]: {}", delegate_id, response.content),
10805 )
10806 .await?;
10807 let final_content = self
10808 .apply_post_loop_result(input, post_result)
10809 .await?
10810 .content;
10811
10812 let mut result = AgentResponse::new(final_content);
10813
10814 let metadata = serde_json::json!({
10815 "orchestration": {
10816 "type": "delegate",
10817 "agent": delegate_id,
10818 "state": state_name,
10819 "response": response.content,
10820 "duration_ms": duration_ms,
10821 }
10822 });
10823 result.metadata = Some(
10824 serde_json::from_value::<std::collections::HashMap<String, serde_json::Value>>(
10825 metadata,
10826 )
10827 .unwrap_or_default(),
10828 );
10829
10830 self.finish_turn_if_root(&result).await?;
10831 Ok(result)
10832 }
10833
10834 async fn handle_concurrent_state(
10836 &self,
10837 input: &str,
10838 config: &ai_agents_state::ConcurrentStateConfig,
10839 ) -> Result<AgentResponse> {
10840 use std::time::Instant;
10841
10842 let registry = self.spawner_registry.as_ref().ok_or_else(|| {
10843 AgentError::Config(
10844 "Concurrent state requires an agent registry. Add a spawner section.".into(),
10845 )
10846 })?;
10847
10848 let context_mode = config.context_mode.clone().unwrap_or_default();
10853 let context_input = self
10854 .observe_purpose(
10855 ObservationPurpose::OrchestrationRouting,
10856 crate::orchestration::context::prepare_delegate_input(
10857 input,
10858 &context_mode,
10859 &*self.memory,
10860 self.llm_registry.get("router").ok().as_deref(),
10861 ),
10862 )
10863 .await?;
10864
10865 let effective_input = if let Some(ref tmpl) = config.input {
10866 render_concurrent_template(tmpl, &context_input, &self.build_context_with_overlays())
10867 .unwrap_or_else(|_| context_input.clone())
10868 } else {
10869 context_input
10870 };
10871
10872 let start = Instant::now();
10873
10874 let llm_name = config
10875 .aggregation
10876 .synthesizer_llm
10877 .as_deref()
10878 .unwrap_or("router");
10879 let llm_provider = self.llm_registry.get(llm_name).ok();
10880
10881 let vote_parallelism = if self.runtime_config.optimization.enabled
10882 && self
10883 .runtime_config
10884 .optimization
10885 .parallel_orchestration_vote_extraction
10886 {
10887 Some(self.runtime_config.optimization.max_parallel_runtime_tasks)
10888 } else {
10889 None
10890 };
10891
10892 let result = self
10893 .observe_purpose(
10894 ObservationPurpose::OrchestrationAggregation,
10895 scope_actor_context(
10896 self.outbound_actor_context(),
10897 crate::orchestration::concurrent(
10898 registry,
10899 &effective_input,
10900 &config.agents,
10901 &config.aggregation,
10902 llm_provider.as_deref(),
10903 config.min_required,
10904 config.timeout_ms,
10905 config.on_partial_failure.clone(),
10906 vote_parallelism,
10907 ),
10908 ),
10909 )
10910 .await?;
10911
10912 let duration_ms = start.elapsed().as_millis() as u64;
10913 let agent_ids: Vec<String> = config.agents.iter().map(|a| a.id().to_string()).collect();
10914 let strategy = format!("{:?}", config.aggregation.strategy);
10915 self.hooks
10916 .on_concurrent_complete(&agent_ids, &strategy, duration_ms)
10917 .await;
10918
10919 let _ = self.context_manager.set(
10921 "concurrent.result",
10922 serde_json::Value::String(result.response.content.clone()),
10923 );
10924
10925 let agents_json: Vec<serde_json::Value> = result
10927 .agent_results
10928 .iter()
10929 .map(|ar| {
10930 serde_json::json!({
10931 "id": ar.agent_id,
10932 "response": ar.response.as_ref().map(|r| r.content.as_str()),
10933 "success": ar.success,
10934 "error": ar.error,
10935 "duration_ms": ar.duration_ms,
10936 })
10937 })
10938 .collect();
10939
10940 let _ = self.context_manager.set(
10942 "orchestration",
10943 serde_json::json!({
10944 "type": "concurrent",
10945 "result": result.response.content,
10946 "strategy": strategy,
10947 "agents": agents_json,
10948 "duration_ms": duration_ms,
10949 }),
10950 );
10951
10952 self.commit_root_user_message(input).await?;
10953
10954 let post_result = self
10955 .post_loop_processing(input, result.response.content.clone())
10956 .await?;
10957 let final_content = self
10958 .apply_post_loop_result(input, post_result)
10959 .await?
10960 .content;
10961
10962 let mut response = AgentResponse::new(final_content);
10963 let metadata = serde_json::json!({
10964 "orchestration": {
10965 "type": "concurrent",
10966 "result": result.response.content,
10967 "strategy": strategy,
10968 "agents": agents_json,
10969 "duration_ms": duration_ms,
10970 }
10971 });
10972 response.metadata = Some(
10973 serde_json::from_value::<std::collections::HashMap<String, serde_json::Value>>(
10974 metadata,
10975 )
10976 .unwrap_or_default(),
10977 );
10978
10979 self.finish_turn_if_root(&response).await?;
10980 Ok(response)
10981 }
10982
10983 async fn handle_group_chat_state(
10985 &self,
10986 input: &str,
10987 config: &ai_agents_state::GroupChatStateConfig,
10988 ) -> Result<AgentResponse> {
10989 use std::time::Instant;
10990
10991 let registry = self.spawner_registry.as_ref().ok_or_else(|| {
10992 AgentError::Config(
10993 "Group chat state requires an agent registry. Add a spawner section.".into(),
10994 )
10995 })?;
10996
10997 let start = Instant::now();
10998
10999 let llm_provider = self.llm_registry.get("router").ok();
11000
11001 let context_mode = config.context_mode.clone().unwrap_or_default();
11003 let context_input = self
11004 .observe_purpose(
11005 ObservationPurpose::OrchestrationRouting,
11006 crate::orchestration::context::prepare_delegate_input(
11007 input,
11008 &context_mode,
11009 &*self.memory,
11010 self.llm_registry.get("router").ok().as_deref(),
11011 ),
11012 )
11013 .await?;
11014
11015 let effective_topic = if let Some(ref tmpl) = config.input {
11017 render_concurrent_template(tmpl, &context_input, &self.build_context_with_overlays())
11018 .unwrap_or_else(|_| context_input.clone())
11019 } else {
11020 context_input
11021 };
11022
11023 let result = self
11024 .observe_purpose(
11025 ObservationPurpose::OrchestrationConversation,
11026 scope_actor_context(
11027 self.outbound_actor_context(),
11028 crate::orchestration::group_chat(
11029 registry,
11030 &effective_topic,
11031 config,
11032 llm_provider.as_deref(),
11033 Some(&*self.hooks),
11034 ),
11035 ),
11036 )
11037 .await?;
11038
11039 let duration_ms = start.elapsed().as_millis() as u64;
11040
11041 let _ = self.context_manager.set(
11043 "group_chat.conclusion",
11044 serde_json::Value::String(result.response.content.clone()),
11045 );
11046
11047 let transcript_json: Vec<serde_json::Value> = result
11049 .transcript
11050 .iter()
11051 .map(|t| {
11052 serde_json::json!({
11053 "speaker": t.speaker,
11054 "round": t.round,
11055 "content": t.content,
11056 })
11057 })
11058 .collect();
11059
11060 let _ = self.context_manager.set(
11062 "orchestration",
11063 serde_json::json!({
11064 "type": "group_chat",
11065 "conclusion": result.response.content,
11066 "transcript": transcript_json,
11067 "rounds": result.rounds_completed,
11068 "termination": result.termination_reason,
11069 "duration_ms": duration_ms,
11070 }),
11071 );
11072
11073 self.commit_root_user_message(input).await?;
11074
11075 let post_result = self
11076 .post_loop_processing(input, result.response.content.clone())
11077 .await?;
11078 let final_content = self
11079 .apply_post_loop_result(input, post_result)
11080 .await?
11081 .content;
11082
11083 let mut response = AgentResponse::new(final_content);
11084 let metadata = serde_json::json!({
11085 "orchestration": {
11086 "type": "group_chat",
11087 "conclusion": result.response.content,
11088 "transcript": transcript_json,
11089 "rounds": result.rounds_completed,
11090 "termination": result.termination_reason,
11091 "duration_ms": duration_ms,
11092 }
11093 });
11094 response.metadata = Some(
11095 serde_json::from_value::<std::collections::HashMap<String, serde_json::Value>>(
11096 metadata,
11097 )
11098 .unwrap_or_default(),
11099 );
11100
11101 self.finish_turn_if_root(&response).await?;
11102 Ok(response)
11103 }
11104
11105 async fn handle_pipeline_state(
11107 &self,
11108 input: &str,
11109 config: &ai_agents_state::PipelineStateConfig,
11110 ) -> Result<AgentResponse> {
11111 use std::time::Instant;
11112
11113 let registry = self.spawner_registry.as_ref().ok_or_else(|| {
11114 AgentError::Config(
11115 "Pipeline state requires an agent registry. Add a spawner section.".into(),
11116 )
11117 })?;
11118
11119 let start = Instant::now();
11120
11121 let stages: Vec<crate::orchestration::PipelineStage> = config
11122 .stages
11123 .iter()
11124 .map(|entry| {
11125 let mut stage = crate::orchestration::PipelineStage::id(entry.id());
11126 if let Some(tmpl) = entry.input() {
11127 stage = stage.with_input(tmpl);
11128 }
11129 stage
11130 })
11131 .collect();
11132
11133 let context_mode = config.context_mode.clone().unwrap_or_default();
11135 let context_input = self
11136 .observe_purpose(
11137 ObservationPurpose::OrchestrationRouting,
11138 crate::orchestration::context::prepare_delegate_input(
11139 input,
11140 &context_mode,
11141 &*self.memory,
11142 self.llm_registry.get("router").ok().as_deref(),
11143 ),
11144 )
11145 .await?;
11146
11147 let context_values = self.build_context_with_overlays();
11148 let result = self
11149 .observe_purpose(
11150 ObservationPurpose::OrchestrationRouting,
11151 scope_actor_context(
11152 self.outbound_actor_context(),
11153 crate::orchestration::pipeline(
11154 registry,
11155 &context_input,
11156 &stages,
11157 config.timeout_ms,
11158 Some(&*self.hooks),
11159 Some(&context_values),
11160 ),
11161 ),
11162 )
11163 .await?;
11164
11165 let duration_ms = start.elapsed().as_millis() as u64;
11166
11167 let _ = self.context_manager.set(
11169 "pipeline.result",
11170 serde_json::Value::String(result.response.content.clone()),
11171 );
11172
11173 let stages_json: Vec<serde_json::Value> = result
11175 .stage_outputs
11176 .iter()
11177 .map(|s| {
11178 serde_json::json!({
11179 "agent_id": s.agent_id,
11180 "output": s.output,
11181 "duration_ms": s.duration_ms,
11182 "skipped": s.skipped,
11183 })
11184 })
11185 .collect();
11186
11187 let _ = self.context_manager.set(
11189 "orchestration",
11190 serde_json::json!({
11191 "type": "pipeline",
11192 "result": result.response.content,
11193 "stages": stages_json,
11194 "duration_ms": duration_ms,
11195 }),
11196 );
11197
11198 self.commit_root_user_message(input).await?;
11199
11200 let post_result = self
11201 .post_loop_processing(input, result.response.content.clone())
11202 .await?;
11203 let final_content = self
11204 .apply_post_loop_result(input, post_result)
11205 .await?
11206 .content;
11207
11208 let mut response = AgentResponse::new(final_content);
11209 let metadata = serde_json::json!({
11210 "orchestration": {
11211 "type": "pipeline",
11212 "result": result.response.content,
11213 "stages": stages_json,
11214 "duration_ms": duration_ms,
11215 }
11216 });
11217 response.metadata = Some(
11218 serde_json::from_value::<std::collections::HashMap<String, serde_json::Value>>(
11219 metadata,
11220 )
11221 .unwrap_or_default(),
11222 );
11223
11224 self.finish_turn_if_root(&response).await?;
11225 Ok(response)
11226 }
11227
11228 async fn handle_handoff_state(
11230 &self,
11231 input: &str,
11232 config: &ai_agents_state::HandoffStateConfig,
11233 ) -> Result<AgentResponse> {
11234 use std::time::Instant;
11235
11236 let registry = self.spawner_registry.as_ref().ok_or_else(|| {
11237 AgentError::Config(
11238 "Handoff state requires an agent registry. Add a spawner section.".into(),
11239 )
11240 })?;
11241
11242 let llm = self
11243 .llm_registry
11244 .get("router")
11245 .map_err(|_| AgentError::Config("Handoff state requires a router LLM.".into()))?;
11246
11247 let start = Instant::now();
11248
11249 let context_mode = config.context_mode.clone().unwrap_or_default();
11251 let context_input = self
11252 .observe_purpose(
11253 ObservationPurpose::OrchestrationRouting,
11254 crate::orchestration::context::prepare_delegate_input(
11255 input,
11256 &context_mode,
11257 &*self.memory,
11258 self.llm_registry.get("router").ok().as_deref(),
11259 ),
11260 )
11261 .await?;
11262
11263 let effective_input = if let Some(ref tmpl) = config.input {
11265 render_concurrent_template(tmpl, &context_input, &self.build_context_with_overlays())
11266 .unwrap_or_else(|_| context_input.clone())
11267 } else {
11268 context_input
11269 };
11270
11271 let result = self
11272 .observe_purpose(
11273 ObservationPurpose::OrchestrationRouting,
11274 scope_actor_context(
11275 self.outbound_actor_context(),
11276 crate::orchestration::handoff(
11277 registry,
11278 &effective_input,
11279 &config.initial_agent,
11280 &config.available_agents,
11281 config.max_handoffs,
11282 llm.as_ref(),
11283 Some(&*self.hooks),
11284 ),
11285 ),
11286 )
11287 .await?;
11288
11289 let duration_ms = start.elapsed().as_millis() as u64;
11290
11291 let _ = self.context_manager.set(
11293 "handoff.result",
11294 serde_json::Value::String(result.response.content.clone()),
11295 );
11296
11297 let chain_json: Vec<serde_json::Value> = result
11299 .handoff_chain
11300 .iter()
11301 .map(|h| {
11302 serde_json::json!({
11303 "from": h.from_agent,
11304 "to": h.to_agent,
11305 "reason": h.reason,
11306 })
11307 })
11308 .collect();
11309
11310 let _ = self.context_manager.set(
11312 "orchestration",
11313 serde_json::json!({
11314 "type": "handoff",
11315 "result": result.response.content,
11316 "final_agent": result.final_agent,
11317 "handoff_chain": chain_json,
11318 "duration_ms": duration_ms,
11319 }),
11320 );
11321
11322 self.commit_root_user_message(input).await?;
11323
11324 let post_result = self
11325 .post_loop_processing(input, result.response.content.clone())
11326 .await?;
11327 let final_content = self
11328 .apply_post_loop_result(input, post_result)
11329 .await?
11330 .content;
11331
11332 let mut response = AgentResponse::new(final_content);
11333 let metadata = serde_json::json!({
11334 "orchestration": {
11335 "type": "handoff",
11336 "result": result.response.content,
11337 "final_agent": result.final_agent,
11338 "handoff_chain": chain_json,
11339 "duration_ms": duration_ms,
11340 }
11341 });
11342 response.metadata = Some(
11343 serde_json::from_value::<std::collections::HashMap<String, serde_json::Value>>(
11344 metadata,
11345 )
11346 .unwrap_or_default(),
11347 );
11348
11349 self.finish_turn_if_root(&response).await?;
11350 Ok(response)
11351 }
11352
11353 async fn run_loop_internal(&self, input: &str) -> Result<AgentResponse> {
11355 self.begin_root_turn();
11356 self.pre_turn_session_lifecycle().await;
11358
11359 let input_data = self.process_input(input).await?;
11360 self.update_active_turn_context(&input_data.content, input_data.context.clone());
11361
11362 for (key, value) in &input_data.context {
11365 let _ = self.context_manager.set(key, value.clone());
11366 }
11367
11368 if input_data.metadata.rejected {
11369 let reason = input_data
11370 .metadata
11371 .rejection_reason
11372 .unwrap_or_else(|| "Input rejected".to_string());
11373 warn!(reason = %reason, "Input rejected");
11374 let response = AgentResponse::new(reason);
11375 self.finish_turn_if_root(&response).await?;
11376 return Ok(response);
11377 }
11378
11379 let processed_input = &input_data.content;
11380
11381 if let Some(response) = self.try_pre_response_transition(processed_input).await? {
11382 return Ok(response);
11383 }
11384
11385 if let Some(ref sm) = self.state_machine
11387 && let Some(def) = sm.current_definition()
11388 {
11389 if let Some(ref delegate_id) = def.delegate {
11390 return self
11391 .handle_delegated_state(processed_input, delegate_id, &def)
11392 .await;
11393 }
11394 if let Some(ref concurrent_config) = def.concurrent {
11395 return self
11396 .handle_concurrent_state(processed_input, concurrent_config)
11397 .await;
11398 }
11399 if let Some(ref group_chat_config) = def.group_chat {
11400 return self
11401 .handle_group_chat_state(processed_input, group_chat_config)
11402 .await;
11403 }
11404 if let Some(ref pipeline_config) = def.pipeline {
11405 return self
11406 .handle_pipeline_state(processed_input, pipeline_config)
11407 .await;
11408 }
11409 if let Some(ref handoff_config) = def.handoff {
11410 return self
11411 .handle_handoff_state(processed_input, handoff_config)
11412 .await;
11413 }
11414 }
11415
11416 if let Some(response) =
11421 Box::pin(self.try_speculative_branches(processed_input, &input_data.context)).await?
11422 {
11423 return Ok(response);
11424 }
11425
11426 match self.try_skill_route(processed_input).await? {
11427 SkillRouteResult::Response { skill_id, content } => {
11428 self.commit_root_user_message(processed_input).await?;
11429 return self
11430 .handle_skill_response(processed_input, &skill_id, content, &input_data.context)
11431 .await;
11432 }
11433 SkillRouteResult::NeedsClarification {
11434 response,
11435 ownership,
11436 } => {
11437 let admission = self
11438 .admit_optional_disambiguation_ownership(ownership)
11439 .await?;
11440 self.commit_root_user_message(processed_input).await?;
11441 if Self::skill_clarification_needs_memory_record(&response) {
11442 self.memory
11445 .add_message(ChatMessage::assistant(&response.content))
11446 .await?;
11447 }
11448 drop(admission);
11449 self.finish_turn_if_root(&response).await?;
11450 return Ok(response);
11451 }
11452 SkillRouteResult::NoMatch => {} }
11454
11455 let effective_reasoning = self.get_effective_reasoning_config();
11456 let reasoning_mode = self.determine_reasoning_mode(processed_input).await?;
11457 let auto_detected = matches!(effective_reasoning.mode, ReasoningMode::Auto);
11458
11459 info!(
11460 reasoning_mode = ?reasoning_mode,
11461 auto_detected = auto_detected,
11462 reflection_enabled = ?self.reflection_config.enabled,
11463 "Reasoning mode determined"
11464 );
11465
11466 if matches!(reasoning_mode, ReasoningMode::PlanAndExecute) {
11467 self.commit_root_user_message(processed_input).await?;
11468 return self
11469 .handle_plan_and_execute(processed_input, &input_data.context, auto_detected)
11470 .await;
11471 }
11472
11473 self.commit_root_user_message(processed_input).await?;
11474
11475 let mut iterations = 0u32;
11476 let mut all_tool_calls: Vec<ToolCall> = Vec::new();
11477 let mut thinking_content: Option<String> = None;
11478
11479 let llm = self.get_state_llm()?;
11480
11481 loop {
11482 let effective_max = if reasoning_mode != ReasoningMode::None {
11484 let rc = self.get_effective_reasoning_config();
11485 self.max_iterations.min(rc.max_iterations)
11486 } else {
11487 self.max_iterations
11488 };
11489
11490 if iterations >= effective_max {
11491 let err = AgentError::Other(format!("Max iterations ({}) exceeded", effective_max));
11492 self.hooks.on_error(&err).await;
11493 error!(iterations = iterations, "Max iterations exceeded");
11494 return Err(err);
11495 }
11496 iterations += 1;
11497 *self.iteration_count.write() = iterations;
11498
11499 debug!(iteration = iterations, max = effective_max, "LLM call");
11500
11501 let protocol = self.main_tool_protocol(llm.as_ref(), false).await?;
11502 let mut messages = self
11503 .build_messages_internal(true, None, protocol.choice.is_none())
11504 .await?;
11505 self.inject_reasoning_prompt(&mut messages, &reasoning_mode, iterations == 1);
11506
11507 self.hooks.on_llm_start(&messages).await;
11508 let llm_start = Instant::now();
11509 let response = self
11510 .complete_main_llm_with_recovery(Arc::clone(&llm), &messages, &protocol)
11511 .await?;
11512
11513 let llm_duration_ms = llm_start.elapsed().as_millis() as u64;
11514 self.hooks.on_llm_complete(&response, llm_duration_ms).await;
11515
11516 let content = response.content.trim();
11517
11518 if let Some(tool_calls) = self.parse_main_tool_calls(content, &protocol)? {
11519 match self
11520 .handle_tool_calls(
11521 processed_input,
11522 content,
11523 tool_calls,
11524 &mut all_tool_calls,
11525 None,
11526 )
11527 .await?
11528 {
11529 ToolCallOutcome::Continue | ToolCallOutcome::TransitionFired => continue,
11530 ToolCallOutcome::Rejected(resp) => {
11531 self.finish_turn_if_root(&resp).await?;
11532 return Ok(resp);
11533 }
11534 }
11535 }
11536
11537 let (extracted_thinking, answer) = self.extract_thinking(content);
11538 if extracted_thinking.is_some() {
11539 thinking_content = extracted_thinking;
11540 }
11541
11542 let output_data = self.process_output(&answer, &input_data.context).await?;
11543
11544 let mut final_content = if output_data.metadata.rejected {
11545 output_data
11546 .metadata
11547 .rejection_reason
11548 .unwrap_or_else(|| answer.to_string())
11549 } else {
11550 output_data.content
11551 };
11552
11553 let reflection_metadata;
11555 (final_content, reflection_metadata) = self
11556 .run_reflection(&*llm, processed_input, final_content)
11557 .await?;
11558
11559 final_content =
11560 self.format_response_with_thinking(thinking_content.as_deref(), &final_content);
11561
11562 let final_content = {
11566 let result = self
11567 .post_loop_processing(processed_input, final_content)
11568 .await?;
11569 self.apply_post_loop_result(processed_input, result)
11570 .await?
11571 .content
11572 };
11573
11574 let reflected = reflection_metadata.is_some();
11575 let reasoning_mode_debug = format!("{:?}", reasoning_mode);
11576
11577 let response = self.build_agent_response(AgentResponseParts {
11578 content: final_content,
11579 all_tool_calls,
11580 reasoning_mode,
11581 auto_detected,
11582 iterations,
11583 thinking: thinking_content,
11584 reflection_metadata,
11585 });
11586
11587 self.finish_turn_if_root(&response).await?;
11588
11589 let tool_call_count = response.tool_calls.as_ref().map(|tc| tc.len()).unwrap_or(0);
11590 info!(
11591 tool_calls = tool_call_count,
11592 response_len = response.content.len(),
11593 reasoning_mode = %reasoning_mode_debug,
11594 reflected = reflected,
11595 "Chat completed"
11596 );
11597 return Ok(response);
11598 }
11599 }
11600
11601 async fn generate_buffered_streaming_draft(
11602 &self,
11603 processed_input: &str,
11604 routing_resolved: Arc<AtomicBool>,
11605 ) -> Result<StreamingDraftResult> {
11606 let llm = self.get_state_llm()?;
11607 if llm.configured_tool_choice().is_some() {
11608 let draft = self
11609 .generate_main_response_draft(processed_input, &ReasoningMode::None)
11610 .await?;
11611 return Ok(StreamingDraftResult::new(draft, Vec::new()));
11612 }
11613 let protocol = self.main_tool_protocol(llm.as_ref(), true).await?;
11615 let messages = self.build_messages_for_draft(processed_input).await?;
11616 let source = self
11617 .open_main_stream_with_recovery(Arc::clone(&llm), &messages, &protocol)
11618 .await?;
11619 let mut buffer = crate::optimization::StreamBranchBuffer::new(self.streaming.buffer_size)?;
11620 let mut chunks = Vec::new();
11621 let mut accumulated = String::new();
11622 match source {
11623 MainStreamSource::StaticResponse(text) => {
11624 accumulated.push_str(&text);
11625 let stream_chunk = StreamChunk::content(text);
11626 if routing_resolved.load(Ordering::SeqCst) {
11627 chunks.push(stream_chunk);
11628 } else {
11629 buffer.push(stream_chunk)?;
11630 }
11631 }
11632 MainStreamSource::Stream(mut stream) => {
11633 while let Some(chunk_result) = stream.next().await {
11634 let chunk = chunk_result.map_err(|e| AgentError::LLM(e.to_string()))?;
11635 accumulated.push_str(&chunk.delta);
11636 let stream_chunk = StreamChunk::content(chunk.delta);
11637 if routing_resolved.load(Ordering::SeqCst) {
11638 chunks.push(stream_chunk);
11639 } else {
11640 buffer.push(stream_chunk)?;
11641 }
11642 }
11643 }
11644 }
11645 chunks.splice(0..0, buffer.drain());
11646 let content = accumulated.trim().to_string();
11647 let draft = if let Some(calls) = self.parse_tool_calls(&content)? {
11648 MainResponseDraft::ToolCalls {
11649 raw_content: content,
11650 calls,
11651 thinking: None,
11652 }
11653 } else {
11654 MainResponseDraft::Text {
11655 raw_content: content,
11656 thinking: None,
11657 }
11658 };
11659 Ok(StreamingDraftResult::new(draft, chunks))
11660 }
11661
11662 async fn try_buffered_streaming_branches(
11663 &self,
11664 processed_input: &str,
11665 input_context: &HashMap<String, Value>,
11666 ) -> Result<Option<(AgentResponse, Vec<StreamChunk>)>> {
11667 let optimization = &self.runtime_config.optimization;
11668 if !optimization.enabled {
11669 return Ok(None);
11670 }
11671 if !matches!(
11677 self.get_effective_reasoning_config().mode,
11678 ReasoningMode::None
11679 ) {
11680 return Ok(None);
11681 }
11682 let transition_enabled =
11683 optimization.speculative_state_transitions && self.has_parallel_transition_candidates();
11684 if !transition_enabled {
11685 return Ok(None);
11686 }
11687 let mut branch_scheduler =
11688 TurnBranchScheduler::new(optimization.max_parallel_runtime_tasks)?;
11689 if !branch_scheduler.reserve_task() {
11690 return Ok(None);
11691 }
11692 if !self
11693 .reserve_active_speculative_llm_call(RuntimeOptimizationKind::BufferedStreamingRouting)
11694 {
11695 branch_scheduler.release_task();
11696 return Ok(None);
11697 }
11698 if !branch_scheduler.reserve_task() {
11699 branch_scheduler.release_task();
11700 return Ok(None);
11701 }
11702 let mut main_branch = RuntimeBranch::new(
11703 RuntimeTaskPurpose::MainResponse,
11704 RuntimeOptimizationKind::BufferedStreamingRouting,
11705 RuntimeTaskPriority::Normal,
11706 RuntimeCommitBehavior::FinalResponse,
11707 );
11708 let mut transition_branch = RuntimeBranch::new(
11709 RuntimeTaskPurpose::StateTransition,
11710 RuntimeOptimizationKind::ParallelStateTransition,
11711 RuntimeTaskPriority::Critical,
11712 RuntimeCommitBehavior::TransitionDecision,
11713 );
11714 let main_id = main_branch.branch_id();
11715 let transition_id = transition_branch.branch_id();
11716 let routing_resolved = Arc::new(AtomicBool::new(false));
11717 let mut main_future =
11718 Box::pin(crate::optimization::observability::with_branch_observation(
11719 &main_id,
11720 RuntimeOptimizationKind::BufferedStreamingRouting,
11721 RuntimeCommitBehavior::FinalResponse,
11722 self.generate_buffered_streaming_draft(
11723 processed_input,
11724 Arc::clone(&routing_resolved),
11725 ),
11726 ));
11727 let mut transition_future =
11728 Box::pin(crate::optimization::observability::with_branch_observation(
11729 &transition_id,
11730 RuntimeOptimizationKind::ParallelStateTransition,
11731 RuntimeCommitBehavior::TransitionDecision,
11732 self.select_parallel_transition_candidate(processed_input),
11733 ));
11734 let mut main_pending = true;
11735 let mut transition_pending = true;
11736 let mut main_result: Option<Result<StreamingDraftResult>> = None;
11737 let mut transition_finalized = false;
11738 let mut transition_candidate: Option<TransitionCandidate> = None;
11739 loop {
11740 if let Some(candidate) = transition_candidate.take() {
11741 if self
11742 .approve_transition_target(&candidate.from_state, candidate.target())
11743 .await?
11744 {
11745 drop(main_future);
11747 drop(transition_future);
11748 self.finalize_branch_loss(
11749 &main_id,
11750 RuntimeOptimizationKind::BufferedStreamingRouting,
11751 RuntimeCommitBehavior::FinalResponse,
11752 main_pending,
11753 main_result.as_ref().map(|result| result.is_err()),
11754 );
11755 if !self
11756 .apply_pre_response_transition_candidate(
11757 &candidate,
11758 &HashMap::new(),
11759 processed_input,
11760 )
11761 .await?
11762 {
11763 self.finalize_optional_branch(
11764 &transition_id,
11765 RuntimeOptimizationKind::ParallelStateTransition,
11766 RuntimeCommitBehavior::TransitionDecision,
11767 "discarded",
11768 false,
11769 );
11770 return Ok(None);
11771 }
11772 self.finalize_optional_branch(
11773 &transition_id,
11774 RuntimeOptimizationKind::ParallelStateTransition,
11775 RuntimeCommitBehavior::TransitionDecision,
11776 "committed",
11777 true,
11778 );
11779 let response = self.redispatch_current_state(processed_input).await?;
11780 return Ok(Some((
11781 response.clone(),
11782 vec![StreamChunk::content(response.content)],
11783 )));
11784 }
11785 self.finalize_optional_branch(
11786 &transition_id,
11787 RuntimeOptimizationKind::ParallelStateTransition,
11788 RuntimeCommitBehavior::TransitionDecision,
11789 "discarded",
11790 false,
11791 );
11792 transition_finalized = true;
11793 }
11794 if transition_finalized && !routing_resolved.load(Ordering::SeqCst) {
11800 match self
11801 .resolve_buffered_skill_after_transition(processed_input, &routing_resolved)
11802 .await
11803 {
11804 Ok(Some(candidate)) => {
11805 drop(main_future);
11807 drop(transition_future);
11808 self.finalize_branch_loss(
11809 &main_id,
11810 RuntimeOptimizationKind::BufferedStreamingRouting,
11811 RuntimeCommitBehavior::FinalResponse,
11812 main_pending,
11813 main_result.as_ref().map(|result| result.is_err()),
11814 );
11815 return match self
11816 .commit_winning_skill_candidate(
11817 candidate,
11818 processed_input,
11819 input_context,
11820 )
11821 .await?
11822 {
11823 Some(response) => Ok(Some((
11824 response.clone(),
11825 vec![StreamChunk::content(response.content)],
11826 ))),
11827 None => Ok(None),
11828 };
11829 }
11830 Ok(None) => {}
11831 Err(error) => {
11832 drop(main_future);
11833 drop(transition_future);
11834 self.finalize_branch_loss(
11835 &main_id,
11836 RuntimeOptimizationKind::BufferedStreamingRouting,
11837 RuntimeCommitBehavior::FinalResponse,
11838 main_pending,
11839 main_result.as_ref().map(|result| result.is_err()),
11840 );
11841 return Err(error);
11842 }
11843 }
11844 }
11845 if transition_finalized
11846 && routing_resolved.load(Ordering::SeqCst)
11847 && let Some(result) = main_result.take()
11848 {
11849 let stream_draft = match result {
11850 Ok(stream_draft) => stream_draft,
11851 Err(error) => {
11852 self.finalize_optional_branch(
11853 &main_id,
11854 RuntimeOptimizationKind::BufferedStreamingRouting,
11855 RuntimeCommitBehavior::FinalResponse,
11856 "failed",
11857 false,
11858 );
11859 return Err(error);
11860 }
11861 };
11862 let raw_draft_content = stream_draft.draft.raw_content().to_string();
11863 let buffered_chunks = stream_draft.chunks;
11864 self.finalize_optional_branch(
11865 &main_id,
11866 RuntimeOptimizationKind::BufferedStreamingRouting,
11867 RuntimeCommitBehavior::FinalResponse,
11868 "committed",
11869 true,
11870 );
11871 let response = self
11872 .commit_main_response_draft(
11873 processed_input,
11874 input_context,
11875 stream_draft.draft,
11876 ReasoningMode::None,
11877 false,
11878 )
11879 .await?;
11880 let chunks = if response.content == raw_draft_content {
11881 buffered_chunks
11882 } else {
11883 vec![StreamChunk::content(response.content.clone())]
11884 };
11885 return Ok(Some((response, chunks)));
11886 }
11887 tokio::select! {
11888 result = &mut main_future, if main_pending => {
11889 main_pending = false;
11890 main_branch.transition_to(RuntimeBranchStatus::Completed)?;
11891 main_result = Some(result);
11892 }
11893 result = &mut transition_future, if transition_pending => {
11894 transition_pending = false;
11895 transition_branch.transition_to(RuntimeBranchStatus::Completed)?;
11896 match result {
11897 Ok(ParallelTransitionSelection::Candidate(candidate)) => {
11898 transition_candidate = Some(candidate)
11899 }
11900 Ok(ParallelTransitionSelection::NoMatch) => {
11901 self.finalize_optional_branch(
11902 &transition_id,
11903 RuntimeOptimizationKind::ParallelStateTransition,
11904 RuntimeCommitBehavior::TransitionDecision,
11905 "discarded",
11906 false,
11907 );
11908 transition_finalized = true;
11909 }
11910 Ok(ParallelTransitionSelection::ReservationExhausted) => {
11911 self.finalize_optional_branch(
11912 &transition_id,
11913 RuntimeOptimizationKind::ParallelStateTransition,
11914 RuntimeCommitBehavior::TransitionDecision,
11915 "cancelled",
11916 false,
11917 );
11918 routing_resolved.store(true, Ordering::SeqCst);
11919 self.finalize_branch_loss(
11920 &main_id,
11921 RuntimeOptimizationKind::BufferedStreamingRouting,
11922 RuntimeCommitBehavior::FinalResponse,
11923 main_pending,
11924 main_result.as_ref().map(|result| result.is_err()),
11925 );
11926 return Ok(None);
11927 }
11928 Err(_) => {
11929 self.finalize_optional_branch(
11930 &transition_id,
11931 RuntimeOptimizationKind::ParallelStateTransition,
11932 RuntimeCommitBehavior::TransitionDecision,
11933 "failed",
11934 false,
11935 );
11936 transition_finalized = true;
11937 }
11938 }
11939 }
11940 }
11941 }
11942 }
11943
11944 async fn resolve_buffered_skill_after_transition(
11950 &self,
11951 processed_input: &str,
11952 routing_resolved: &AtomicBool,
11953 ) -> Result<Option<SkillCandidate>> {
11954 let candidate = if self.skill_router.is_some() {
11955 self.select_skill_candidate(processed_input).await?
11956 } else {
11957 None
11958 };
11959 if candidate.is_none() {
11960 routing_resolved.store(true, Ordering::SeqCst);
11961 }
11962 Ok(candidate)
11963 }
11964
11965 fn run_loop_internal_stream<'a>(
11969 &'a self,
11970 input: &'a str,
11971 terminal: RuntimeStreamTerminalSlot,
11972 ) -> Pin<Box<dyn Stream<Item = StreamChunk> + Send + 'a>> {
11973 let include_state_events = self.streaming.include_state_events;
11974
11975 Box::pin(async_stream::stream! {
11976 self.begin_root_turn();
11977 self.pre_turn_session_lifecycle().await;
11979
11980 let input_data = match self.process_input(input).await {
11981 Ok(data) => data,
11982 Err(e) => {
11983 yield StreamChunk::error(e.to_string());
11984 return;
11985 }
11986 };
11987 self.update_active_turn_context(&input_data.content, input_data.context.clone());
11988
11989 for (key, value) in &input_data.context {
11991 let _ = self.context_manager.set(key, value.clone());
11992 }
11993
11994 if input_data.metadata.rejected {
11995 let reason = input_data
11996 .metadata
11997 .rejection_reason
11998 .unwrap_or_else(|| "Input rejected".to_string());
11999 warn!(reason = %reason, "Input rejected (stream)");
12000 let response = AgentResponse::new(&reason);
12003 if let Err(e) = self.finish_turn_if_root(&response).await {
12004 yield StreamChunk::error(e.to_string());
12005 return;
12006 }
12007 yield StreamChunk::content(&reason);
12008 record_runtime_stream_final(&terminal, response);
12009 yield StreamChunk::Done {};
12010 return;
12011 }
12012
12013 let processed_input = &input_data.content;
12014
12015 let streaming_policy = self.runtime_config.optimization.streaming_policy;
12016
12017 if self.runtime_config.optimization.enabled
12024 && !matches!(
12025 streaming_policy,
12026 crate::optimization::StreamingOptimizationPolicy::Disabled
12027 )
12028 {
12029 match self.try_pre_response_transition(processed_input).await {
12030 Ok(Some(response)) => {
12031 yield StreamChunk::content(&response.content);
12032 record_runtime_stream_final(&terminal, response);
12033 yield StreamChunk::Done {};
12034 return;
12035 }
12036 Ok(None) => {}
12037 Err(e) => {
12038 yield StreamChunk::error(e.to_string());
12039 return;
12040 }
12041 }
12042 }
12043
12044 if self.runtime_config.optimization.enabled
12045 && matches!(
12046 streaming_policy,
12047 crate::optimization::StreamingOptimizationPolicy::BufferUntilRoutingDone
12048 )
12049 {
12050 match Box::pin(self.try_buffered_streaming_branches(processed_input, &input_data.context)).await {
12055 Ok(Some((response, chunks))) => {
12056 for chunk in chunks {
12057 yield chunk;
12058 }
12059 record_runtime_stream_final(&terminal, response);
12060 yield StreamChunk::Done {};
12061 return;
12062 }
12063 Ok(None) => {}
12064 Err(e) => {
12065 yield StreamChunk::error(e.to_string());
12066 return;
12067 }
12068 }
12069 }
12070
12071 if let Some(ref sm) = self.state_machine
12073 && let Some(def) = sm.current_definition()
12074 {
12075 let orchestration_result = if let Some(ref delegate_id) = def.delegate {
12076 Some(self.handle_delegated_state(processed_input, delegate_id, &def).await)
12077 } else if let Some(ref concurrent_config) = def.concurrent {
12078 Some(self.handle_concurrent_state(processed_input, concurrent_config).await)
12079 } else if let Some(ref group_chat_config) = def.group_chat {
12080 Some(self.handle_group_chat_state(processed_input, group_chat_config).await)
12081 } else if let Some(ref pipeline_config) = def.pipeline {
12082 Some(self.handle_pipeline_state(processed_input, pipeline_config).await)
12083 } else if let Some(ref handoff_config) = def.handoff {
12084 Some(self.handle_handoff_state(processed_input, handoff_config).await)
12085 } else {
12086 None
12087 };
12088
12089 if let Some(result) = orchestration_result {
12090 match result {
12091 Ok(response) => {
12092 yield StreamChunk::content(&response.content);
12093 record_runtime_stream_final(&terminal, response);
12094 yield StreamChunk::Done {};
12095 }
12096 Err(e) => {
12097 yield StreamChunk::error(e.to_string());
12098 }
12099 }
12100 return;
12101 }
12102 }
12103
12104 match self.try_skill_route(processed_input).await {
12106 Ok(SkillRouteResult::Response { skill_id, content }) => {
12107 if let Err(e) = self.commit_root_user_message(processed_input).await {
12108 yield StreamChunk::error(e.to_string());
12109 return;
12110 }
12111 match self.handle_skill_response(processed_input, &skill_id, content, &input_data.context).await {
12112 Ok(resp) => {
12113 yield StreamChunk::content(&resp.content);
12114 record_runtime_stream_final(&terminal, resp);
12115 yield StreamChunk::Done {};
12116 return;
12117 }
12118 Err(e) => {
12119 yield StreamChunk::error(e.to_string());
12120 return;
12121 }
12122 }
12123 }
12124 Ok(SkillRouteResult::NeedsClarification {
12125 response,
12126 ownership,
12127 }) => {
12128 let admission = match self
12129 .admit_optional_disambiguation_ownership(ownership)
12130 .await
12131 {
12132 Ok(admission) => admission,
12133 Err(e) => {
12134 yield StreamChunk::error(e.to_string());
12135 return;
12136 }
12137 };
12138 if let Err(e) = self.commit_root_user_message(processed_input).await {
12139 yield StreamChunk::error(e.to_string());
12140 return;
12141 }
12142 if Self::skill_clarification_needs_memory_record(&response)
12144 && let Err(e) = self.memory.add_message(ChatMessage::assistant(&response.content)).await
12145 {
12146 yield StreamChunk::error(e.to_string());
12147 return;
12148 }
12149 drop(admission);
12150 if let Err(e) = self.finish_turn_if_root(&response).await {
12151 yield StreamChunk::error(e.to_string());
12152 return;
12153 }
12154 yield StreamChunk::content(&response.content);
12155 record_runtime_stream_final(&terminal, response);
12156 yield StreamChunk::Done {};
12157 return;
12158 }
12159 Ok(SkillRouteResult::NoMatch) => {} Err(e) => {
12161 yield StreamChunk::error(e.to_string());
12162 return;
12163 }
12164 }
12165
12166 let effective_reasoning = self.get_effective_reasoning_config();
12168 let reasoning_mode = match self.determine_reasoning_mode(processed_input).await {
12169 Ok(mode) => mode,
12170 Err(e) => {
12171 yield StreamChunk::error(e.to_string());
12172 return;
12173 }
12174 };
12175 let auto_detected = matches!(effective_reasoning.mode, ReasoningMode::Auto);
12176
12177 info!(
12178 reasoning_mode = ?reasoning_mode,
12179 auto_detected = auto_detected,
12180 "Reasoning mode determined (stream)"
12181 );
12182
12183 if matches!(reasoning_mode, ReasoningMode::PlanAndExecute) {
12185 if let Err(e) = self.commit_root_user_message(processed_input).await {
12186 yield StreamChunk::error(e.to_string());
12187 return;
12188 }
12189 match self.handle_plan_and_execute(processed_input, &input_data.context, auto_detected).await {
12190 Ok(resp) => {
12191 yield StreamChunk::content(&resp.content);
12192 record_runtime_stream_final(&terminal, resp);
12193 yield StreamChunk::Done {};
12194 return;
12195 }
12196 Err(e) => {
12197 yield StreamChunk::error(e.to_string());
12198 return;
12199 }
12200 }
12201 }
12202
12203 if let Err(e) = self.commit_root_user_message(processed_input).await {
12204 yield StreamChunk::error(e.to_string());
12205 return;
12206 }
12207
12208 let llm = match self.get_state_llm() {
12209 Ok(llm) => llm,
12210 Err(e) => {
12211 yield StreamChunk::error(e.to_string());
12212 return;
12213 }
12214 };
12215
12216 let mut iterations = 0u32;
12217 let mut all_tool_calls: Vec<ToolCall> = Vec::new();
12218 let mut thinking_content: Option<String> = None;
12219
12220 loop {
12221 let effective_max = if reasoning_mode != ReasoningMode::None {
12223 let rc = self.get_effective_reasoning_config();
12224 self.max_iterations.min(rc.max_iterations)
12225 } else {
12226 self.max_iterations
12227 };
12228
12229 if iterations >= effective_max {
12230 let err_msg = format!("Max iterations ({}) exceeded", effective_max);
12231 let err = AgentError::Other(err_msg.clone());
12232 self.hooks.on_error(&err).await;
12233 error!(iterations = iterations, "Max iterations exceeded (stream)");
12234 yield StreamChunk::error(err_msg);
12235 return;
12236 }
12237 iterations += 1;
12238 *self.iteration_count.write() = iterations;
12239
12240 debug!(iteration = iterations, max = effective_max, "LLM call (stream)");
12241
12242 let protocol = match self.main_tool_protocol(llm.as_ref(), false).await {
12243 Ok(protocol) => protocol,
12244 Err(e) => {
12245 yield StreamChunk::error(e.to_string());
12246 return;
12247 }
12248 };
12249 let mut messages = match self
12250 .build_messages_internal(true, None, protocol.choice.is_none())
12251 .await
12252 {
12253 Ok(m) => m,
12254 Err(e) => {
12255 yield StreamChunk::error(e.to_string());
12256 return;
12257 }
12258 };
12259 self.inject_reasoning_prompt(&mut messages, &reasoning_mode, iterations == 1);
12260
12261 self.hooks.on_llm_start(&messages).await;
12262 let llm_start = Instant::now();
12263
12264 let buffered_decision = self.main_stream_must_buffer(&reasoning_mode, &protocol);
12265 let content = if buffered_decision {
12266 let response = match self
12270 .complete_main_llm_with_recovery(
12271 Arc::clone(&llm),
12272 &messages,
12273 &protocol,
12274 )
12275 .await
12276 {
12277 Ok(r) => r,
12278 Err(e) => {
12279 yield StreamChunk::error(e.to_string());
12280 return;
12281 }
12282 };
12283 let llm_duration_ms = llm_start.elapsed().as_millis() as u64;
12284 self.hooks.on_llm_complete(&response, llm_duration_ms).await;
12285 response.content.trim().to_string()
12286 } else {
12287 let source = match self
12289 .open_main_stream_with_recovery(Arc::clone(&llm), &messages, &protocol)
12290 .await
12291 {
12292 Ok(source) => source,
12293 Err(e) => {
12294 yield StreamChunk::error(e.to_string());
12295 return;
12296 }
12297 };
12298 let mut accumulated = String::new();
12299 match source {
12300 MainStreamSource::StaticResponse(text) => {
12301 accumulated.push_str(&text);
12302 yield StreamChunk::content(text);
12303 }
12304 MainStreamSource::Stream(mut stream_inner) => {
12305 while let Some(chunk_result) = stream_inner.next().await {
12306 match chunk_result {
12307 Ok(chunk) => {
12308 accumulated.push_str(&chunk.delta);
12309 yield StreamChunk::content(chunk.delta);
12310 }
12311 Err(e) => {
12312 yield StreamChunk::error(e.to_string());
12314 return;
12315 }
12316 }
12317 }
12318 }
12319 }
12320 let llm_duration_ms = llm_start.elapsed().as_millis() as u64;
12321 let llm_response = ai_agents_core::LLMResponse::new(
12323 accumulated.trim(),
12324 ai_agents_core::FinishReason::Stop,
12325 );
12326 self.hooks.on_llm_complete(&llm_response, llm_duration_ms).await;
12327 accumulated.trim().to_string()
12328 };
12329
12330 let parsed_tool_calls = match self.parse_main_tool_calls(&content, &protocol) {
12332 Ok(calls) => calls,
12333 Err(error) => {
12334 yield StreamChunk::error(error.to_string());
12335 return;
12336 }
12337 };
12338 if let Some(tool_calls) = parsed_tool_calls {
12339 let mut events = Vec::new();
12342 let outcome = self
12343 .handle_tool_calls(
12344 processed_input,
12345 &content,
12346 tool_calls,
12347 &mut all_tool_calls,
12348 Some(&mut events),
12349 )
12350 .await;
12351 for chunk in events.drain(..) {
12352 yield chunk;
12353 }
12354 match outcome {
12355 Ok(ToolCallOutcome::Continue) | Ok(ToolCallOutcome::TransitionFired) => continue,
12356 Ok(ToolCallOutcome::Rejected(response)) => {
12357 if let Err(finalize_error) = self.finish_turn_if_root(&response).await {
12358 yield StreamChunk::error(finalize_error.to_string());
12359 return;
12360 }
12361 let legacy_error = response.content.clone();
12362 record_runtime_stream_final(&terminal, response);
12363 yield StreamChunk::error(legacy_error);
12364 yield StreamChunk::Done {};
12365 return;
12366 }
12367 Err(e) => {
12368 yield StreamChunk::error(e.to_string());
12369 return;
12370 }
12371 }
12372 }
12373
12374 let (extracted_thinking, answer) = self.extract_thinking(&content);
12376 if extracted_thinking.is_some() {
12377 thinking_content = extracted_thinking;
12378 }
12379
12380 let output_data = match self.process_output(&answer, &input_data.context).await {
12381 Ok(d) => d,
12382 Err(e) => {
12383 yield StreamChunk::error(e.to_string());
12384 return;
12385 }
12386 };
12387
12388 let final_content = if output_data.metadata.rejected {
12389 output_data
12390 .metadata
12391 .rejection_reason
12392 .unwrap_or_else(|| answer.to_string())
12393 } else {
12394 output_data.content
12395 };
12396
12397 let (final_content, reflection_metadata) = match self
12399 .run_reflection(&*llm, processed_input, final_content)
12400 .await
12401 {
12402 Ok(r) => r,
12403 Err(e) => {
12404 yield StreamChunk::error(e.to_string());
12405 return;
12406 }
12407 };
12408
12409 let final_content = self.format_response_with_thinking(
12410 thinking_content.as_deref(),
12411 &final_content,
12412 );
12413
12414 if buffered_decision {
12416 yield StreamChunk::content(&final_content);
12417 }
12418
12419 let post_result = match self
12423 .post_loop_processing(processed_input, final_content)
12424 .await
12425 {
12426 Ok(r) => r,
12427 Err(e) => {
12428 yield StreamChunk::error(e.to_string());
12429 return;
12430 }
12431 };
12432
12433 let applied = match self.apply_post_loop_result(processed_input, post_result).await {
12434 Ok(applied) => applied,
12435 Err(e) => {
12436 yield StreamChunk::error(e.to_string());
12437 return;
12438 }
12439 };
12440
12441 if applied.transitioned {
12442 if include_state_events
12443 && let Some(state) = self.current_state()
12444 {
12445 yield StreamChunk::state_transition(None, state);
12446 }
12447 if applied.regenerated {
12453 yield StreamChunk::content(&applied.content);
12454 }
12455 }
12456 let final_content = applied.content;
12457
12458 let final_response = self.build_agent_response(AgentResponseParts {
12460 content: final_content,
12461 all_tool_calls,
12462 reasoning_mode,
12463 auto_detected,
12464 iterations,
12465 thinking: thinking_content,
12466 reflection_metadata,
12467 });
12468 if let Err(e) = self.finish_turn_if_root(&final_response).await {
12469 yield StreamChunk::error(e.to_string());
12470 return;
12471 }
12472
12473 record_runtime_stream_final(&terminal, final_response);
12474 yield StreamChunk::Done {};
12475 return;
12476 }
12477 })
12478 }
12479
12480 fn run_loop_stream<'a>(
12483 &'a self,
12484 input: &'a str,
12485 terminal: RuntimeStreamTerminalSlot,
12486 ) -> Pin<Box<dyn Stream<Item = StreamChunk> + Send + 'a>> {
12487 Box::pin(async_stream::stream! {
12488 self.begin_root_turn();
12489 let _root_cleanup = RootTurnCleanup::new(self);
12490 self.hooks.on_message_received(input).await;
12491
12492 if !self.context_initialized.swap(true, Ordering::SeqCst) {
12494 if let Err(e) = self.context_manager.initialize().await {
12495 yield StreamChunk::error(e.to_string());
12496 return;
12497 }
12498 debug!("Context manager initialized (defaults, env, builtins)");
12499 }
12500
12501 if let Err(e) = self.check_turn_timeout().await {
12502 yield StreamChunk::error(e.to_string());
12503 return;
12504 }
12505 if let Err(e) = self.context_manager.refresh_per_turn().await {
12506 yield StreamChunk::error(e.to_string());
12507 return;
12508 }
12509
12510 self.clear_disambiguation_context();
12512
12513 let input_to_run = match self.resolve_disambiguation(input).await {
12517 Err(e) => {
12518 yield StreamChunk::error(e.to_string());
12519 return;
12520 }
12521 Ok(DisambiguationDispatch::Terminal(response)) => {
12522 yield StreamChunk::content(&response.content);
12523 record_runtime_stream_final(&terminal, response);
12524 yield StreamChunk::Done {};
12525 return;
12526 }
12527 Ok(DisambiguationDispatch::RecheckSkill {
12528 skill_id,
12529 enriched_input,
12530 disambiguation_epoch,
12531 state_generation,
12532 }) => {
12533 match self
12534 .recheck_skill_disambiguation(
12535 &skill_id,
12536 &enriched_input,
12537 disambiguation_epoch,
12538 state_generation,
12539 )
12540 .await
12541 {
12542 Ok(resp) => {
12543 yield StreamChunk::content(&resp.content);
12544 record_runtime_stream_final(&terminal, resp);
12545 yield StreamChunk::Done {};
12546 return;
12547 }
12548 Err(e) => {
12549 yield StreamChunk::error(e.to_string());
12550 return;
12551 }
12552 }
12553 }
12554 Ok(DisambiguationDispatch::Proceed(input)) => input,
12555 };
12556
12557 let mut inner = self.run_loop_internal_stream(&input_to_run, Arc::clone(&terminal));
12558 while let Some(chunk) = inner.next().await {
12559 yield chunk;
12560 }
12561 })
12562 }
12563
12564 pub fn info(&self) -> AgentInfo {
12565 self.info.clone()
12566 }
12567
12568 pub fn skills(&self) -> &[SkillDefinition] {
12569 &self.skills
12570 }
12571
12572 async fn reset_runtime_state(&self) -> Result<()> {
12574 let _admission = self.disambiguation_admission.write().await;
12575 if self.state_transition_reserved.load(Ordering::SeqCst) {
12576 return Err(AgentError::Other(
12577 "Cannot reset while a state transition is in progress".to_string(),
12578 ));
12579 }
12580 self.disambiguation_epoch.fetch_add(1, Ordering::SeqCst);
12581 *self.pending_skill_id.write() = None;
12582 if let Some(disambiguator) = self.disambiguation_manager.as_ref() {
12583 disambiguator.clear_pending().await;
12584 }
12585 self.memory.clear().await?;
12586 self.active_native_exchanges.write().clear();
12587 *self.iteration_count.write() = 0;
12588 self.tool_call_history.write().clear();
12589 if let Some(ref sm) = self.state_machine {
12590 sm.reset();
12591 }
12592 Ok(())
12593 }
12594
12595 pub async fn reset(&self) -> Result<()> {
12597 self.reset_runtime_state().await
12598 }
12599
12600 pub fn max_context_tokens(&self) -> u32 {
12601 self.max_context_tokens
12602 }
12603
12604 pub fn llm_registry(&self) -> &Arc<LLMRegistry> {
12605 &self.llm_registry
12606 }
12607
12608 pub fn state_machine(&self) -> Option<&Arc<StateMachine>> {
12609 self.state_machine.as_ref()
12610 }
12611
12612 pub fn context_manager(&self) -> &Arc<ContextManager> {
12613 &self.context_manager
12614 }
12615
12616 pub fn tool_call_history(&self) -> Vec<ToolCallRecord> {
12617 self.tool_call_history.read().clone()
12618 }
12619
12620 pub fn memory_token_budget(&self) -> Option<&MemoryTokenBudget> {
12621 self.memory_token_budget.as_ref()
12622 }
12623
12624 pub fn parallel_tools_config(&self) -> &ParallelToolsConfig {
12625 &self.parallel_tools
12626 }
12627
12628 pub fn streaming_config(&self) -> &StreamingConfig {
12629 &self.streaming
12630 }
12631
12632 pub fn hooks(&self) -> &Arc<dyn AgentHooks> {
12633 &self.hooks
12634 }
12635
12636 pub fn hitl_engine(&self) -> Option<&HITLEngine> {
12637 self.hitl_engine.as_ref()
12638 }
12639
12640 pub fn approval_handler(&self) -> &Arc<dyn ApprovalHandler> {
12641 &self.approval_handler
12642 }
12643
12644 fn build_hitl_language_context(&self) -> HashMap<String, Value> {
12646 let mut ctx = HashMap::new();
12647 for key in &["user.language", "input.detected.language", "language"] {
12648 if let Some(val) = self.context_manager.get(key) {
12649 ctx.insert(key.to_string(), val);
12650 }
12651 }
12652 ctx
12653 }
12654
12655 async fn request_hitl_approval(&self, check_result: HITLCheckResult) -> Result<ApprovalResult> {
12657 let Some(request) = check_result.into_request() else {
12658 return Ok(ApprovalResult::Approved);
12659 };
12660
12661 self.hooks.on_approval_requested(&request).await;
12662
12663 let timeout = request.timeout;
12664
12665 let raw_result = if let Some(duration) = timeout {
12666 match tokio::time::timeout(
12667 duration,
12668 self.approval_handler.request_approval(request.clone()),
12669 )
12670 .await
12671 {
12672 Ok(result) => result,
12673 Err(_) => ApprovalResult::timeout(),
12674 }
12675 } else {
12676 self.approval_handler
12677 .request_approval(request.clone())
12678 .await
12679 };
12680
12681 self.hooks
12682 .on_approval_result(&request.id, &raw_result)
12683 .await;
12684
12685 let (outcome, effective_result): (ApprovalResolvedOutcome, Result<ApprovalResult>) =
12686 match &raw_result {
12687 ApprovalResult::Approved => (
12688 ApprovalResolvedOutcome::Approved,
12689 Ok(ApprovalResult::Approved),
12690 ),
12691 ApprovalResult::Rejected { reason } => (
12692 ApprovalResolvedOutcome::Rejected {
12693 reason: reason.clone(),
12694 },
12695 Ok(ApprovalResult::Rejected {
12696 reason: reason.clone(),
12697 }),
12698 ),
12699 ApprovalResult::Modified { changes } => (
12700 ApprovalResolvedOutcome::Modified {
12701 changes: changes.clone(),
12702 },
12703 Ok(ApprovalResult::Modified {
12704 changes: changes.clone(),
12705 }),
12706 ),
12707 ApprovalResult::Timeout => {
12708 if let Some(ref engine) = self.hitl_engine {
12709 match engine.config().on_timeout {
12710 TimeoutAction::Approve => (
12711 ApprovalResolvedOutcome::Approved,
12712 Ok(ApprovalResult::Approved),
12713 ),
12714 TimeoutAction::Reject => {
12715 let reason = Some("Timeout".to_string());
12716 (
12717 ApprovalResolvedOutcome::Rejected {
12718 reason: reason.clone(),
12719 },
12720 Ok(ApprovalResult::Rejected { reason }),
12721 )
12722 }
12723 TimeoutAction::Error => {
12724 let message = "HITL approval timeout".to_string();
12725 (
12726 ApprovalResolvedOutcome::Error {
12727 message: message.clone(),
12728 },
12729 Err(AgentError::Other(message)),
12730 )
12731 }
12732 }
12733 } else {
12734 let reason = Some("Timeout (no engine)".to_string());
12735 (
12736 ApprovalResolvedOutcome::Rejected {
12737 reason: reason.clone(),
12738 },
12739 Ok(ApprovalResult::Rejected { reason }),
12740 )
12741 }
12742 }
12743 };
12744
12745 self.hooks
12746 .on_approval_resolved(&request, &raw_result, &outcome)
12747 .await;
12748
12749 effective_result
12750 }
12751
12752 pub async fn check_state_hitl(&self, from: Option<&str>, to: &str) -> Result<bool> {
12753 if let Some(ref hitl_engine) = self.hitl_engine {
12754 let hitl_lang_ctx = self.build_hitl_language_context();
12755 let check_result = self
12756 .observe_purpose(
12757 ObservationPurpose::HitlLocalization,
12758 hitl_engine.check_state_transition_with_localization(
12759 from,
12760 to,
12761 &hitl_lang_ctx,
12762 self.approval_handler.as_ref(),
12763 Some(&self.llm_registry),
12764 ),
12765 )
12766 .await?;
12767 if check_result.is_required() {
12768 let result = self.request_hitl_approval(check_result).await?;
12769 return Ok(matches!(
12770 result,
12771 ApprovalResult::Approved | ApprovalResult::Modified { .. }
12772 ));
12773 }
12774 }
12775 Ok(true)
12776 }
12777
12778 async fn execute_tools_parallel(
12780 &self,
12781 tool_calls: &[ToolCall],
12782 ) -> Vec<(String, Result<String>)> {
12783 let can_run_parallel = tool_calls.iter().all(|tc| {
12784 self.tools
12785 .resolve(&tc.name)
12786 .map(|resolved| resolved.tool.classify_call(&tc.arguments).concurrency_safe)
12787 .unwrap_or(false)
12788 });
12789
12790 if !self.parallel_tools.enabled || tool_calls.len() <= 1 || !can_run_parallel {
12791 let mut results = Vec::new();
12792 for tc in tool_calls {
12793 let result = self
12794 .observe_purpose(
12795 current_observation_context()
12796 .map(|context| context.purpose)
12797 .unwrap_or_default(),
12798 self.execute_tool_smart(tc),
12799 )
12800 .await;
12801 results.push((tc.id.clone(), result));
12802 }
12803 return results;
12804 }
12805
12806 let chunks: Vec<_> = tool_calls
12807 .chunks(self.parallel_tools.max_parallel)
12808 .collect();
12809
12810 let mut all_results = Vec::new();
12811
12812 for chunk in chunks {
12813 let futures: Vec<_> = chunk
12814 .iter()
12815 .map(|tc| {
12816 let tc = tc.clone();
12817 async move {
12818 let result = self.execute_tool_smart(&tc).await;
12819 (tc.id.clone(), result)
12820 }
12821 })
12822 .collect();
12823
12824 let results = futures::future::join_all(futures).await;
12825 all_results.extend(results);
12826 }
12827
12828 all_results
12829 }
12830
12831 pub async fn chat_stream<'a>(
12835 &'a self,
12836 input: &'a str,
12837 ) -> Result<Pin<Box<dyn Stream<Item = StreamChunk> + Send + 'a>>> {
12838 let RootTurnAdmission {
12839 guard: root_turn_guard,
12840 identity_stack,
12841 } = self.acquire_root_turn().await?;
12842 scope_runtime_gate_identity_stack(&identity_stack, self.init_storage()).await?;
12846 info!(input_len = input.len(), "Starting streaming chat");
12847 let terminal = new_runtime_stream_terminal_slot();
12848 let inner = self.run_loop_stream(input, terminal);
12849 let observation_context = self.build_observation_context(None);
12850 let stream: Pin<Box<dyn Stream<Item = StreamChunk> + Send + 'a>> =
12851 Box::pin(async_stream::stream! {
12852 let mut root_turn_guard = Some(root_turn_guard);
12853 let mut inner = inner;
12854 loop {
12855 let next = scope_runtime_gate_identity_stack(&identity_stack, async {
12856 if let Some(context) = observation_context.as_ref() {
12857 with_observation_context(context.clone(), inner.next()).await
12858 } else {
12859 inner.next().await
12860 }
12861 })
12862 .await;
12863 match next {
12864 Some(StreamChunk::Done {}) => {
12865 while scope_runtime_gate_identity_stack(&identity_stack, async {
12866 if let Some(context) = observation_context.as_ref() {
12867 with_observation_context(context.clone(), inner.next())
12868 .await
12869 .is_some()
12870 } else {
12871 inner.next().await.is_some()
12872 }
12873 })
12874 .await
12875 {}
12876 if observation_context.is_some() {
12877 scope_runtime_gate_identity_stack(
12878 &identity_stack,
12879 self.export_observability_if_configured(),
12880 )
12881 .await;
12882 }
12883 drop(root_turn_guard.take());
12884 yield StreamChunk::Done {};
12885 return;
12886 }
12887 Some(chunk) => yield chunk,
12888 None => {
12889 if observation_context.is_some() {
12890 scope_runtime_gate_identity_stack(
12891 &identity_stack,
12892 self.export_observability_if_configured(),
12893 )
12894 .await;
12895 }
12896 drop(root_turn_guard.take());
12897 return;
12898 }
12899 }
12900 }
12901 });
12902 Ok(stream)
12903 }
12904
12905 pub async fn chat_stream_events<'a>(
12909 &'a self,
12910 input: &'a str,
12911 ) -> Result<Pin<Box<dyn Stream<Item = AgentStreamEvent> + Send + 'a>>> {
12912 let RootTurnAdmission {
12913 guard,
12914 identity_stack,
12915 } = self.acquire_root_turn().await?;
12916 scope_runtime_gate_identity_stack(&identity_stack, self.init_storage()).await?;
12920 info!(input_len = input.len(), "Starting streaming chat events");
12921 let terminal = new_runtime_stream_terminal_slot();
12922 let inner = self.run_loop_stream(input, Arc::clone(&terminal));
12923 let observation_context = self.build_observation_context(None);
12924 Ok(self.drive_event_stream(
12925 inner,
12926 terminal,
12927 guard,
12928 identity_stack,
12929 observation_context,
12930 None,
12931 ))
12932 }
12933
12934 pub async fn chat_stream_events_with_actor_context<'a>(
12940 &'a self,
12941 input: &'a str,
12942 actor_context: crate::TurnActorContext,
12943 ) -> Result<Pin<Box<dyn Stream<Item = AgentStreamEvent> + Send + 'a>>> {
12944 let RootTurnAdmission {
12945 guard,
12946 identity_stack,
12947 } = self.acquire_root_turn().await?;
12948 scope_runtime_gate_identity_stack(&identity_stack, self.init_storage()).await?;
12949 info!(
12950 input_len = input.len(),
12951 "Starting streaming chat events with actor context"
12952 );
12953 let actor_id = actor_context.effective_actor_id().map(str::to_string);
12954 let terminal = new_runtime_stream_terminal_slot();
12955 let inner = self.run_loop_stream(input, Arc::clone(&terminal));
12956 let observation_context = self.build_observation_context(actor_id);
12957 Ok(self.drive_event_stream(
12958 inner,
12959 terminal,
12960 guard,
12961 identity_stack,
12962 observation_context,
12963 Some(actor_context),
12964 ))
12965 }
12966
12967 fn drive_event_stream<'a>(
12976 &'a self,
12977 mut inner: Pin<Box<dyn Stream<Item = StreamChunk> + Send + 'a>>,
12978 terminal: RuntimeStreamTerminalSlot,
12979 root_turn_guard: tokio::sync::OwnedMutexGuard<()>,
12980 identity_stack: RootTurnGateIdentityStack,
12981 observation_context: Option<SpanContext>,
12982 actor_context: Option<crate::TurnActorContext>,
12983 ) -> Pin<Box<dyn Stream<Item = AgentStreamEvent> + Send + 'a>> {
12984 Box::pin(async_stream::stream! {
12985 let mut root_turn_guard = Some(root_turn_guard);
12986 loop {
12987 let next = poll_scoped_chunk(
12988 &mut inner,
12989 &identity_stack,
12990 observation_context.as_ref(),
12991 actor_context.as_ref(),
12992 )
12993 .await;
12994 match next {
12995 Some(StreamChunk::Done {}) => {
12996 let terminal_event = { terminal.write().take() };
12997 if let Some(response) = terminal_event {
12998 while poll_scoped_chunk(
12999 &mut inner,
13000 &identity_stack,
13001 observation_context.as_ref(),
13002 actor_context.as_ref(),
13003 )
13004 .await
13005 .is_some()
13006 {}
13007 if observation_context.is_some() {
13008 scope_runtime_gate_identity_stack(
13009 &identity_stack,
13010 self.export_observability_if_configured(),
13011 )
13012 .await;
13013 }
13014 drop(root_turn_guard.take());
13015 yield AgentStreamEvent::Final(response);
13016 return;
13017 }
13018 }
13019 Some(StreamChunk::Error { message }) => {
13020 let finalized = { terminal.read().is_some() };
13021 if finalized {
13022 continue;
13023 }
13024 while poll_scoped_chunk(
13025 &mut inner,
13026 &identity_stack,
13027 observation_context.as_ref(),
13028 actor_context.as_ref(),
13029 )
13030 .await
13031 .is_some()
13032 {}
13033 if observation_context.is_some() {
13034 scope_runtime_gate_identity_stack(
13035 &identity_stack,
13036 self.export_observability_if_configured(),
13037 )
13038 .await;
13039 }
13040 drop(root_turn_guard.take());
13041 yield AgentStreamEvent::Chunk(StreamChunk::Error { message });
13042 return;
13043 }
13044 Some(chunk) => yield AgentStreamEvent::Chunk(chunk),
13045 None => {
13046 if observation_context.is_some() {
13047 scope_runtime_gate_identity_stack(
13048 &identity_stack,
13049 self.export_observability_if_configured(),
13050 )
13051 .await;
13052 }
13053 drop(root_turn_guard.take());
13054 return;
13055 }
13056 }
13057 }
13058 })
13059 }
13060}
13061
13062async fn poll_scoped_chunk<'a>(
13068 inner: &mut Pin<Box<dyn Stream<Item = StreamChunk> + Send + 'a>>,
13069 identity_stack: &RootTurnGateIdentityStack,
13070 observation_context: Option<&SpanContext>,
13071 actor_context: Option<&crate::TurnActorContext>,
13072) -> Option<StreamChunk> {
13073 scope_runtime_gate_identity_stack(identity_stack, async {
13074 let next = inner.next();
13075 match (observation_context, actor_context) {
13076 (Some(observation), Some(actor)) => {
13077 with_observation_context(
13078 observation.clone(),
13079 scope_actor_context(actor.clone(), next),
13080 )
13081 .await
13082 }
13083 (Some(observation), None) => with_observation_context(observation.clone(), next).await,
13084 (None, Some(actor)) => scope_actor_context(actor.clone(), next).await,
13085 (None, None) => next.await,
13086 }
13087 })
13088 .await
13089}
13090
13091#[async_trait]
13092impl ToolInvoker for RuntimeAgent {
13093 async fn invoke_tool(&self, request: ToolExecutionRequest) -> Result<ToolExecutionRecord> {
13094 self.execute_tool_record(request).await
13095 }
13096}
13097
13098#[async_trait]
13099impl Agent for RuntimeAgent {
13100 async fn chat(&self, input: &str) -> Result<AgentResponse> {
13102 let RootTurnAdmission {
13103 guard,
13104 identity_stack,
13105 } = self.acquire_root_turn().await?;
13106 let result = scope_runtime_gate_identity_stack(&identity_stack, async {
13107 let result = if let Some(context) = self.build_observation_context(None) {
13108 with_observation_context(context, self.run_loop(input)).await
13109 } else {
13110 self.run_loop(input).await
13111 };
13112 self.export_observability_if_configured().await;
13113 result
13114 })
13115 .await;
13116 drop(guard);
13117 result
13118 }
13119
13120 fn info(&self) -> AgentInfo {
13121 self.info.clone()
13122 }
13123
13124 async fn reset(&self) -> Result<()> {
13126 self.reset_runtime_state().await
13127 }
13128}
13129
13130fn background_maintenance_tags(
13140 label: &str,
13141 stage: &str,
13142 reason: Option<&str>,
13143 policy: Option<&crate::optimization::config::MaintenanceTaskPolicy>,
13144) -> HashMap<String, String> {
13145 let mut tags = HashMap::new();
13146 tags.insert("runtime.background".to_string(), "true".to_string());
13147 tags.insert("runtime.maintenance".to_string(), label.to_string());
13148 tags.insert("runtime.maintenance_stage".to_string(), stage.to_string());
13149 if let Some(policy) = policy {
13150 tags.insert(
13151 "runtime.await_before_next_turn".to_string(),
13152 await_before_next_turn_label(policy.await_before_next_turn).to_string(),
13153 );
13154 tags.insert(
13155 "runtime.maintenance_mode".to_string(),
13156 maintenance_mode_label(policy.mode).to_string(),
13157 );
13158 }
13159 if let Some(reason) = reason {
13160 tags.insert("runtime.reason".to_string(), reason.to_string());
13161 }
13162 tags
13163}
13164
13165fn await_before_next_turn_label(policy: AwaitBeforeNextTurn) -> &'static str {
13166 match policy {
13167 AwaitBeforeNextTurn::Never => "never",
13168 AwaitBeforeNextTurn::SameActor => "same_actor",
13169 AwaitBeforeNextTurn::Always => "always",
13170 }
13171}
13172
13173fn maintenance_mode_label(mode: MaintenanceMode) -> &'static str {
13174 match mode {
13175 MaintenanceMode::InlineSerial => "inline_serial",
13176 MaintenanceMode::InlineParallel => "inline_parallel",
13177 MaintenanceMode::Background => "background",
13178 }
13179}
13180
13181fn record_background_maintenance_event(
13183 manager: Option<&Arc<ObservabilityManager>>,
13184 label: &str,
13185 status: EventStatus,
13186 duration_ms: u64,
13187 stage: &str,
13188 reason: Option<String>,
13189 policy: Option<&crate::optimization::config::MaintenanceTaskPolicy>,
13190) {
13191 if let Some(manager) = manager {
13192 manager.record_lifecycle_event(
13193 EventType::MemoryOperation {
13194 operation: format!("{}_background_{}", label, stage),
13195 },
13196 ObservationPurpose::Other(format!("{}_maintenance", label)),
13197 status,
13198 duration_ms,
13199 background_maintenance_tags(label, stage, reason.as_deref(), policy),
13200 None,
13201 );
13202 }
13203}
13204
13205fn effective_maintenance_mode(mode: MaintenanceMode, force_parallel: bool) -> MaintenanceMode {
13206 if force_parallel && matches!(mode, MaintenanceMode::InlineSerial) {
13207 MaintenanceMode::InlineParallel
13208 } else {
13209 mode
13210 }
13211}
13212
13213fn observation_purpose_for_process(hint: ProcessPurposeHint) -> ObservationPurpose {
13214 match hint {
13215 ProcessPurposeHint::Detect => ObservationPurpose::ProcessDetect,
13216 ProcessPurposeHint::Extract => ObservationPurpose::ProcessExtract,
13217 ProcessPurposeHint::Validate => ObservationPurpose::ProcessValidate,
13218 ProcessPurposeHint::Transform | ProcessPurposeHint::Other => {
13219 ObservationPurpose::ProcessTransform
13220 }
13221 }
13222}
13223
13224fn new_tool_resource_locks() -> ToolResourceLocks {
13225 Arc::new(RwLock::new(HashMap::new()))
13226}
13227
13228fn tool_resource_lock_keys(
13233 _canonical_id: &str,
13234 args: &Value,
13235 bindings: &ai_agents_core::ToolPolicyBindings,
13236 classification: &ai_agents_core::ToolCallClassification,
13237) -> Vec<String> {
13238 if classification.concurrency_safe {
13239 return Vec::new();
13240 }
13241
13242 let mut keys = Vec::new();
13243 let mut has_path_resource = false;
13244 for binding in &bindings.path_fields {
13245 let value = value_at_argument_path(args, &binding.field)
13246 .cloned()
13247 .or_else(|| {
13248 binding
13249 .default_path
13250 .as_ref()
13251 .map(|path| Value::String(path.clone()))
13252 });
13253 if let Some(value) = value {
13254 collect_resource_strings(&value, |_| {
13255 has_path_resource = true;
13256 });
13257 }
13258 }
13259 for binding in &bindings.domain_fields {
13260 if let Some(value) = value_at_argument_path(args, &binding.field) {
13261 collect_resource_strings(value, |domain| {
13262 let normalized = if binding.is_url {
13263 normalized_url_resource_key(domain)
13264 } else {
13265 domain.trim().trim_end_matches('.').to_ascii_lowercase()
13266 };
13267 keys.push(format!("domain:{}", normalized));
13268 });
13269 }
13270 }
13271 for binding in &bindings.command_fields {
13272 if !matches!(binding.kind, ai_agents_core::CommandBindingKind::Cwd) {
13273 continue;
13274 }
13275 if let Some(value) = value_at_argument_path(args, &binding.field) {
13276 collect_resource_strings(value, |_| {
13277 has_path_resource = true;
13278 });
13279 }
13280 }
13281 if has_path_resource {
13282 keys.push("path-mutation:global".to_string());
13283 }
13284 if keys.is_empty() {
13285 keys.push("side-effect:unbound".to_string());
13286 }
13287 keys.sort();
13288 keys.dedup();
13289 keys
13290}
13291
13292fn value_at_argument_path<'a>(value: &'a Value, field: &str) -> Option<&'a Value> {
13293 let mut current = value;
13294 for segment in field.split('.') {
13295 if segment.is_empty() {
13296 return None;
13297 }
13298 current = current.get(segment)?;
13299 }
13300 Some(current)
13301}
13302
13303fn collect_resource_strings(value: &Value, mut collect: impl FnMut(&str)) {
13304 match value {
13305 Value::String(value) => collect(value),
13306 Value::Array(values) => {
13307 for value in values {
13308 if let Some(value) = value.as_str() {
13309 collect(value);
13310 }
13311 }
13312 }
13313 _ => {}
13314 }
13315}
13316
13317fn normalized_url_resource_key(value: &str) -> String {
13318 let value = value.trim();
13319 let Some((scheme, remainder)) = value.split_once("://") else {
13320 return value.to_ascii_lowercase();
13321 };
13322 let authority_end = remainder.find(['/', '?', '#']).unwrap_or(remainder.len());
13323 let (authority, suffix) = remainder.split_at(authority_end);
13324 format!(
13325 "{}://{}{}",
13326 scheme.to_ascii_lowercase(),
13327 authority.to_ascii_lowercase(),
13328 suffix
13329 )
13330}
13331
13332fn render_concurrent_template(
13333 template: &str,
13334 user_input: &str,
13335 context_values: &std::collections::HashMap<String, serde_json::Value>,
13336) -> Result<String> {
13337 let mut env = minijinja::Environment::new();
13338 env.add_template("concurrent", template)
13339 .map_err(|e| AgentError::Other(format!("Concurrent template parse error: {}", e)))?;
13340
13341 let mut ctx = std::collections::BTreeMap::new();
13342 ctx.insert("user_input".to_string(), minijinja::Value::from(user_input));
13343
13344 let context_obj = minijinja::Value::from_serialize(context_values);
13346 ctx.insert("context".to_string(), context_obj);
13347
13348 let tmpl = env
13349 .get_template("concurrent")
13350 .map_err(|e| AgentError::Other(format!("Concurrent template error: {}", e)))?;
13351
13352 tmpl.render(minijinja::Value::from_serialize(&ctx))
13353 .map_err(|e| AgentError::Other(format!("Concurrent template render error: {}", e)))
13354}
13355
13356#[cfg(test)]
13357mod tests {
13358 use super::*;
13359 use crate::AgentBuilder;
13360 use ai_agents_core::{LLMChunk, LLMConfig, LLMError, LLMFeature, Tool};
13361 use ai_agents_llm::mock::MockLLMProvider;
13362 use ai_agents_skills::{SkillDefinition, SkillStep};
13363 use ai_agents_tools::{
13364 CalculatorTool, CopyPathTool, DeletePathTool, FileWriteTool, MovePathTool, ToolAliases,
13365 ToolDescriptor, ToolProvider, ToolProviderError, ToolProviderType, WebFetchResolver,
13366 WebFetchTool, WebFetchTransport, WebFetchTransportRequest, WebFetchTransportResponse,
13367 };
13368
13369 fn mock_with_response(response: &str) -> MockLLMProvider {
13370 let mut mock = MockLLMProvider::new("test");
13371 mock.set_response(response);
13372 mock
13373 }
13374
13375 fn mock_with_responses(responses: Vec<&str>) -> MockLLMProvider {
13376 let mut mock = MockLLMProvider::new("test");
13377 mock.set_responses(responses.into_iter().map(String::from).collect(), true);
13378 mock
13379 }
13380
13381 async fn collect_stream_events(
13383 agent: &RuntimeAgent,
13384 input: &str,
13385 ) -> (String, Vec<StreamChunk>, Option<AgentResponse>) {
13386 use futures::StreamExt;
13387 let mut events = agent.chat_stream_events(input).await.expect("stream opens");
13388 let mut content = String::new();
13389 let mut chunks = Vec::new();
13390 let mut final_response = None;
13391 while let Some(event) = events.next().await {
13392 match event {
13393 AgentStreamEvent::Chunk(chunk) => {
13394 if let StreamChunk::Content { text } = &chunk {
13395 content.push_str(text);
13396 }
13397 chunks.push(chunk);
13398 }
13399 AgentStreamEvent::Final(response) => final_response = Some(response),
13400 }
13401 }
13402 (content, chunks, final_response)
13403 }
13404
13405 fn metadata_keys(response: &AgentResponse) -> std::collections::BTreeSet<String> {
13406 response
13407 .metadata
13408 .as_ref()
13409 .map(|m| m.keys().cloned().collect())
13410 .unwrap_or_default()
13411 }
13412
13413 async fn assert_blocking_streaming_parity<F>(
13416 build: F,
13417 input: &str,
13418 ) -> (AgentResponse, AgentResponse, Vec<StreamChunk>)
13419 where
13420 F: Fn() -> RuntimeAgent,
13421 {
13422 let blocking_agent = build();
13423 let streaming_agent = build();
13424
13425 let blocking = blocking_agent
13426 .chat(input)
13427 .await
13428 .expect("blocking chat succeeds");
13429 let (_, chunks, final_response) = collect_stream_events(&streaming_agent, input).await;
13430 let streamed = final_response.expect("streaming must emit Final when blocking succeeds");
13431
13432 assert_eq!(
13433 blocking.content, streamed.content,
13434 "committed content differs"
13435 );
13436 assert_eq!(
13437 metadata_keys(&blocking),
13438 metadata_keys(&streamed),
13439 "metadata key sets differ"
13440 );
13441 assert_eq!(
13442 blocking.tool_calls.as_ref().map(Vec::len),
13443 streamed.tool_calls.as_ref().map(Vec::len),
13444 "tool call counts differ"
13445 );
13446 assert_eq!(
13447 blocking_agent.current_state(),
13448 streaming_agent.current_state(),
13449 "final states differ"
13450 );
13451 (blocking, streamed, chunks)
13452 }
13453
13454 fn signed_calculator_response(
13455 exchange_id: &str,
13456 call_id: &str,
13457 expression: &str,
13458 ) -> LLMResponse {
13459 let call = ToolCall {
13460 id: call_id.to_string(),
13461 name: "calculator".to_string(),
13462 arguments: serde_json::json!({"expression": expression}),
13463 };
13464 let state = ai_agents_core::NativeProviderState::new(
13465 exchange_id,
13466 "fixture",
13467 "native-tools",
13468 ai_agents_core::NativeProviderTarget::new("https://fixture.invalid/", "fixture-model")
13469 .unwrap(),
13470 serde_json::json!({
13471 "role": "model",
13472 "parts": [{
13473 "functionCall": {"name": "calculator", "args": {"expression": expression}},
13474 "thoughtSignature": format!("signature-{exchange_id}")
13475 }]
13476 }),
13477 vec![ai_agents_core::NativeCallBinding::new(call_id, 0).unwrap()],
13478 )
13479 .unwrap();
13480 LLMResponse::new("", FinishReason::ToolCall)
13481 .with_provider_state(state)
13482 .unwrap()
13483 .with_tool_calls(vec![call])
13484 .unwrap()
13485 }
13486
13487 struct TerminalHistoryProvider {
13488 calls: Arc<std::sync::atomic::AtomicU32>,
13489 }
13490
13491 struct DroppingSignedAssistantMemory {
13492 messages: RwLock<Vec<ChatMessage>>,
13493 }
13494
13495 struct DroppingEarlierSequentialMemory {
13496 messages: RwLock<Vec<ChatMessage>>,
13497 signed_seen: std::sync::atomic::AtomicUsize,
13498 }
13499
13500 #[async_trait]
13501 impl ai_agents_core::Memory for DroppingSignedAssistantMemory {
13502 async fn add_message(&self, message: ChatMessage) -> Result<()> {
13503 let signed = message.role == ai_agents_core::Role::Assistant
13504 && ai_agents_core::decode_native_tool_call_markers(&message.content)
13505 .map_err(|error| AgentError::LLM(error.to_string()))?
13506 .is_some_and(|batch| batch.provider_state().is_some());
13507 if !signed {
13508 self.messages.write().push(message);
13509 }
13510 Ok(())
13511 }
13512
13513 async fn get_messages(&self, limit: Option<usize>) -> Result<Vec<ChatMessage>> {
13514 let messages = self.messages.read();
13515 let start = limit
13516 .map(|limit| messages.len().saturating_sub(limit))
13517 .unwrap_or(0);
13518 Ok(messages[start..].to_vec())
13519 }
13520
13521 async fn clear(&self) -> Result<()> {
13522 self.messages.write().clear();
13523 Ok(())
13524 }
13525
13526 fn len(&self) -> usize {
13527 self.messages.read().len()
13528 }
13529
13530 async fn restore(&self, snapshot: ai_agents_core::MemorySnapshot) -> Result<()> {
13531 *self.messages.write() = snapshot.messages;
13532 Ok(())
13533 }
13534 }
13535
13536 #[async_trait]
13537 impl ai_agents_memory::Memory for DroppingSignedAssistantMemory {}
13538
13539 #[async_trait]
13540 impl ai_agents_core::Memory for DroppingEarlierSequentialMemory {
13541 async fn add_message(&self, message: ChatMessage) -> Result<()> {
13542 let signed = message.role == ai_agents_core::Role::Assistant
13543 && ai_agents_core::decode_native_tool_call_markers(&message.content)
13544 .map_err(|error| AgentError::LLM(error.to_string()))?
13545 .is_some_and(|batch| batch.provider_state().is_some());
13546 let mut messages = self.messages.write();
13547 if signed && self.signed_seen.fetch_add(1, Ordering::SeqCst) == 1 {
13548 messages.retain(|stored| !stored.content.contains("seq-call-1"));
13549 }
13550 messages.push(message);
13551 Ok(())
13552 }
13553
13554 async fn get_messages(&self, limit: Option<usize>) -> Result<Vec<ChatMessage>> {
13555 let messages = self.messages.read();
13556 let start = limit
13557 .map(|limit| messages.len().saturating_sub(limit))
13558 .unwrap_or(0);
13559 Ok(messages[start..].to_vec())
13560 }
13561
13562 async fn clear(&self) -> Result<()> {
13563 self.messages.write().clear();
13564 self.signed_seen.store(0, Ordering::SeqCst);
13565 Ok(())
13566 }
13567
13568 fn len(&self) -> usize {
13569 self.messages.read().len()
13570 }
13571
13572 async fn restore(&self, snapshot: ai_agents_core::MemorySnapshot) -> Result<()> {
13573 *self.messages.write() = snapshot.messages;
13574 self.signed_seen.store(0, Ordering::SeqCst);
13575 Ok(())
13576 }
13577 }
13578
13579 #[async_trait]
13580 impl ai_agents_memory::Memory for DroppingEarlierSequentialMemory {}
13581
13582 #[async_trait]
13583 impl LLMProvider for TerminalHistoryProvider {
13584 async fn complete(
13585 &self,
13586 _messages: &[ChatMessage],
13587 _config: Option<&LLMConfig>,
13588 ) -> std::result::Result<LLMResponse, LLMError> {
13589 self.calls.fetch_add(1, Ordering::SeqCst);
13590 Err(LLMError::Serialization(
13591 "native history integrity failure".to_string(),
13592 ))
13593 }
13594
13595 async fn complete_stream(
13596 &self,
13597 _messages: &[ChatMessage],
13598 _config: Option<&LLMConfig>,
13599 ) -> std::result::Result<
13600 Box<dyn Stream<Item = std::result::Result<LLMChunk, LLMError>> + Unpin + Send>,
13601 LLMError,
13602 > {
13603 Err(LLMError::Serialization(
13604 "native history integrity failure".to_string(),
13605 ))
13606 }
13607
13608 fn provider_name(&self) -> &str {
13609 "terminal-history"
13610 }
13611
13612 fn supports(&self, _feature: LLMFeature) -> bool {
13613 false
13614 }
13615
13616 fn is_terminal_error(&self, error: &LLMError) -> bool {
13617 matches!(error, LLMError::Serialization(_))
13618 }
13619 }
13620
13621 fn disambiguation_state_machine(
13623 state_enabled: Option<bool>,
13624 require_confirmation: bool,
13625 ) -> Arc<StateMachine> {
13626 let definition = ai_agents_state::StateDefinition {
13627 prompt: Some("Handle the resolved request.".to_string()),
13628 disambiguation: Some(ai_agents_disambiguation::StateDisambiguationOverride {
13629 enabled: state_enabled,
13630 require_confirmation,
13631 ..Default::default()
13632 }),
13633 ..Default::default()
13634 };
13635 let review = ai_agents_state::StateDefinition {
13636 prompt: Some("Review a fresh request.".to_string()),
13637 ..Default::default()
13638 };
13639 Arc::new(
13640 StateMachine::new(ai_agents_state::StateConfig {
13641 initial: "active".to_string(),
13642 states: std::collections::HashMap::from([
13643 ("active".to_string(), definition),
13644 ("review".to_string(), review),
13645 ]),
13646 global_transitions: Vec::new(),
13647 fallback: None,
13648 max_no_transition: None,
13649 regenerate_on_transition: true,
13650 })
13651 .unwrap(),
13652 )
13653 }
13654
13655 fn state_disambiguation_agent(
13657 responses: Vec<&str>,
13658 manager_enabled: bool,
13659 state_enabled: Option<bool>,
13660 require_confirmation: bool,
13661 ) -> (RuntimeAgent, MockLLMProvider) {
13662 state_disambiguation_agent_with_skills(
13663 responses,
13664 manager_enabled,
13665 state_enabled,
13666 require_confirmation,
13667 Vec::new(),
13668 )
13669 }
13670
13671 fn state_disambiguation_agent_with_skills(
13673 responses: Vec<&str>,
13674 manager_enabled: bool,
13675 state_enabled: Option<bool>,
13676 require_confirmation: bool,
13677 skills: Vec<SkillDefinition>,
13678 ) -> (RuntimeAgent, MockLLMProvider) {
13679 let mut mock = MockLLMProvider::new("state-confirmation");
13680 mock.set_responses(responses.into_iter().map(String::from).collect(), false);
13681 let observed = mock.clone();
13682 let agent = AgentBuilder::new()
13683 .system_prompt("Handle requests.")
13684 .llm(Arc::new(mock.clone()))
13685 .llm_alias("router", Arc::new(mock))
13686 .state_machine(disambiguation_state_machine(
13687 state_enabled,
13688 require_confirmation,
13689 ))
13690 .skills(skills)
13691 .build()
13692 .unwrap()
13693 .with_disambiguation(DisambiguationConfig {
13694 enabled: manager_enabled,
13695 ..Default::default()
13696 });
13697 (agent, observed)
13698 }
13699
13700 fn confirmation_skill() -> SkillDefinition {
13702 SkillDefinition {
13703 id: "send_report".to_string(),
13704 description: "Send a report after clarification".to_string(),
13705 trigger: "When the user asks to send a report".to_string(),
13706 steps: vec![SkillStep::Prompt {
13707 prompt: "Execute confirmed report skill for: {{ input }}".to_string(),
13708 llm: None,
13709 }],
13710 reasoning: None,
13711 reflection: None,
13712 disambiguation: Some(ai_agents_disambiguation::SkillDisambiguationOverride {
13713 enabled: Some(true),
13714 ..Default::default()
13715 }),
13716 }
13717 }
13718
13719 fn confirmation_skill_call_count(observed: &MockLLMProvider) -> usize {
13721 observed
13722 .call_history()
13723 .iter()
13724 .filter(|call| {
13725 call.messages
13726 .iter()
13727 .any(|message| message.content.contains("Execute confirmed report skill"))
13728 })
13729 .count()
13730 }
13731
13732 struct BlockingRuntimeConfirmationObserver {
13733 entered: tokio::sync::Barrier,
13734 release: tokio::sync::Notify,
13735 }
13736
13737 impl BlockingRuntimeConfirmationObserver {
13738 fn new() -> Self {
13739 Self {
13740 entered: tokio::sync::Barrier::new(2),
13741 release: tokio::sync::Notify::new(),
13742 }
13743 }
13744 }
13745
13746 struct ResetOnTransitionHooks {
13747 agent: parking_lot::Mutex<Option<Weak<RuntimeAgent>>>,
13748 invoked: AtomicBool,
13749 }
13750
13751 #[async_trait]
13752 impl AgentHooks for ResetOnTransitionHooks {
13753 async fn on_state_transition(&self, _from: Option<&str>, _to: &str, _reason: &str) {
13754 if self.invoked.swap(true, Ordering::SeqCst) {
13755 return;
13756 }
13757 let agent = self.agent.lock().as_ref().and_then(Weak::upgrade);
13758 if let Some(agent) = agent {
13759 agent.reset().await.unwrap();
13760 }
13761 }
13762 }
13763
13764 impl ClarificationObserver for BlockingRuntimeConfirmationObserver {
13765 fn observe_question<'a>(
13766 &'a self,
13767 future: ClarificationQuestionFuture<'a>,
13768 ) -> ClarificationQuestionFuture<'a> {
13769 future
13770 }
13771
13772 fn observe_parse<'a>(
13773 &'a self,
13774 future: ClarificationParseFuture<'a>,
13775 ) -> ClarificationParseFuture<'a> {
13776 future
13777 }
13778
13779 fn observe_confirmation_parse<'a>(
13780 &'a self,
13781 future: ConfirmationParseFuture<'a>,
13782 ) -> ConfirmationParseFuture<'a> {
13783 Box::pin(async move {
13784 self.entered.wait().await;
13785 self.release.notified().await;
13786 future.await
13787 })
13788 }
13789 }
13790
13791 #[tokio::test]
13792 async fn state_confirmation_blocks_redispatch_until_explicit_agreement() {
13793 let (agent, observed) = state_disambiguation_agent(
13794 vec![
13795 r#"{"is_ambiguous":true,"confidence":0.2,"ambiguity_type":"missing_target","reasoning":"target missing","what_is_unclear":["target"],"detected_language":"en"}"#,
13796 r#"{"question":"What should I send?","options":null}"#,
13797 r#"{"status":"answered","selected_option":null,"enriched_input":"Send the report to Ada","resolved":{"intent":"send_report"}}"#,
13798 r#"{"question":"Should I send the report to Ada?"}"#,
13799 r#"{"status":"confirmed"}"#,
13800 "Request executed.",
13801 ],
13802 true,
13803 None,
13804 true,
13805 );
13806
13807 let clarification = agent.chat("Send it").await.unwrap();
13808 assert_eq!(clarification.content, "What should I send?");
13809 assert_eq!(observed.call_count(), 2);
13810
13811 let confirmation = agent.chat("The report to Ada").await.unwrap();
13812 assert_eq!(confirmation.content, "Should I send the report to Ada?");
13813 assert_eq!(
13814 confirmation
13815 .metadata
13816 .as_ref()
13817 .and_then(|metadata| metadata.get("disambiguation"))
13818 .and_then(|metadata| metadata.get("status"))
13819 .and_then(Value::as_str),
13820 Some("awaiting_confirmation")
13821 );
13822 assert_eq!(observed.call_count(), 4);
13823
13824 let completed = agent.chat("Yes").await.unwrap();
13825 assert_eq!(completed.content, "Request executed.");
13826 assert_eq!(observed.call_count(), 6);
13827 }
13828
13829 #[tokio::test]
13830 async fn streaming_state_confirmation_ends_the_turn_before_redispatch() {
13831 let (agent, observed) = state_disambiguation_agent(
13832 vec![
13833 r#"{"is_ambiguous":true,"confidence":0.2,"ambiguity_type":"missing_target","reasoning":"target missing","what_is_unclear":["target"],"detected_language":"en"}"#,
13834 r#"{"question":"What should I send?","options":null}"#,
13835 r#"{"status":"answered","selected_option":null,"enriched_input":"Send the report to Ada","resolved":{"intent":"send_report"}}"#,
13836 r#"{"question":"Should I send the report to Ada?"}"#,
13837 r#"{"status":"confirmed"}"#,
13838 "Request executed.",
13839 ],
13840 true,
13841 None,
13842 true,
13843 );
13844
13845 let mut clarification_stream = agent.chat_stream("Send it").await.unwrap();
13846 let mut clarification = String::new();
13847 while let Some(chunk) = clarification_stream.next().await {
13848 match chunk {
13849 StreamChunk::Content { text } => clarification.push_str(&text),
13850 StreamChunk::Done {} => break,
13851 StreamChunk::Error { message } => panic!("unexpected stream error: {message}"),
13852 _ => {}
13853 }
13854 }
13855 assert_eq!(clarification, "What should I send?");
13856 assert_eq!(observed.call_count(), 2);
13857
13858 let mut confirmation_stream = agent.chat_stream_events("The report to Ada").await.unwrap();
13859 let mut confirmation = None;
13860 while let Some(event) = confirmation_stream.next().await {
13861 match event {
13862 AgentStreamEvent::Final(response) => confirmation = Some(response),
13863 AgentStreamEvent::Chunk(StreamChunk::Error { message }) => {
13864 panic!("unexpected stream error: {message}")
13865 }
13866 AgentStreamEvent::Chunk(_) => {}
13867 }
13868 }
13869 let confirmation = confirmation.expect("confirmation must finalize");
13870 assert_eq!(confirmation.content, "Should I send the report to Ada?");
13871 assert_eq!(
13872 confirmation
13873 .metadata
13874 .as_ref()
13875 .and_then(|metadata| metadata.get("disambiguation"))
13876 .and_then(|metadata| metadata.get("status"))
13877 .and_then(Value::as_str),
13878 Some("awaiting_confirmation")
13879 );
13880 assert_eq!(observed.call_count(), 4);
13881
13882 let mut completed_stream = agent.chat_stream("Yes").await.unwrap();
13883 let mut completed = String::new();
13884 while let Some(chunk) = completed_stream.next().await {
13885 match chunk {
13886 StreamChunk::Content { text } => completed.push_str(&text),
13887 StreamChunk::Done {} => break,
13888 StreamChunk::Error { message } => panic!("unexpected stream error: {message}"),
13889 _ => {}
13890 }
13891 }
13892 assert_eq!(completed, "Request executed.");
13893 assert_eq!(observed.call_count(), 6);
13894 }
13895
13896 #[tokio::test]
13898 async fn root_turn_gate_serializes_blocking_and_streaming_entry_points() {
13899 let (complete_entered, mut complete_events) = tokio::sync::mpsc::unbounded_channel();
13900 let agent = Arc::new(
13901 AgentBuilder::new()
13902 .system_prompt("Serialize root turns.")
13903 .llm(Arc::new(RootTurnProbeProvider { complete_entered }))
13904 .build()
13905 .unwrap(),
13906 );
13907 let blocking_agent = Arc::clone(&agent);
13908
13909 let legacy_stream = agent.chat_stream("stream owner").await.unwrap();
13910 assert!(agent.root_turn_gate.try_lock().is_err());
13911 let blocking = tokio::spawn(async move { blocking_agent.chat("blocked").await.unwrap() });
13912 assert!(
13913 tokio::time::timeout(std::time::Duration::from_millis(50), complete_events.recv())
13914 .await
13915 .is_err(),
13916 "blocking turn reached the provider while the legacy stream owned the root gate"
13917 );
13918
13919 drop(legacy_stream);
13920 assert_eq!(
13921 tokio::time::timeout(std::time::Duration::from_secs(2), complete_events.recv())
13922 .await
13923 .expect("blocking turn did not enter after stream drop"),
13924 Some(())
13925 );
13926 let response = tokio::time::timeout(std::time::Duration::from_secs(2), blocking)
13927 .await
13928 .expect("blocking turn did not finish after stream drop")
13929 .unwrap();
13930 assert_eq!(response.content, "blocking complete");
13931
13932 let mut event_stream = agent.chat_stream_events("event terminal").await.unwrap();
13933 assert!(agent.root_turn_gate.try_lock().is_err());
13934 let mut saw_final = false;
13935 while let Some(event) = event_stream.next().await {
13936 if matches!(event, AgentStreamEvent::Final(_)) {
13937 saw_final = true;
13938 break;
13939 }
13940 }
13941 assert!(saw_final);
13942 assert!(
13943 agent.root_turn_gate.try_lock().is_ok(),
13944 "authoritative terminal event retained the root gate"
13945 );
13946 }
13947
13948 #[tokio::test]
13950 async fn response_hook_rejects_same_runtime_chat_reentry() {
13951 let hooks = Arc::new(ResponseChatHooks {
13952 target: parking_lot::Mutex::new(None),
13953 invoked: AtomicBool::new(false),
13954 nested_result: parking_lot::Mutex::new(None),
13955 });
13956 let agent = Arc::new(
13957 AgentBuilder::new()
13958 .system_prompt("Reject response hook reentry.")
13959 .llm(Arc::new(mock_with_response("outer response")))
13960 .hooks(hooks.clone())
13961 .build()
13962 .unwrap(),
13963 );
13964 *hooks.target.lock() = Some(Arc::downgrade(&agent));
13965
13966 let response = tokio::time::timeout(
13967 std::time::Duration::from_secs(2),
13968 agent.chat("outer request"),
13969 )
13970 .await
13971 .expect("same-runtime response hook reentry must fail without deadlocking")
13972 .unwrap();
13973
13974 assert_eq!(response.content, "outer response");
13975 let nested_result = hooks
13976 .nested_result
13977 .lock()
13978 .clone()
13979 .expect("response hook must record its nested call");
13980 let error = nested_result.expect_err("same-runtime nested chat must be rejected");
13981 assert!(error.contains("reentrant root turn ownership"));
13982 }
13983
13984 #[tokio::test]
13986 async fn root_turn_gate_allows_nested_runtime_and_rejects_cycles() {
13987 let agent_a = AgentBuilder::new()
13988 .system_prompt("Runtime A.")
13989 .llm(Arc::new(mock_with_response("response A")))
13990 .build()
13991 .unwrap();
13992 let agent_b = AgentBuilder::new()
13993 .system_prompt("Runtime B.")
13994 .llm(Arc::new(mock_with_response("response B")))
13995 .build()
13996 .unwrap();
13997 let RootTurnAdmission {
13998 guard: guard_a,
13999 identity_stack: stack_a,
14000 } = agent_a.acquire_root_turn().await.unwrap();
14001
14002 let cycle_error = scope_runtime_gate_identity_stack(&stack_a, async {
14003 let RootTurnAdmission {
14004 guard: guard_b,
14005 identity_stack: stack_b,
14006 } = agent_b
14007 .acquire_root_turn()
14008 .await
14009 .expect("runtime B must acquire a different gate");
14010 let result =
14011 scope_runtime_gate_identity_stack(&stack_b, agent_a.acquire_root_turn()).await;
14012 drop(guard_b);
14013 match result {
14014 Err(error) => error,
14015 Ok(_) => panic!("runtime A accepted a repeated gate identity"),
14016 }
14017 })
14018 .await;
14019 drop(guard_a);
14020
14021 assert!(
14022 cycle_error
14023 .to_string()
14024 .contains("reentrant root turn ownership")
14025 );
14026 }
14027
14028 #[tokio::test]
14030 async fn concurrent_orchestration_propagates_root_gate_ancestry() {
14031 let registry = Arc::new(crate::spawner::AgentRegistry::new());
14032 let hooks_a = Arc::new(ConcurrentResponseHooks {
14033 registry: Arc::downgrade(®istry),
14034 child_id: "runtime-b".to_string(),
14035 invoked: AtomicBool::new(false),
14036 nested_result: parking_lot::Mutex::new(None),
14037 });
14038 let hooks_b = Arc::new(ResponseChatHooks {
14039 target: parking_lot::Mutex::new(None),
14040 invoked: AtomicBool::new(false),
14041 nested_result: parking_lot::Mutex::new(None),
14042 });
14043 let agent_a = AgentBuilder::new()
14044 .system_prompt("Runtime A dispatches runtime B concurrently.")
14045 .llm(Arc::new(mock_with_response("response A")))
14046 .hooks(hooks_a.clone())
14047 .build()
14048 .unwrap();
14049 let agent_b = AgentBuilder::new()
14050 .system_prompt("Runtime B attempts to re-enter runtime A.")
14051 .llm(Arc::new(mock_with_response("response B")))
14052 .hooks(hooks_b.clone())
14053 .build()
14054 .unwrap();
14055 let spec_a = crate::spec::AgentSpec {
14056 name: "runtime-a".to_string(),
14057 system_prompt: "Runtime A dispatches runtime B concurrently.".to_string(),
14058 ..crate::spec::AgentSpec::default()
14059 };
14060 let spec_b = crate::spec::AgentSpec {
14061 name: "runtime-b".to_string(),
14062 system_prompt: "Runtime B attempts to re-enter runtime A.".to_string(),
14063 ..crate::spec::AgentSpec::default()
14064 };
14065 registry
14066 .register(crate::spawner::SpawnedAgent::from_runtime(
14067 "runtime-a".to_string(),
14068 agent_a,
14069 spec_a,
14070 ))
14071 .await
14072 .unwrap();
14073 registry
14074 .register(crate::spawner::SpawnedAgent::from_runtime(
14075 "runtime-b".to_string(),
14076 agent_b,
14077 spec_b,
14078 ))
14079 .await
14080 .unwrap();
14081 let runtime_a = registry.get("runtime-a").unwrap();
14082 *hooks_b.target.lock() = Some(Arc::downgrade(&runtime_a));
14083
14084 let response = tokio::time::timeout(
14085 std::time::Duration::from_secs(2),
14086 runtime_a.chat("outer concurrent request"),
14087 )
14088 .await
14089 .expect("concurrent orchestration cycle must fail without deadlocking")
14090 .unwrap();
14091
14092 assert_eq!(response.content, "response A");
14093 let child_result = hooks_a
14094 .nested_result
14095 .lock()
14096 .clone()
14097 .expect("runtime A hook must record runtime B completion");
14098 assert_eq!(child_result.unwrap(), "response B");
14099 let cycle_result = hooks_b
14100 .nested_result
14101 .lock()
14102 .clone()
14103 .expect("runtime B hook must record runtime A reentry");
14104 assert!(
14105 cycle_result
14106 .expect_err("runtime A accepted a repeated gate identity")
14107 .contains("reentrant root turn ownership")
14108 );
14109 }
14110
14111 fn skill_clarification_responses() -> Vec<&'static str> {
14114 vec![
14115 r#"{"is_ambiguous":false,"confidence":0.99,"ambiguity_type":null,"reasoning":"top-level clear","what_is_unclear":[],"detected_language":"en"}"#,
14116 "send_report",
14117 r#"{"is_ambiguous":true,"confidence":0.2,"ambiguity_type":"missing_target","reasoning":"target missing","what_is_unclear":["target"],"detected_language":"en"}"#,
14118 r#"{"question":"What should I send?","options":null}"#,
14119 ]
14120 }
14121
14122 #[tokio::test]
14128 async fn test_stream_skill_clarification_memory_matches_blocking() {
14129 let (blocking_agent, _) = state_disambiguation_agent_with_skills(
14130 skill_clarification_responses(),
14131 true,
14132 None,
14133 true,
14134 vec![confirmation_skill()],
14135 );
14136 let blocking = blocking_agent.chat("Send it").await.unwrap();
14137 let blocking_messages = blocking_agent.memory.get_messages(None).await.unwrap();
14138
14139 let (streaming_agent, _) = state_disambiguation_agent_with_skills(
14140 skill_clarification_responses(),
14141 true,
14142 None,
14143 true,
14144 vec![confirmation_skill()],
14145 );
14146 let (content, chunks, streamed) = collect_stream_events(&streaming_agent, "Send it").await;
14147 assert!(!chunks.iter().any(StreamChunk::is_error), "{chunks:?}");
14148 let streamed = streamed.expect("skill clarification must finalize as Final");
14149 let streaming_messages = streaming_agent.memory.get_messages(None).await.unwrap();
14150
14151 assert_eq!(blocking.content, "What should I send?");
14152 assert_eq!(streamed.content, blocking.content);
14153 assert_eq!(content, streamed.content);
14154 assert_eq!(
14155 blocking
14156 .metadata
14157 .as_ref()
14158 .and_then(|m| m.get("disambiguation")),
14159 streamed
14160 .metadata
14161 .as_ref()
14162 .and_then(|m| m.get("disambiguation")),
14163 );
14164 assert_eq!(
14165 streamed
14166 .metadata
14167 .as_ref()
14168 .and_then(|m| m.get("disambiguation"))
14169 .and_then(|d| d.get("status"))
14170 .and_then(Value::as_str),
14171 Some("awaiting_clarification"),
14172 );
14173 let shape = |messages: &[ChatMessage]| {
14174 messages
14175 .iter()
14176 .map(|m| (format!("{:?}", m.role), m.content.clone()))
14177 .collect::<Vec<_>>()
14178 };
14179 assert_eq!(shape(&blocking_messages), shape(&streaming_messages));
14180 assert_eq!(
14181 shape(&streaming_messages),
14182 vec![
14183 ("User".to_string(), "Send it".to_string()),
14184 ("Assistant".to_string(), "What should I send?".to_string()),
14185 ],
14186 );
14187 assert_eq!(
14188 *streaming_agent.pending_skill_id.read(),
14189 Some("send_report".to_string()),
14190 );
14191 }
14192
14193 #[tokio::test]
14195 async fn test_stream_skill_clarification_memory_failure_surfaces_as_error() {
14196 let build = || {
14198 let mut mock = MockLLMProvider::new("skill-clarification");
14199 mock.set_responses(
14200 skill_clarification_responses()
14201 .into_iter()
14202 .map(String::from)
14203 .collect(),
14204 false,
14205 );
14206 AgentBuilder::new()
14207 .system_prompt("Handle requests.")
14208 .llm(Arc::new(mock.clone()))
14209 .llm_alias("router", Arc::new(mock))
14210 .state_machine(disambiguation_state_machine(None, true))
14211 .skills(vec![confirmation_skill()])
14212 .memory(Arc::new(FailingMemory {
14213 messages: parking_lot::RwLock::new(Vec::new()),
14214 fail_on_add: 2,
14215 adds: std::sync::atomic::AtomicUsize::new(0),
14216 }))
14217 .build()
14218 .unwrap()
14219 .with_disambiguation(DisambiguationConfig {
14220 enabled: true,
14221 ..Default::default()
14222 })
14223 };
14224
14225 let blocking = build().chat("Send it").await;
14226 assert!(
14227 blocking.is_err(),
14228 "blocking must surface the failed clarification write: {blocking:?}"
14229 );
14230
14231 let (_, chunks, streamed) = collect_stream_events(&build(), "Send it").await;
14232 assert!(
14233 streamed.is_none(),
14234 "a failed write must not finalize the turn"
14235 );
14236 assert!(
14237 chunks.iter().any(|chunk| matches!(
14238 chunk,
14239 StreamChunk::Error { message } if message.contains("simulated memory failure")
14240 )),
14241 "streaming must surface the failed clarification write: {chunks:?}"
14242 );
14243 }
14244
14245 #[tokio::test]
14247 async fn confirmed_skill_route_executes_exactly_once() {
14248 let (agent, observed) = state_disambiguation_agent_with_skills(
14249 vec![
14250 r#"{"is_ambiguous":false,"confidence":0.99,"ambiguity_type":null,"reasoning":"top-level clear","what_is_unclear":[],"detected_language":"en"}"#,
14251 "send_report",
14252 r#"{"is_ambiguous":true,"confidence":0.2,"ambiguity_type":"missing_target","reasoning":"target missing","what_is_unclear":["target"],"detected_language":"en"}"#,
14253 r#"{"question":"What should I send?","options":null}"#,
14254 r#"{"status":"answered","selected_option":null,"enriched_input":"Send the report to Ada","resolved":{"intent":"send_report"}}"#,
14255 r#"{"question":"Should I send the report to Ada?"}"#,
14256 r#"{"status":"confirmed"}"#,
14257 r#"{"is_ambiguous":false,"confidence":0.99,"ambiguity_type":null,"reasoning":"resolved","what_is_unclear":[],"detected_language":"en"}"#,
14258 "Report skill executed.",
14259 ],
14260 true,
14261 None,
14262 true,
14263 vec![confirmation_skill()],
14264 );
14265
14266 let clarification = agent.chat("Send it").await.unwrap();
14267 assert_eq!(clarification.content, "What should I send?");
14268 assert_eq!(confirmation_skill_call_count(&observed), 0);
14269
14270 let confirmation = agent.chat("The report to Ada").await.unwrap();
14271 assert_eq!(confirmation.content, "Should I send the report to Ada?");
14272 assert_eq!(
14273 confirmation
14274 .metadata
14275 .as_ref()
14276 .and_then(|metadata| metadata.get("disambiguation"))
14277 .and_then(|metadata| metadata.get("status"))
14278 .and_then(Value::as_str),
14279 Some("awaiting_confirmation")
14280 );
14281 assert_eq!(confirmation_skill_call_count(&observed), 0);
14282
14283 let completed = agent.chat("Yes").await.unwrap();
14284 assert_eq!(completed.content, "Report skill executed.");
14285 assert_eq!(confirmation_skill_call_count(&observed), 1);
14286 assert!(agent.pending_skill_id.read().is_none());
14287 let messages = agent.memory.get_messages(None).await.unwrap();
14288 assert!(!messages.iter().any(|message| message.content == "Yes"));
14289 }
14290
14291 #[tokio::test]
14293 async fn confirmed_skill_recheck_preserves_new_clarification_metadata() {
14294 let (agent, observed) = state_disambiguation_agent_with_skills(
14295 vec![
14296 r#"{"is_ambiguous":false,"confidence":0.99,"ambiguity_type":null,"reasoning":"top-level clear","what_is_unclear":[],"detected_language":"en"}"#,
14297 "send_report",
14298 r#"{"is_ambiguous":true,"confidence":0.2,"ambiguity_type":"missing_target","reasoning":"target missing","what_is_unclear":["target"],"detected_language":"en"}"#,
14299 r#"{"question":"What should I send?","options":null}"#,
14300 r#"{"status":"answered","selected_option":null,"enriched_input":"Send the report to Ada","resolved":{"intent":"send_report"}}"#,
14301 r#"{"question":"Should I send the report to Ada?"}"#,
14302 r#"{"status":"confirmed"}"#,
14303 r#"{"is_ambiguous":true,"confidence":0.3,"ambiguity_type":"missing_parameters","reasoning":"timing missing","what_is_unclear":["timing"],"detected_language":"en"}"#,
14304 r#"{"question":"When should I send it?","options":null}"#,
14305 ],
14306 true,
14307 None,
14308 true,
14309 vec![confirmation_skill()],
14310 );
14311
14312 agent.chat("Send it").await.unwrap();
14313 agent.chat("The report to Ada").await.unwrap();
14314 let follow_up = agent.chat("Yes").await.unwrap();
14315
14316 assert_eq!(follow_up.content, "When should I send it?");
14317 let metadata = follow_up
14318 .metadata
14319 .as_ref()
14320 .and_then(|metadata| metadata.get("disambiguation"))
14321 .unwrap();
14322 assert_eq!(
14323 metadata.get("status").and_then(Value::as_str),
14324 Some("awaiting_clarification")
14325 );
14326 assert_eq!(
14327 metadata.get("skill_id").and_then(Value::as_str),
14328 Some("send_report")
14329 );
14330 assert!(metadata.get("detection").is_some());
14331 assert_eq!(confirmation_skill_call_count(&observed), 0);
14332 }
14333
14334 #[tokio::test]
14336 async fn rejected_skill_confirmation_never_executes() {
14337 let (agent, observed) = state_disambiguation_agent_with_skills(
14338 vec![
14339 r#"{"is_ambiguous":false,"confidence":0.99,"ambiguity_type":null,"reasoning":"top-level clear","what_is_unclear":[],"detected_language":"en"}"#,
14340 "send_report",
14341 r#"{"is_ambiguous":true,"confidence":0.2,"ambiguity_type":"missing_target","reasoning":"target missing","what_is_unclear":["target"],"detected_language":"en"}"#,
14342 r#"{"question":"What should I send?","options":null}"#,
14343 r#"{"status":"answered","selected_option":null,"enriched_input":"Send the report to Ada","resolved":{"intent":"send_report"}}"#,
14344 r#"{"question":"Should I send the report to Ada?"}"#,
14345 r#"{"status":"rejected"}"#,
14346 "Confirmation rejected.",
14347 ],
14348 true,
14349 None,
14350 true,
14351 vec![confirmation_skill()],
14352 );
14353
14354 agent.chat("Send it").await.unwrap();
14355 agent.chat("The report to Ada").await.unwrap();
14356 let rejected = agent.chat("No").await.unwrap();
14357
14358 assert_eq!(rejected.content, "Confirmation rejected.");
14359 assert_eq!(confirmation_skill_call_count(&observed), 0);
14360 assert!(agent.pending_skill_id.read().is_none());
14361 }
14362
14363 #[tokio::test]
14365 async fn reset_invalidates_pending_skill_confirmation_before_streaming_input() {
14366 let (agent, observed) = state_disambiguation_agent_with_skills(
14367 vec![
14368 r#"{"is_ambiguous":false,"confidence":0.99,"ambiguity_type":null,"reasoning":"top-level clear","what_is_unclear":[],"detected_language":"en"}"#,
14369 "send_report",
14370 r#"{"is_ambiguous":true,"confidence":0.2,"ambiguity_type":"missing_target","reasoning":"target missing","what_is_unclear":["target"],"detected_language":"en"}"#,
14371 r#"{"question":"What should I send?","options":null}"#,
14372 r#"{"status":"answered","selected_option":null,"enriched_input":"Send the report to Ada","resolved":{"intent":"send_report"}}"#,
14373 r#"{"question":"Should I send the report to Ada?"}"#,
14374 r#"{"is_ambiguous":false,"confidence":0.99,"ambiguity_type":null,"reasoning":"fresh input","what_is_unclear":[],"detected_language":"en"}"#,
14375 "none",
14376 "Fresh response.",
14377 ],
14378 true,
14379 None,
14380 true,
14381 vec![confirmation_skill()],
14382 );
14383
14384 agent.chat("Send it").await.unwrap();
14385 agent.chat("The report to Ada").await.unwrap();
14386 agent.reset().await.unwrap();
14387 assert!(agent.pending_skill_id.read().is_none());
14388 assert!(
14389 !agent
14390 .disambiguation_manager()
14391 .unwrap()
14392 .has_pending_clarification()
14393 .await
14394 );
14395
14396 let mut stream = agent.chat_stream("Yes").await.unwrap();
14397 let mut content = String::new();
14398 while let Some(chunk) = stream.next().await {
14399 match chunk {
14400 StreamChunk::Content { text } => content.push_str(&text),
14401 StreamChunk::Done {} => break,
14402 StreamChunk::Error { message } => panic!("unexpected stream error: {message}"),
14403 _ => {}
14404 }
14405 }
14406
14407 assert_eq!(content, "Fresh response.");
14408 assert_eq!(confirmation_skill_call_count(&observed), 0);
14409 }
14410
14411 #[tokio::test]
14413 async fn trait_reset_clears_pending_skill_confirmation() {
14414 let (agent, _) = state_disambiguation_agent_with_skills(
14415 vec![
14416 r#"{"is_ambiguous":false,"confidence":0.99,"ambiguity_type":null,"reasoning":"top-level clear","what_is_unclear":[],"detected_language":"en"}"#,
14417 "send_report",
14418 r#"{"is_ambiguous":true,"confidence":0.2,"ambiguity_type":"missing_target","reasoning":"target missing","what_is_unclear":["target"],"detected_language":"en"}"#,
14419 r#"{"question":"What should I send?","options":null}"#,
14420 r#"{"status":"answered","selected_option":null,"enriched_input":"Send the report to Ada","resolved":{"intent":"send_report"}}"#,
14421 r#"{"question":"Should I send the report to Ada?"}"#,
14422 ],
14423 true,
14424 None,
14425 true,
14426 vec![confirmation_skill()],
14427 );
14428
14429 agent.chat("Send it").await.unwrap();
14430 agent.chat("The report to Ada").await.unwrap();
14431 <RuntimeAgent as Agent>::reset(&agent).await.unwrap();
14432
14433 assert!(agent.pending_skill_id.read().is_none());
14434 assert!(
14435 !agent
14436 .disambiguation_manager()
14437 .unwrap()
14438 .has_pending_clarification()
14439 .await
14440 );
14441 }
14442
14443 #[tokio::test]
14445 async fn state_change_invalidates_pending_skill_confirmation() {
14446 let (agent, observed) = state_disambiguation_agent_with_skills(
14447 vec![
14448 r#"{"is_ambiguous":false,"confidence":0.99,"ambiguity_type":null,"reasoning":"top-level clear","what_is_unclear":[],"detected_language":"en"}"#,
14449 "send_report",
14450 r#"{"is_ambiguous":true,"confidence":0.2,"ambiguity_type":"missing_target","reasoning":"target missing","what_is_unclear":["target"],"detected_language":"en"}"#,
14451 r#"{"question":"What should I send?","options":null}"#,
14452 r#"{"status":"answered","selected_option":null,"enriched_input":"Send the report to Ada","resolved":{"intent":"send_report"}}"#,
14453 r#"{"question":"Should I send the report to Ada?"}"#,
14454 r#"{"is_ambiguous":false,"confidence":0.99,"ambiguity_type":null,"reasoning":"fresh input","what_is_unclear":[],"detected_language":"en"}"#,
14455 "none",
14456 "Fresh response.",
14457 ],
14458 true,
14459 None,
14460 true,
14461 vec![confirmation_skill()],
14462 );
14463
14464 agent.chat("Send it").await.unwrap();
14465 agent.chat("The report to Ada").await.unwrap();
14466 agent.transition_to("review").await.unwrap();
14467 let cancelled = agent.chat("Yes").await.unwrap();
14468
14469 assert_eq!(cancelled.content, "Fresh response.");
14470 assert_eq!(confirmation_skill_call_count(&observed), 0);
14471 assert!(agent.pending_skill_id.read().is_none());
14472 }
14473
14474 #[tokio::test]
14476 async fn in_flight_confirmation_cannot_redispatch_after_reset() {
14477 let (mut agent, observed) = state_disambiguation_agent_with_skills(
14478 vec![
14479 r#"{"is_ambiguous":false,"confidence":0.99,"ambiguity_type":null,"reasoning":"top-level clear","what_is_unclear":[],"detected_language":"en"}"#,
14480 "send_report",
14481 r#"{"is_ambiguous":true,"confidence":0.2,"ambiguity_type":"missing_target","reasoning":"target missing","what_is_unclear":["target"],"detected_language":"en"}"#,
14482 r#"{"question":"What should I send?","options":null}"#,
14483 r#"{"status":"answered","selected_option":null,"enriched_input":"Send the report to Ada","resolved":{"intent":"send_report"}}"#,
14484 r#"{"question":"Should I send the report to Ada?"}"#,
14485 r#"{"status":"confirmed"}"#,
14486 "Confirmation cancelled.",
14487 ],
14488 true,
14489 None,
14490 true,
14491 vec![confirmation_skill()],
14492 );
14493 let observer = Arc::new(BlockingRuntimeConfirmationObserver::new());
14494 let manager = agent
14495 .disambiguation_manager
14496 .take()
14497 .unwrap()
14498 .with_clarification_observer(observer.clone());
14499 agent.disambiguation_manager = Some(manager);
14500 let agent = Arc::new(agent);
14501
14502 agent.chat("Send it").await.unwrap();
14503 agent.chat("The report to Ada").await.unwrap();
14504
14505 let confirming_agent = Arc::clone(&agent);
14506 let confirmation = tokio::spawn(async move { confirming_agent.chat("Yes").await });
14507 observer.entered.wait().await;
14508 agent.reset().await.unwrap();
14509 observer.release.notify_one();
14510
14511 let response = confirmation.await.unwrap().unwrap();
14512 assert_eq!(response.content, "Confirmation cancelled.");
14513 assert_eq!(confirmation_skill_call_count(&observed), 0);
14514 assert!(agent.pending_skill_id.read().is_none());
14515 }
14516
14517 #[tokio::test]
14519 async fn queued_reset_prevents_stale_confirmation_question_publication() {
14520 let (agent, observed) = state_disambiguation_agent(
14521 vec![
14522 r#"{"is_ambiguous":true,"confidence":0.2,"ambiguity_type":"missing_target","reasoning":"target missing","what_is_unclear":["target"],"detected_language":"en"}"#,
14523 r#"{"question":"What should I send?","options":null}"#,
14524 r#"{"status":"answered","selected_option":null,"enriched_input":"Send the report to Ada","resolved":{"intent":"send_report"}}"#,
14525 r#"{"question":"Should I send the report to Ada?"}"#,
14526 ],
14527 true,
14528 None,
14529 true,
14530 );
14531 let agent = Arc::new(agent);
14532 agent.chat("Send it").await.unwrap();
14533
14534 let admission = agent.disambiguation_admission.write().await;
14535 let (started_tx, started_rx) = tokio::sync::oneshot::channel();
14536 let resetting_agent = Arc::clone(&agent);
14537 let reset = tokio::spawn(async move {
14538 let _ = started_tx.send(());
14539 resetting_agent.reset().await
14540 });
14541 started_rx.await.unwrap();
14542 tokio::task::yield_now().await;
14543
14544 let responding_agent = Arc::clone(&agent);
14545 let response =
14546 tokio::spawn(async move { responding_agent.chat("The report to Ada").await });
14547 tokio::time::timeout(std::time::Duration::from_secs(2), async {
14548 while observed.call_count() < 4 {
14549 tokio::task::yield_now().await;
14550 }
14551 })
14552 .await
14553 .expect("clarification processing must reach terminal publication");
14554 drop(admission);
14555
14556 reset.await.unwrap().unwrap();
14557 let error = response.await.unwrap().unwrap_err();
14558 assert!(error.to_string().contains("ownership changed"));
14559 assert!(
14560 !agent
14561 .disambiguation_manager()
14562 .unwrap()
14563 .has_pending_clarification()
14564 .await
14565 );
14566 assert!(agent.memory.get_messages(None).await.unwrap().is_empty());
14567 }
14568
14569 #[tokio::test]
14571 async fn queued_reset_prevents_stale_skill_clarification_publication() {
14572 let (agent, observed) = state_disambiguation_agent_with_skills(
14573 vec![
14574 r#"{"is_ambiguous":false,"confidence":0.99,"ambiguity_type":null,"reasoning":"top-level clear","what_is_unclear":[],"detected_language":"en"}"#,
14575 "send_report",
14576 r#"{"is_ambiguous":true,"confidence":0.2,"ambiguity_type":"missing_target","reasoning":"target missing","what_is_unclear":["target"],"detected_language":"en"}"#,
14577 r#"{"question":"What should I send?","options":null}"#,
14578 ],
14579 true,
14580 None,
14581 true,
14582 vec![confirmation_skill()],
14583 );
14584 let agent = Arc::new(agent);
14585 let admission = agent.disambiguation_admission.write().await;
14586 let (started_tx, started_rx) = tokio::sync::oneshot::channel();
14587 let resetting_agent = Arc::clone(&agent);
14588 let reset = tokio::spawn(async move {
14589 let _ = started_tx.send(());
14590 resetting_agent.reset().await
14591 });
14592 started_rx.await.unwrap();
14593 tokio::task::yield_now().await;
14594
14595 let responding_agent = Arc::clone(&agent);
14596 let response = tokio::spawn(async move { responding_agent.chat("Send it").await });
14597 tokio::time::timeout(std::time::Duration::from_secs(2), async {
14598 while observed.call_count() < 4 {
14599 tokio::task::yield_now().await;
14600 }
14601 })
14602 .await
14603 .expect("skill clarification must reach terminal publication");
14604 drop(admission);
14605
14606 reset.await.unwrap().unwrap();
14607 let error = response.await.unwrap().unwrap_err();
14608 assert!(error.to_string().contains("ownership changed"));
14609 assert_eq!(confirmation_skill_call_count(&observed), 0);
14610 assert!(agent.pending_skill_id.read().is_none());
14611 assert!(agent.memory.get_messages(None).await.unwrap().is_empty());
14612 }
14613
14614 #[tokio::test]
14616 async fn transition_hook_can_reset_without_admission_deadlock() {
14617 let hooks = Arc::new(ResetOnTransitionHooks {
14618 agent: parking_lot::Mutex::new(None),
14619 invoked: AtomicBool::new(false),
14620 });
14621 let agent = Arc::new(
14622 AgentBuilder::new()
14623 .system_prompt("Test transition hook reentrancy.")
14624 .llm(Arc::new(mock_with_response("done")))
14625 .state_machine(disambiguation_state_machine(None, false))
14626 .build()
14627 .unwrap()
14628 .with_hooks(hooks.clone()),
14629 );
14630 *hooks.agent.lock() = Some(Arc::downgrade(&agent));
14631
14632 let transitioned = tokio::time::timeout(
14633 std::time::Duration::from_secs(2),
14634 agent.apply_transition_target("active", "review", "test transition", None),
14635 )
14636 .await
14637 .expect("transition hook reset must not deadlock")
14638 .unwrap();
14639
14640 assert!(transitioned);
14641 assert!(hooks.invoked.load(Ordering::SeqCst));
14642 assert_eq!(agent.current_state().as_deref(), Some("active"));
14643 }
14644
14645 #[tokio::test]
14647 async fn concurrent_transition_cannot_duplicate_exit_actions() {
14648 let gate = PathMutationGate::new();
14649 let active = ai_agents_state::StateDefinition {
14650 on_exit: vec![StateAction::Tool {
14651 tool: "transition_exit".to_string(),
14652 args: Some(serde_json::json!({"path": "./transition-exit.txt"})),
14653 }],
14654 ..Default::default()
14655 };
14656 let state_machine = Arc::new(
14657 StateMachine::new(ai_agents_state::StateConfig {
14658 initial: "active".to_string(),
14659 states: HashMap::from([
14660 ("active".to_string(), active),
14661 (
14662 "review".to_string(),
14663 ai_agents_state::StateDefinition::default(),
14664 ),
14665 ]),
14666 global_transitions: Vec::new(),
14667 fallback: None,
14668 max_no_transition: None,
14669 regenerate_on_transition: true,
14670 })
14671 .unwrap(),
14672 );
14673 let agent = Arc::new(
14674 AgentBuilder::new()
14675 .system_prompt("Test transition reservation.")
14676 .llm(Arc::new(mock_with_response("done")))
14677 .tool(Arc::new(BlockingPathMutationTool {
14678 id: "transition_exit",
14679 path_fields: vec![ai_agents_core::PathPolicyBinding::write("path")],
14680 gate: gate.clone(),
14681 }))
14682 .state_machine(state_machine)
14683 .build()
14684 .unwrap(),
14685 );
14686
14687 let first_agent = Arc::clone(&agent);
14688 let first = tokio::spawn(async move { first_agent.transition_to("review").await });
14689 tokio::time::timeout(std::time::Duration::from_secs(2), gate.wait_until_entered())
14690 .await
14691 .expect("reserved transition must enter its exit action");
14692
14693 let second = tokio::time::timeout(
14694 std::time::Duration::from_secs(2),
14695 agent.transition_to("review"),
14696 )
14697 .await
14698 .expect("competing transition must fail without waiting for the exit action")
14699 .unwrap_err();
14700 assert!(second.to_string().contains("already in progress"));
14701
14702 gate.release();
14703 first.await.unwrap().unwrap();
14704 assert_eq!(agent.current_state().as_deref(), Some("review"));
14705 }
14706
14707 #[tokio::test]
14709 async fn concurrent_transition_cannot_overtake_enter_actions() {
14710 let gate = PathMutationGate::new();
14711 let review = ai_agents_state::StateDefinition {
14712 on_enter: vec![StateAction::Tool {
14713 tool: "transition_enter".to_string(),
14714 args: Some(serde_json::json!({"path": "./transition-enter.txt"})),
14715 }],
14716 ..Default::default()
14717 };
14718 let state_machine = Arc::new(
14719 StateMachine::new(ai_agents_state::StateConfig {
14720 initial: "active".to_string(),
14721 states: HashMap::from([
14722 (
14723 "active".to_string(),
14724 ai_agents_state::StateDefinition::default(),
14725 ),
14726 ("review".to_string(), review),
14727 ]),
14728 global_transitions: Vec::new(),
14729 fallback: None,
14730 max_no_transition: None,
14731 regenerate_on_transition: true,
14732 })
14733 .unwrap(),
14734 );
14735 let agent = Arc::new(
14736 AgentBuilder::new()
14737 .system_prompt("Test transition lifecycle reservation.")
14738 .llm(Arc::new(mock_with_response("done")))
14739 .tool(Arc::new(BlockingPathMutationTool {
14740 id: "transition_enter",
14741 path_fields: vec![ai_agents_core::PathPolicyBinding::write("path")],
14742 gate: gate.clone(),
14743 }))
14744 .state_machine(state_machine)
14745 .build()
14746 .unwrap(),
14747 );
14748
14749 let first_agent = Arc::clone(&agent);
14750 let first = tokio::spawn(async move { first_agent.transition_to("review").await });
14751 tokio::time::timeout(std::time::Duration::from_secs(2), gate.wait_until_entered())
14752 .await
14753 .expect("committed transition must enter its destination action");
14754
14755 let second = agent.transition_to("active").await.unwrap_err();
14756 assert!(second.to_string().contains("already in progress"));
14757 assert!(agent.reset().await.is_err());
14758
14759 gate.release();
14760 first.await.unwrap().unwrap();
14761 assert_eq!(agent.current_state().as_deref(), Some("review"));
14762 }
14763
14764 #[tokio::test]
14766 async fn same_state_restore_invalidates_pending_skill_confirmation() {
14767 let (agent, observed) = state_disambiguation_agent_with_skills(
14768 vec![
14769 r#"{"is_ambiguous":false,"confidence":0.99,"ambiguity_type":null,"reasoning":"top-level clear","what_is_unclear":[],"detected_language":"en"}"#,
14770 "send_report",
14771 r#"{"is_ambiguous":true,"confidence":0.2,"ambiguity_type":"missing_target","reasoning":"target missing","what_is_unclear":["target"],"detected_language":"en"}"#,
14772 r#"{"question":"What should I send?","options":null}"#,
14773 r#"{"status":"answered","selected_option":null,"enriched_input":"Send the report to Ada","resolved":{"intent":"send_report"}}"#,
14774 r#"{"question":"Should I send the report to Ada?"}"#,
14775 ],
14776 true,
14777 None,
14778 true,
14779 vec![confirmation_skill()],
14780 );
14781
14782 agent.chat("Send it").await.unwrap();
14783 agent.chat("The report to Ada").await.unwrap();
14784 let snapshot = agent.save_state().await.unwrap();
14785 assert_eq!(agent.current_state().as_deref(), Some("active"));
14786
14787 agent.restore_state(snapshot).await.unwrap();
14788
14789 assert_eq!(agent.current_state().as_deref(), Some("active"));
14790 assert!(agent.pending_skill_id.read().is_none());
14791 assert!(
14792 !agent
14793 .disambiguation_manager()
14794 .unwrap()
14795 .has_pending_clarification()
14796 .await
14797 );
14798 assert_eq!(confirmation_skill_call_count(&observed), 0);
14799 }
14800
14801 #[tokio::test]
14803 async fn direct_state_generation_change_invalidates_confirmation() {
14804 let (agent, observed) = state_disambiguation_agent_with_skills(
14805 vec![
14806 r#"{"is_ambiguous":false,"confidence":0.99,"ambiguity_type":null,"reasoning":"top-level clear","what_is_unclear":[],"detected_language":"en"}"#,
14807 "send_report",
14808 r#"{"is_ambiguous":true,"confidence":0.2,"ambiguity_type":"missing_target","reasoning":"target missing","what_is_unclear":["target"],"detected_language":"en"}"#,
14809 r#"{"question":"What should I send?","options":null}"#,
14810 r#"{"status":"answered","selected_option":null,"enriched_input":"Send the report to Ada","resolved":{"intent":"send_report"}}"#,
14811 r#"{"question":"Should I send the report to Ada?"}"#,
14812 "Confirmation cancelled.",
14813 ],
14814 true,
14815 None,
14816 true,
14817 vec![confirmation_skill()],
14818 );
14819
14820 agent.chat("Send it").await.unwrap();
14821 agent.chat("The report to Ada").await.unwrap();
14822 let state_machine = agent.state_machine().unwrap();
14823 state_machine
14824 .transition_to("review", "external test")
14825 .unwrap();
14826 state_machine
14827 .transition_to("active", "external test")
14828 .unwrap();
14829
14830 let response = agent.chat("Yes").await.unwrap();
14831
14832 assert_eq!(response.content, "Confirmation cancelled.");
14833 assert_eq!(confirmation_skill_call_count(&observed), 0);
14834 assert!(agent.pending_skill_id.read().is_none());
14835 }
14836
14837 #[tokio::test]
14838 async fn state_confirmation_does_not_add_a_question_for_clear_input() {
14839 let (agent, observed) = state_disambiguation_agent(
14840 vec![
14841 r#"{"is_ambiguous":false,"confidence":0.99,"ambiguity_type":null,"reasoning":"clear","what_is_unclear":[],"detected_language":"en"}"#,
14842 "Request executed.",
14843 ],
14844 true,
14845 None,
14846 true,
14847 );
14848
14849 let response = agent.chat("Send the report to Ada").await.unwrap();
14850
14851 assert_eq!(response.content, "Request executed.");
14852 assert_eq!(observed.call_count(), 2);
14853 }
14854
14855 #[tokio::test]
14856 async fn state_override_cannot_activate_a_disabled_top_level_manager() {
14857 let (agent, observed) =
14858 state_disambiguation_agent(vec!["Request executed."], false, Some(true), true);
14859
14860 assert!(!agent.has_disambiguation());
14861 let response = agent.chat("Send it").await.unwrap();
14862
14863 assert_eq!(response.content, "Request executed.");
14864 assert_eq!(observed.call_count(), 1);
14865 }
14866
14867 #[tokio::test]
14868 async fn native_required_choice_executes_through_the_shared_tool_path() {
14869 let mut mock = MockLLMProvider::new("native-required");
14870 mock.set_tool_choice(Some(ToolChoice::Required));
14871 let native_call = ToolCall {
14872 id: "provider-call-1".to_string(),
14873 name: "calculator".to_string(),
14874 arguments: serde_json::json!({"expression": "2 + 2"}),
14875 };
14876 let provider_state = ai_agents_core::NativeProviderState::new(
14877 "fixture-exchange-1",
14878 "fixture",
14879 "native-tools",
14880 ai_agents_core::NativeProviderTarget::new("https://fixture.invalid/", "fixture-model")
14881 .unwrap(),
14882 serde_json::json!({
14883 "role": "model",
14884 "parts": [{
14885 "functionCall": {"name": "calculator", "args": {"expression": "2 + 2"}},
14886 "thoughtSignature": "fixture-signature"
14887 }]
14888 }),
14889 vec![ai_agents_core::NativeCallBinding::new("provider-call-1", 0).unwrap()],
14890 )
14891 .unwrap();
14892 mock.add_response(
14893 LLMResponse::new("", FinishReason::ToolCall)
14894 .with_provider_state(provider_state)
14895 .unwrap()
14896 .with_tool_calls(vec![native_call])
14897 .unwrap(),
14898 );
14899 mock.add_response(LLMResponse::new("The answer is 4.", FinishReason::Stop));
14900 let observed = mock.clone();
14901 let agent = AgentBuilder::new()
14902 .system_prompt("Use the calculator when needed.")
14903 .llm(Arc::new(mock))
14904 .tool(Arc::new(CalculatorTool::new()))
14905 .build()
14906 .unwrap();
14907
14908 let response = agent.chat("What is 2 + 2?").await.unwrap();
14909
14910 assert_eq!(response.content, "The answer is 4.");
14911 assert_eq!(
14912 response.tool_calls.as_ref().unwrap()[0].id,
14913 "provider-call-1"
14914 );
14915 let calls = observed.call_history();
14916 assert_eq!(calls.len(), 2);
14917 assert!(matches!(
14918 calls[0].request.as_ref().map(|request| &request.choice),
14919 Some(ToolChoice::Required)
14920 ));
14921 assert!(matches!(
14922 calls[1].request.as_ref().map(|request| &request.choice),
14923 Some(ToolChoice::Auto)
14924 ));
14925 let replay_batch = calls[1]
14926 .messages
14927 .iter()
14928 .find_map(|message| {
14929 ai_agents_core::decode_native_tool_call_markers(&message.content).unwrap()
14930 })
14931 .expect("signed native call marker must be replayed");
14932 assert_eq!(
14933 replay_batch.provider_state().unwrap().exchange_id(),
14934 "fixture-exchange-1"
14935 );
14936 assert!(calls[1].messages.iter().any(|message| {
14937 ai_agents_core::decode_native_tool_result_markers(&message.content)
14938 .is_ok_and(|results| results.is_some())
14939 }));
14940 }
14941
14942 #[tokio::test]
14943 async fn custom_memory_loss_stops_before_signed_tool_execution() {
14944 let mut mock = MockLLMProvider::new("native-custom-memory");
14945 mock.set_tool_choice(Some(ToolChoice::Required));
14946 let call = ToolCall {
14947 id: "provider-call-drop".to_string(),
14948 name: "calculator".to_string(),
14949 arguments: serde_json::json!({"expression": "3 + 4"}),
14950 };
14951 let state = ai_agents_core::NativeProviderState::new(
14952 "fixture-exchange-drop",
14953 "fixture",
14954 "native-tools",
14955 ai_agents_core::NativeProviderTarget::new("https://fixture.invalid/", "fixture-model")
14956 .unwrap(),
14957 serde_json::json!({
14958 "role": "model",
14959 "parts": [{
14960 "functionCall": {"name": "calculator", "args": {"expression": "3 + 4"}},
14961 "thoughtSignature": "fixture-signature-drop"
14962 }]
14963 }),
14964 vec![ai_agents_core::NativeCallBinding::new("provider-call-drop", 0).unwrap()],
14965 )
14966 .unwrap();
14967 mock.add_response(
14968 LLMResponse::new("", FinishReason::ToolCall)
14969 .with_provider_state(state)
14970 .unwrap()
14971 .with_tool_calls(vec![call])
14972 .unwrap(),
14973 );
14974 let agent = AgentBuilder::new()
14975 .system_prompt("Use the calculator.")
14976 .llm(Arc::new(mock))
14977 .memory(Arc::new(DroppingSignedAssistantMemory {
14978 messages: RwLock::new(Vec::new()),
14979 }))
14980 .tool(Arc::new(CalculatorTool::new()))
14981 .build()
14982 .unwrap();
14983
14984 let error = agent.chat("What is 3 + 4?").await.unwrap_err();
14985
14986 assert!(
14987 error
14988 .to_string()
14989 .contains("removed before provider continuation")
14990 );
14991 assert!(agent.tool_call_history.read().is_empty());
14992 }
14993
14994 #[tokio::test]
14995 async fn sequential_signed_history_validates_every_prior_exchange() {
14996 let mut mock = MockLLMProvider::new("native-sequential-memory");
14997 mock.set_tool_choice(Some(ToolChoice::Required));
14998 mock.add_response(signed_calculator_response(
14999 "seq-exchange-1",
15000 "seq-call-1",
15001 "1 + 1",
15002 ));
15003 mock.add_response(signed_calculator_response(
15004 "seq-exchange-2",
15005 "seq-call-2",
15006 "2 + 2",
15007 ));
15008 let agent = AgentBuilder::new()
15009 .system_prompt("Use the calculator sequentially.")
15010 .llm(Arc::new(mock))
15011 .memory(Arc::new(DroppingEarlierSequentialMemory {
15012 messages: RwLock::new(Vec::new()),
15013 signed_seen: std::sync::atomic::AtomicUsize::new(0),
15014 }))
15015 .tool(Arc::new(CalculatorTool::new()))
15016 .build()
15017 .unwrap();
15018
15019 let error = agent.chat("Calculate twice.").await.unwrap_err();
15020
15021 assert!(error.to_string().contains("seq-exchange-1"));
15022 assert_eq!(agent.tool_call_history.read().len(), 1);
15023 }
15024
15025 #[tokio::test]
15026 async fn post_transition_signed_hitl_rejection_stops_before_continuation() {
15027 let mut native = MockLLMProvider::new("post-transition-native");
15028 native.set_tool_choice(Some(ToolChoice::Auto));
15029 let call = ToolCall {
15030 id: "post-transition-call".to_string(),
15031 name: "echo".to_string(),
15032 arguments: serde_json::json!({"message": "hello"}),
15033 };
15034 let state = ai_agents_core::NativeProviderState::new(
15035 "post-transition-exchange",
15036 "fixture",
15037 "native-tools",
15038 ai_agents_core::NativeProviderTarget::new("https://fixture.invalid/", "fixture-model")
15039 .unwrap(),
15040 serde_json::json!({
15041 "role": "model",
15042 "parts": [{
15043 "functionCall": {"name": "echo", "args": {"message": "hello"}},
15044 "thoughtSignature": "post-transition-signature"
15045 }]
15046 }),
15047 vec![ai_agents_core::NativeCallBinding::new("post-transition-call", 0).unwrap()],
15048 )
15049 .unwrap();
15050 native.add_response(
15051 LLMResponse::new("", FinishReason::ToolCall)
15052 .with_provider_state(state)
15053 .unwrap()
15054 .with_tool_calls(vec![call])
15055 .unwrap(),
15056 );
15057 let observed_native = native.clone();
15058 let yaml = r#"
15059name: PostTransitionNativeReject
15060system_prompt: test
15061tools: [echo]
15062hitl:
15063 tools:
15064 echo:
15065 require_approval: true
15066states:
15067 initial: intake
15068 states:
15069 intake:
15070 prompt: intake
15071 transitions:
15072 - to: active
15073 guard:
15074 context:
15075 route:
15076 eq: active
15077 active:
15078 prompt: active
15079 llm: native
15080"#;
15081 let agent = AgentBuilder::from_yaml(yaml)
15082 .unwrap()
15083 .llm(Arc::new(mock_with_response("stale intake response")))
15084 .llm_alias("native", Arc::new(native))
15085 .auto_configure_features()
15086 .unwrap()
15087 .build()
15088 .unwrap();
15089 agent
15090 .set_context("route", serde_json::json!("active"))
15091 .unwrap();
15092
15093 let error = agent.chat("move to active").await.unwrap_err();
15094
15095 assert!(matches!(error, AgentError::HITLRejected(_)));
15096 assert_eq!(observed_native.call_count(), 1);
15097 }
15098
15099 #[test]
15100 fn runtime_overflow_removes_a_past_signed_user_turn_as_one_prefix() {
15101 let call = ToolCall {
15102 id: "overflow-call".to_string(),
15103 name: "calculator".to_string(),
15104 arguments: serde_json::json!({"expression": "1 + 1"}),
15105 };
15106 let state = ai_agents_core::NativeProviderState::new(
15107 "overflow-exchange",
15108 "google",
15109 "generateContent",
15110 ai_agents_core::NativeProviderTarget::new("https://example.invalid/", "gemini-3")
15111 .unwrap(),
15112 serde_json::json!({
15113 "role": "model",
15114 "parts": [{
15115 "functionCall": {"name": "calculator", "args": {"expression": "1 + 1"}},
15116 "thoughtSignature": "overflow-signature"
15117 }]
15118 }),
15119 vec![ai_agents_core::NativeCallBinding::new("overflow-call", 0).unwrap()],
15120 )
15121 .unwrap();
15122 let call_marker = ai_agents_core::encode_native_tool_call_markers(
15123 std::slice::from_ref(&call),
15124 Some(&state),
15125 )
15126 .unwrap();
15127 let result_marker = ai_agents_core::encode_native_tool_result_marker(
15128 &call,
15129 serde_json::json!({"result": 2}),
15130 )
15131 .unwrap();
15132 let history = vec![
15133 ChatMessage::user("old question"),
15134 ChatMessage::assistant(call_marker),
15135 ChatMessage::function("calculator", result_marker),
15136 ChatMessage::assistant("old answer"),
15137 ChatMessage::user("new question"),
15138 ];
15139
15140 let removable = RuntimeAgent::native_safe_prefix_at_least(&history, 1).unwrap();
15141
15142 assert_eq!(removable, 4);
15143 }
15144
15145 #[test]
15146 fn auxiliary_projection_does_not_interpret_user_marker_text() {
15147 let user_text = serde_json::json!({
15148 "_ai_agents_native_tool_call": true,
15149 "id": "",
15150 "tool": "user-data",
15151 "arguments": {}
15152 })
15153 .to_string();
15154
15155 let projected =
15156 RuntimeAgent::readable_native_messages(vec![ChatMessage::user(&user_text)]).unwrap();
15157
15158 assert_eq!(projected[0].content, user_text);
15159 }
15160
15161 #[tokio::test]
15162 async fn terminal_provider_history_error_skips_retry_and_static_fallback() {
15163 let calls = Arc::new(std::sync::atomic::AtomicU32::new(0));
15164 let recovery = RecoveryManager::new(ai_agents_recovery::ErrorRecoveryConfig {
15165 default: ai_agents_recovery::RetryConfig {
15166 max_retries: 3,
15167 ..Default::default()
15168 },
15169 llm: ai_agents_recovery::LLMRecoveryConfig {
15170 on_failure: LLMFailureAction::FallbackResponse {
15171 message: "must not be returned".to_string(),
15172 },
15173 ..Default::default()
15174 },
15175 ..Default::default()
15176 });
15177 let agent = AgentBuilder::new()
15178 .system_prompt("Reject corrupted native history.")
15179 .llm(Arc::new(TerminalHistoryProvider {
15180 calls: Arc::clone(&calls),
15181 }))
15182 .recovery_manager(recovery)
15183 .build()
15184 .unwrap();
15185
15186 let error = agent.chat("continue").await.unwrap_err();
15187
15188 assert!(
15189 error
15190 .to_string()
15191 .contains("native history integrity failure")
15192 );
15193 assert_eq!(calls.load(Ordering::SeqCst), 1);
15194 }
15195
15196 #[tokio::test]
15197 async fn prompt_fallback_uses_one_corrective_retry() {
15198 let mut mock = MockLLMProvider::new("prompt-required");
15199 mock.set_tool_choice(Some(ToolChoice::Required));
15200 mock.set_native_tool_support(false);
15201 mock.set_responses(
15202 vec![
15203 "I can calculate that.".to_string(),
15204 r#"{"tool":"calculator","arguments":{"expression":"2 + 2"}}"#.to_string(),
15205 "The answer is 4.".to_string(),
15206 ],
15207 false,
15208 );
15209 let observed = mock.clone();
15210 let agent = AgentBuilder::new()
15211 .system_prompt("Use tools.")
15212 .llm(Arc::new(mock))
15213 .tool(Arc::new(CalculatorTool::new()))
15214 .build()
15215 .unwrap();
15216
15217 let response = agent.chat("What is 2 + 2?").await.unwrap();
15218
15219 assert_eq!(response.content, "The answer is 4.");
15220 assert_eq!(observed.call_count(), 3);
15221 let corrective = &observed.call_history()[1].messages;
15222 assert!(
15223 corrective
15224 .last()
15225 .unwrap()
15226 .content
15227 .contains("previous response")
15228 );
15229 }
15230
15231 #[tokio::test]
15232 async fn prompt_fallback_fails_after_one_noncompliant_retry() {
15233 let mut mock = MockLLMProvider::new("prompt-required-failure");
15234 mock.set_tool_choice(Some(ToolChoice::Required));
15235 mock.set_native_tool_support(false);
15236 mock.set_responses(
15237 vec!["No tool.".to_string(), "Still no tool.".to_string()],
15238 false,
15239 );
15240 let observed = mock.clone();
15241 let agent = AgentBuilder::new()
15242 .system_prompt("Use tools.")
15243 .llm(Arc::new(mock))
15244 .tool(Arc::new(CalculatorTool::new()))
15245 .build()
15246 .unwrap();
15247
15248 let error = agent.chat("What is 2 + 2?").await.unwrap_err();
15249
15250 assert!(error.to_string().contains("one corrective retry"));
15251 assert_eq!(observed.call_count(), 2);
15252 }
15253
15254 #[tokio::test]
15255 async fn specific_choice_cannot_widen_the_effective_grant() {
15256 let mut mock = MockLLMProvider::new("specific-outside-grant");
15257 mock.set_tool_choice(Some(ToolChoice::Specific("random".to_string())));
15258 let observed = mock.clone();
15259 let agent = AgentBuilder::new()
15260 .system_prompt("Use tools.")
15261 .llm(Arc::new(mock))
15262 .tool(Arc::new(CalculatorTool::new()))
15263 .build()
15264 .unwrap();
15265
15266 let error = agent.chat("Generate a value.").await.unwrap_err();
15267
15268 assert!(error.to_string().contains("is not registered"));
15269 assert_eq!(observed.call_count(), 0);
15270 }
15271
15272 #[tokio::test]
15273 async fn none_choice_exposes_no_tool_protocol() {
15274 let mut mock = MockLLMProvider::new("no-tools");
15275 mock.set_tool_choice(Some(ToolChoice::None));
15276 mock.set_response(r#"{"tool":"calculator","arguments":{"expression":"2 + 2"}}"#);
15277 let observed = mock.clone();
15278 let agent = AgentBuilder::new()
15279 .system_prompt("Answer directly.")
15280 .llm(Arc::new(mock))
15281 .tool(Arc::new(CalculatorTool::new()))
15282 .build()
15283 .unwrap();
15284
15285 let response = agent.chat("Hello").await.unwrap();
15286
15287 assert!(response.tool_calls.is_none());
15288 assert_eq!(observed.call_count(), 1);
15289 let call = observed.last_call().unwrap();
15290 assert!(call.request.is_none());
15291 assert!(
15292 call.messages
15293 .iter()
15294 .all(|message| !message.content.contains("Available tools:"))
15295 );
15296 }
15297
15298 struct RuntimeStorage {
15299 capabilities: Box<[StorageCapability]>,
15300 snapshots: RwLock<HashMap<String, AgentSnapshot>>,
15301 metadata: RwLock<HashMap<String, ai_agents_core::SessionMetadata>>,
15302 metadata_save_calls: AtomicU64,
15303 metadata_load_calls: AtomicU64,
15304 fail_metadata_save: AtomicBool,
15305 fail_metadata_load: AtomicBool,
15306 }
15307
15308 impl RuntimeStorage {
15309 fn new(capabilities: impl IntoIterator<Item = StorageCapability>) -> Self {
15310 Self {
15311 capabilities: capabilities.into_iter().collect(),
15312 snapshots: RwLock::new(HashMap::new()),
15313 metadata: RwLock::new(HashMap::new()),
15314 metadata_save_calls: AtomicU64::new(0),
15315 metadata_load_calls: AtomicU64::new(0),
15316 fail_metadata_save: AtomicBool::new(false),
15317 fail_metadata_load: AtomicBool::new(false),
15318 }
15319 }
15320 }
15321
15322 #[async_trait]
15323 impl AgentStorage for RuntimeStorage {
15324 fn supports(&self, capability: StorageCapability) -> bool {
15325 self.capabilities.contains(&capability)
15326 }
15327
15328 async fn save(&self, session_id: &str, snapshot: &AgentSnapshot) -> Result<()> {
15329 self.snapshots
15330 .write()
15331 .insert(session_id.to_string(), snapshot.clone());
15332 Ok(())
15333 }
15334
15335 async fn load(&self, session_id: &str) -> Result<Option<AgentSnapshot>> {
15336 Ok(self.snapshots.read().get(session_id).cloned())
15337 }
15338
15339 async fn delete(&self, session_id: &str) -> Result<()> {
15340 self.snapshots.write().remove(session_id);
15341 Ok(())
15342 }
15343
15344 async fn list_sessions(&self) -> Result<Vec<String>> {
15345 Ok(self.snapshots.read().keys().cloned().collect())
15346 }
15347
15348 async fn save_snapshot_with_metadata(
15349 &self,
15350 session_id: &str,
15351 snapshot: &AgentSnapshot,
15352 metadata: &ai_agents_core::SessionMetadata,
15353 ) -> Result<()> {
15354 self.metadata_save_calls.fetch_add(1, Ordering::SeqCst);
15355 if self.fail_metadata_save.load(Ordering::SeqCst) {
15356 return Err(AgentError::Persistence("metadata save failed".into()));
15357 }
15358 self.snapshots
15359 .write()
15360 .insert(session_id.to_string(), snapshot.clone());
15361 self.metadata
15362 .write()
15363 .insert(session_id.to_string(), metadata.clone());
15364 Ok(())
15365 }
15366
15367 async fn save_metadata(
15368 &self,
15369 session_id: &str,
15370 metadata: &ai_agents_core::SessionMetadata,
15371 ) -> Result<()> {
15372 self.metadata_save_calls.fetch_add(1, Ordering::SeqCst);
15373 if self.fail_metadata_save.load(Ordering::SeqCst) {
15374 return Err(AgentError::Persistence("metadata save failed".into()));
15375 }
15376 self.metadata
15377 .write()
15378 .insert(session_id.to_string(), metadata.clone());
15379 Ok(())
15380 }
15381
15382 async fn load_metadata(
15383 &self,
15384 session_id: &str,
15385 ) -> Result<Option<ai_agents_core::SessionMetadata>> {
15386 self.metadata_load_calls.fetch_add(1, Ordering::SeqCst);
15387 if self.fail_metadata_load.load(Ordering::SeqCst) {
15388 return Err(AgentError::Persistence("metadata load failed".into()));
15389 }
15390 Ok(self.metadata.read().get(session_id).cloned())
15391 }
15392 }
15393
15394 fn runtime_storage_agent() -> RuntimeAgent {
15395 AgentBuilder::new()
15396 .system_prompt("Test runtime storage integration.")
15397 .llm(Arc::new(mock_with_response("done")))
15398 .build()
15399 .unwrap()
15400 }
15401
15402 fn restore_spec(id: &str) -> crate::spec::AgentSpec {
15403 crate::spec::AgentSpec {
15404 name: id.to_string(),
15405 system_prompt: format!("Restore child {id}."),
15406 ..crate::spec::AgentSpec::default()
15407 }
15408 }
15409
15410 fn restore_entry(id: &str) -> ai_agents_core::SpawnedAgentEntry {
15411 ai_agents_core::SpawnedAgentEntry {
15412 id: id.to_string(),
15413 name: id.to_string(),
15414 spec_yaml: serde_yaml::to_string(&restore_spec(id)).unwrap(),
15415 }
15416 }
15417
15418 fn restore_spawner(
15419 storage: Arc<RuntimeStorage>,
15420 max_agents: usize,
15421 ) -> (
15422 Arc<crate::spawner::AgentSpawner>,
15423 Arc<crate::spawner::AgentRegistry>,
15424 ) {
15425 let mut llms = LLMRegistry::new();
15426 llms.register("default", Arc::new(mock_with_response("done")));
15427 (
15428 Arc::new(
15429 crate::spawner::AgentSpawner::new()
15430 .with_shared_llms(llms)
15431 .with_shared_storage(storage)
15432 .with_max_agents(max_agents),
15433 ),
15434 Arc::new(crate::spawner::AgentRegistry::new()),
15435 )
15436 }
15437
15438 async fn save_restore_target(
15439 parent: &RuntimeAgent,
15440 storage: &RuntimeStorage,
15441 session_id: &str,
15442 entries: Vec<ai_agents_core::SpawnedAgentEntry>,
15443 ) {
15444 let mut snapshot = parent.save_state().await.unwrap();
15445 snapshot.spawned_agents = Some(entries);
15446 storage.save(session_id, &snapshot).await.unwrap();
15447 storage
15448 .save_metadata(session_id, &ai_agents_core::SessionMetadata::default())
15449 .await
15450 .unwrap();
15451 }
15452
15453 #[tokio::test]
15454 async fn storage_init_requires_storage_for_actor_facts() {
15455 let facts = ai_agents_facts::FactsConfig {
15456 enabled: true,
15457 ..Default::default()
15458 };
15459 let agent = runtime_storage_agent().with_facts_config(None, Some(facts));
15460
15461 let error = agent.init_storage().await.unwrap_err();
15462 assert!(matches!(
15463 error,
15464 AgentError::Config(message)
15465 if message.contains("actor facts or actor memory")
15466 && message.contains("none is configured or injected")
15467 ));
15468 }
15469
15470 #[tokio::test]
15471 async fn storage_init_validates_actor_facts_capability() {
15472 let storage = Arc::new(RuntimeStorage::new([StorageCapability::Snapshot]));
15473 let actor_memory = ai_agents_facts::ActorMemoryConfig {
15474 enabled: true,
15475 ..Default::default()
15476 };
15477 let agent = runtime_storage_agent()
15478 .with_storage(storage)
15479 .with_facts_config(Some(actor_memory), None);
15480
15481 assert!(matches!(
15482 agent.init_storage().await,
15483 Err(AgentError::UnsupportedStorageCapability(
15484 StorageCapability::ActorFacts
15485 ))
15486 ));
15487 }
15488
15489 #[tokio::test]
15490 async fn blocking_chat_rejects_unsupported_required_storage() {
15491 let storage = Arc::new(RuntimeStorage::new([StorageCapability::Snapshot]));
15492 let facts = ai_agents_facts::FactsConfig {
15493 enabled: true,
15494 ..Default::default()
15495 };
15496 let agent = runtime_storage_agent()
15497 .with_storage(storage)
15498 .with_facts_config(None, Some(facts));
15499
15500 assert!(matches!(
15501 agent.chat("hello").await,
15502 Err(AgentError::UnsupportedStorageCapability(
15503 StorageCapability::ActorFacts
15504 ))
15505 ));
15506 }
15507
15508 #[tokio::test]
15509 async fn streaming_chat_rejects_unsupported_required_storage_before_stream_creation() {
15510 let storage = Arc::new(RuntimeStorage::new([StorageCapability::Snapshot]));
15511 let config = ai_agents_relationships::RelationshipConfig {
15512 enabled: true,
15513 ..Default::default()
15514 };
15515 let manager = Arc::new(RelationshipManager::from_config(config).unwrap());
15516 let agent = runtime_storage_agent()
15517 .with_storage(storage)
15518 .with_relationships(manager);
15519
15520 assert!(matches!(
15521 agent.chat_stream("hello").await,
15522 Err(AgentError::UnsupportedStorageCapability(
15523 StorageCapability::ActorRelationships
15524 ))
15525 ));
15526 }
15527
15528 #[tokio::test]
15529 async fn storage_init_completes_facts_for_injected_storage() {
15530 let storage = Arc::new(RuntimeStorage::new([
15531 StorageCapability::Snapshot,
15532 StorageCapability::ActorFacts,
15533 ]));
15534 let facts = ai_agents_facts::FactsConfig {
15535 enabled: true,
15536 ..Default::default()
15537 };
15538 let agent = runtime_storage_agent()
15539 .with_storage(storage)
15540 .with_facts_config(None, Some(facts));
15541
15542 agent.init_storage().await.unwrap();
15543 assert!(agent.fact_store().is_some());
15544 }
15545
15546 #[tokio::test]
15547 async fn storage_init_requires_storage_for_persistent_relationships() {
15548 let config = ai_agents_relationships::RelationshipConfig {
15549 enabled: true,
15550 ..Default::default()
15551 };
15552 let manager = Arc::new(RelationshipManager::from_config(config).unwrap());
15553 let agent = runtime_storage_agent().with_relationships(manager);
15554
15555 let error = agent.init_storage().await.unwrap_err();
15556 assert!(matches!(
15557 error,
15558 AgentError::Config(message)
15559 if message.contains("persistent relationships")
15560 && message.contains("none is configured or injected")
15561 ));
15562 }
15563
15564 #[tokio::test]
15565 async fn storage_init_validates_persistent_relationships_capability() {
15566 let storage = Arc::new(RuntimeStorage::new([StorageCapability::Snapshot]));
15567 let config = ai_agents_relationships::RelationshipConfig {
15568 enabled: true,
15569 ..Default::default()
15570 };
15571 let manager = Arc::new(RelationshipManager::from_config(config).unwrap());
15572 let agent = runtime_storage_agent()
15573 .with_storage(storage)
15574 .with_relationships(manager);
15575
15576 assert!(matches!(
15577 agent.init_storage().await,
15578 Err(AgentError::UnsupportedStorageCapability(
15579 StorageCapability::ActorRelationships
15580 ))
15581 ));
15582 }
15583
15584 #[tokio::test]
15585 async fn session_restore_updates_identity_and_clears_stale_actor_binding() {
15586 let storage = Arc::new(RuntimeStorage::new([
15587 StorageCapability::Snapshot,
15588 StorageCapability::SessionMetadata,
15589 ]));
15590 let agent = runtime_storage_agent().with_storage(storage.clone());
15591 agent.set_actor_id("old-actor").unwrap();
15592 agent.save_session("old").await.unwrap();
15593 storage
15594 .save("target", &agent.save_state().await.unwrap())
15595 .await
15596 .unwrap();
15597 storage
15598 .save_metadata("target", &ai_agents_core::SessionMetadata::default())
15599 .await
15600 .unwrap();
15601
15602 assert!(agent.load_session("target").await.unwrap());
15603
15604 assert_eq!(agent.current_session_id.read().as_deref(), Some("target"));
15605 assert_eq!(agent.actor_id(), None);
15606 }
15607
15608 #[tokio::test]
15609 async fn complete_restore_reconciles_growth_shrink_and_empty_topologies() {
15610 let storage = Arc::new(RuntimeStorage::new([
15611 StorageCapability::Snapshot,
15612 StorageCapability::SessionMetadata,
15613 ]));
15614 let (spawner, registry) = restore_spawner(storage.clone(), 3);
15615 let parent = runtime_storage_agent()
15616 .with_storage(storage.clone())
15617 .with_spawner_handles(Arc::clone(&spawner), Arc::clone(®istry));
15618
15619 for id in ["a", "b"] {
15620 let spawned = spawner
15621 .spawn_with_id(id.to_string(), restore_spec(id))
15622 .await
15623 .unwrap();
15624 spawned.agent.save_session("grow").await.unwrap();
15625 registry.register(spawned).await.unwrap();
15626 }
15627 let staged_c = crate::spawner::storage::NamespacedStorage::new(storage.clone(), "c");
15628 staged_c
15629 .save("grow", &AgentSnapshot::new("c".into()))
15630 .await
15631 .unwrap();
15632 staged_c
15633 .save_metadata("grow", &ai_agents_core::SessionMetadata::default())
15634 .await
15635 .unwrap();
15636 save_restore_target(
15637 &parent,
15638 storage.as_ref(),
15639 "grow",
15640 vec![restore_entry("a"), restore_entry("b"), restore_entry("c")],
15641 )
15642 .await;
15643
15644 assert_eq!(parent.restore_session_full("grow").await.unwrap(), 3);
15645 assert_eq!(registry.count(), 3);
15646 assert!(registry.contains("c"));
15647 assert_eq!(spawner.spawned_count(), 3);
15648
15649 for id in ["a", "b"] {
15650 registry
15651 .get(id)
15652 .unwrap()
15653 .save_session("shrink")
15654 .await
15655 .unwrap();
15656 }
15657 save_restore_target(
15658 &parent,
15659 storage.as_ref(),
15660 "shrink",
15661 vec![restore_entry("a"), restore_entry("b")],
15662 )
15663 .await;
15664
15665 assert_eq!(parent.restore_session_full("shrink").await.unwrap(), 2);
15666 assert_eq!(registry.count(), 2);
15667 assert!(!registry.contains("c"));
15668 assert_eq!(spawner.spawned_count(), 2);
15669
15670 save_restore_target(&parent, storage.as_ref(), "empty", Vec::new()).await;
15671
15672 assert_eq!(parent.restore_session_full("empty").await.unwrap(), 0);
15673 assert_eq!(registry.count(), 0);
15674 assert_eq!(spawner.spawned_count(), 0);
15675 assert_eq!(parent.current_session_id.read().as_deref(), Some("empty"));
15676 }
15677
15678 #[tokio::test]
15679 async fn storage_session_metadata_is_called_only_when_advertised() {
15680 let storage = Arc::new(RuntimeStorage::new([StorageCapability::Snapshot]));
15681 storage.fail_metadata_save.store(true, Ordering::SeqCst);
15682 storage.fail_metadata_load.store(true, Ordering::SeqCst);
15683 let agent = runtime_storage_agent().with_storage(storage.clone());
15684
15685 agent.save_session("session").await.unwrap();
15686 assert!(agent.load_session("session").await.unwrap());
15687 assert_eq!(storage.metadata_save_calls.load(Ordering::SeqCst), 0);
15688 assert_eq!(storage.metadata_load_calls.load(Ordering::SeqCst), 0);
15689 }
15690
15691 #[cfg(feature = "sqlite")]
15692 #[tokio::test]
15693 async fn sqlite_runtime_save_filter_reopen_and_reload_stay_consistent() {
15694 let directory =
15695 std::env::temp_dir().join(format!("ai-agents-runtime-sqlite-{}", uuid::Uuid::new_v4()));
15696 let path = directory.join("sessions.sqlite");
15697 let path_string = path.to_string_lossy().into_owned();
15698 let storage = Arc::new(
15699 ai_agents_storage::SqliteStorage::new(&path_string)
15700 .await
15701 .unwrap(),
15702 );
15703 let agent = runtime_storage_agent().with_storage(storage.clone());
15704 agent.set_session_metadata(ai_agents_core::SessionMetadata {
15705 tags: vec!["initial".into()],
15706 ..Default::default()
15707 });
15708 agent.chat("persist this turn").await.unwrap();
15709 agent.save_session("session").await.unwrap();
15710
15711 agent.set_session_metadata(ai_agents_core::SessionMetadata {
15712 tags: vec!["updated".into()],
15713 ..Default::default()
15714 });
15715 agent.save_session("session").await.unwrap();
15716 assert!(
15717 agent
15718 .list_sessions_filtered(&ai_agents_core::SessionFilter {
15719 tags: Some(vec!["initial".into()]),
15720 ..Default::default()
15721 })
15722 .await
15723 .unwrap()
15724 .is_empty()
15725 );
15726 assert_eq!(
15727 agent
15728 .list_sessions_filtered(&ai_agents_core::SessionFilter {
15729 tags: Some(vec!["updated".into()]),
15730 ..Default::default()
15731 })
15732 .await
15733 .unwrap()
15734 .len(),
15735 1
15736 );
15737 drop(agent);
15738 storage.close().await;
15739 drop(storage);
15740
15741 let reopened_storage = Arc::new(
15742 ai_agents_storage::SqliteStorage::new(&path_string)
15743 .await
15744 .unwrap(),
15745 );
15746 let restored = runtime_storage_agent().with_storage(reopened_storage.clone());
15747 assert!(restored.load_session("session").await.unwrap());
15748 assert_eq!(restored.session_metadata().tags, vec!["updated"]);
15749 assert_eq!(
15750 restored.current_session_id.read().as_deref(),
15751 Some("session")
15752 );
15753 assert!(restored.save_state().await.unwrap().memory.messages.len() >= 2);
15754 assert_eq!(
15755 restored
15756 .list_sessions_filtered(&ai_agents_core::SessionFilter {
15757 tags: Some(vec!["updated".into()]),
15758 ..Default::default()
15759 })
15760 .await
15761 .unwrap()
15762 .len(),
15763 1
15764 );
15765
15766 drop(restored);
15767 reopened_storage.close().await;
15768 drop(reopened_storage);
15769 crate::remove_sqlite_test_directory(&directory)
15770 .await
15771 .unwrap();
15772 }
15773
15774 #[tokio::test]
15775 async fn storage_session_metadata_backend_failures_propagate() {
15776 let storage = Arc::new(RuntimeStorage::new([
15777 StorageCapability::Snapshot,
15778 StorageCapability::SessionMetadata,
15779 ]));
15780 let agent = runtime_storage_agent().with_storage(storage.clone());
15781
15782 agent.save_session("session").await.unwrap();
15783 storage
15784 .save("target", &agent.save_state().await.unwrap())
15785 .await
15786 .unwrap();
15787 storage.fail_metadata_load.store(true, Ordering::SeqCst);
15788 assert!(matches!(
15789 agent.load_session("target").await,
15790 Err(AgentError::Persistence(message)) if message == "metadata load failed"
15791 ));
15792 assert_eq!(agent.current_session_id.read().as_deref(), Some("session"));
15793
15794 storage.fail_metadata_save.store(true, Ordering::SeqCst);
15795 assert!(matches!(
15796 agent.save_session("session").await,
15797 Err(AgentError::Persistence(message)) if message == "metadata save failed"
15798 ));
15799 }
15800
15801 struct ProviderFutureDropSignal {
15802 dropped: Arc<AtomicBool>,
15803 }
15804
15805 impl Drop for ProviderFutureDropSignal {
15806 fn drop(&mut self) {
15807 self.dropped.store(true, Ordering::SeqCst);
15808 }
15809 }
15810
15811 struct BufferedLockingProvider {
15812 lock: Arc<tokio::sync::Mutex<()>>,
15813 stream_started: Arc<tokio::sync::Notify>,
15814 stream_dropped: Arc<AtomicBool>,
15815 committed_after_drop: Arc<AtomicBool>,
15816 }
15817
15818 #[async_trait]
15819 impl LLMProvider for BufferedLockingProvider {
15820 async fn complete(
15821 &self,
15822 _messages: &[ChatMessage],
15823 _config: Option<&LLMConfig>,
15824 ) -> std::result::Result<LLMResponse, LLMError> {
15825 let _guard = self.lock.lock().await;
15826 self.committed_after_drop
15827 .store(self.stream_dropped.load(Ordering::SeqCst), Ordering::SeqCst);
15828 Ok(LLMResponse::new(
15829 "Committed technical response.",
15830 FinishReason::Stop,
15831 ))
15832 }
15833
15834 async fn complete_stream(
15835 &self,
15836 _messages: &[ChatMessage],
15837 _config: Option<&LLMConfig>,
15838 ) -> std::result::Result<
15839 Box<dyn Stream<Item = std::result::Result<LLMChunk, LLMError>> + Unpin + Send>,
15840 LLMError,
15841 > {
15842 let _guard = self.lock.lock().await;
15843 let _drop_signal = ProviderFutureDropSignal {
15844 dropped: Arc::clone(&self.stream_dropped),
15845 };
15846 self.stream_started.notify_one();
15847 std::future::pending().await
15848 }
15849
15850 fn provider_name(&self) -> &str {
15851 "buffered-locking"
15852 }
15853
15854 fn supports(&self, _feature: LLMFeature) -> bool {
15855 false
15856 }
15857 }
15858
15859 struct PendingDropStream {
15860 dropped: Arc<AtomicBool>,
15861 dropped_notify: Arc<tokio::sync::Notify>,
15862 }
15863
15864 impl Stream for PendingDropStream {
15865 type Item = std::result::Result<LLMChunk, LLMError>;
15866
15867 fn poll_next(
15868 self: Pin<&mut Self>,
15869 _cx: &mut std::task::Context<'_>,
15870 ) -> std::task::Poll<Option<Self::Item>> {
15871 std::task::Poll::Pending
15872 }
15873 }
15874
15875 impl Drop for PendingDropStream {
15876 fn drop(&mut self) {
15877 self.dropped.store(true, Ordering::SeqCst);
15878 self.dropped_notify.notify_one();
15879 }
15880 }
15881
15882 struct EstablishedStreamProvider {
15883 stream_started: Arc<tokio::sync::Notify>,
15884 stream_dropped: Arc<AtomicBool>,
15885 stream_dropped_notify: Arc<tokio::sync::Notify>,
15886 committed_after_drop: Arc<AtomicBool>,
15887 }
15888
15889 #[async_trait]
15890 impl LLMProvider for EstablishedStreamProvider {
15891 async fn complete(
15892 &self,
15893 _messages: &[ChatMessage],
15894 _config: Option<&LLMConfig>,
15895 ) -> std::result::Result<LLMResponse, LLMError> {
15896 if !self.stream_dropped.load(Ordering::SeqCst) {
15897 self.stream_dropped_notify.notified().await;
15898 }
15899 self.committed_after_drop
15900 .store(self.stream_dropped.load(Ordering::SeqCst), Ordering::SeqCst);
15901 Ok(LLMResponse::new(
15902 "Committed technical response.",
15903 FinishReason::Stop,
15904 ))
15905 }
15906
15907 async fn complete_stream(
15908 &self,
15909 _messages: &[ChatMessage],
15910 _config: Option<&LLMConfig>,
15911 ) -> std::result::Result<
15912 Box<dyn Stream<Item = std::result::Result<LLMChunk, LLMError>> + Unpin + Send>,
15913 LLMError,
15914 > {
15915 self.stream_started.notify_one();
15916 Ok(Box::new(PendingDropStream {
15917 dropped: Arc::clone(&self.stream_dropped),
15918 dropped_notify: Arc::clone(&self.stream_dropped_notify),
15919 }))
15920 }
15921
15922 fn provider_name(&self) -> &str {
15923 "established-stream"
15924 }
15925
15926 fn supports(&self, _feature: LLMFeature) -> bool {
15927 false
15928 }
15929 }
15930
15931 struct FirstCallLockingProvider {
15932 lock: Arc<tokio::sync::Mutex<()>>,
15933 first_started: Arc<tokio::sync::Notify>,
15934 first_dropped: Arc<AtomicBool>,
15935 committed_after_drop: Arc<AtomicBool>,
15936 calls: AtomicU64,
15937 }
15938
15939 #[async_trait]
15940 impl LLMProvider for FirstCallLockingProvider {
15941 async fn complete(
15942 &self,
15943 _messages: &[ChatMessage],
15944 _config: Option<&LLMConfig>,
15945 ) -> std::result::Result<LLMResponse, LLMError> {
15946 let _guard = self.lock.lock().await;
15947 let call = self.calls.fetch_add(1, Ordering::SeqCst);
15948 if call == 0 {
15949 let _drop_signal = ProviderFutureDropSignal {
15950 dropped: Arc::clone(&self.first_dropped),
15951 };
15952 self.first_started.notify_one();
15953 return std::future::pending().await;
15954 }
15955 self.committed_after_drop
15956 .store(self.first_dropped.load(Ordering::SeqCst), Ordering::SeqCst);
15957 Ok(LLMResponse::new(
15958 "Committed technical response.",
15959 FinishReason::Stop,
15960 ))
15961 }
15962
15963 async fn complete_stream(
15964 &self,
15965 _messages: &[ChatMessage],
15966 _config: Option<&LLMConfig>,
15967 ) -> std::result::Result<
15968 Box<dyn Stream<Item = std::result::Result<LLMChunk, LLMError>> + Unpin + Send>,
15969 LLMError,
15970 > {
15971 Err(LLMError::Other(
15972 "streaming is not used in this test".to_string(),
15973 ))
15974 }
15975
15976 fn provider_name(&self) -> &str {
15977 "first-call-locking"
15978 }
15979
15980 fn supports(&self, _feature: LLMFeature) -> bool {
15981 false
15982 }
15983 }
15984
15985 struct RoutingAfterProviderStart {
15986 provider_started: Arc<tokio::sync::Notify>,
15987 }
15988
15989 #[async_trait]
15990 impl LLMProvider for RoutingAfterProviderStart {
15991 async fn complete(
15992 &self,
15993 _messages: &[ChatMessage],
15994 _config: Option<&LLMConfig>,
15995 ) -> std::result::Result<LLMResponse, LLMError> {
15996 self.provider_started.notified().await;
15997 Ok(LLMResponse::new("1", FinishReason::Stop))
15998 }
15999
16000 async fn complete_stream(
16001 &self,
16002 _messages: &[ChatMessage],
16003 _config: Option<&LLMConfig>,
16004 ) -> std::result::Result<
16005 Box<dyn Stream<Item = std::result::Result<LLMChunk, LLMError>> + Unpin + Send>,
16006 LLMError,
16007 > {
16008 Err(LLMError::Other(
16009 "streaming is not used in this test".to_string(),
16010 ))
16011 }
16012
16013 fn provider_name(&self) -> &str {
16014 "routing-after-start"
16015 }
16016
16017 fn supports(&self, _feature: LLMFeature) -> bool {
16018 false
16019 }
16020 }
16021
16022 struct ResponseCountingHooks {
16024 responses: Arc<std::sync::atomic::AtomicUsize>,
16025 }
16026
16027 struct RootTurnProbeProvider {
16029 complete_entered: tokio::sync::mpsc::UnboundedSender<()>,
16030 }
16031
16032 struct ResponseChatHooks {
16034 target: parking_lot::Mutex<Option<Weak<RuntimeAgent>>>,
16035 invoked: AtomicBool,
16036 nested_result: parking_lot::Mutex<Option<std::result::Result<String, String>>>,
16037 }
16038
16039 struct ConcurrentResponseHooks {
16041 registry: Weak<crate::spawner::AgentRegistry>,
16042 child_id: String,
16043 invoked: AtomicBool,
16044 nested_result: parking_lot::Mutex<Option<std::result::Result<String, String>>>,
16045 }
16046
16047 struct RetryDeadlineTool {
16049 calls: Arc<std::sync::atomic::AtomicUsize>,
16050 deadlines: Arc<parking_lot::Mutex<Vec<chrono::DateTime<chrono::Utc>>>>,
16051 remaining_ms: Arc<parking_lot::Mutex<Vec<i64>>>,
16052 }
16053
16054 struct ToolLifecycleRecordingHooks {
16056 events: parking_lot::Mutex<Vec<String>>,
16057 records: parking_lot::Mutex<Vec<ToolExecutionRecord>>,
16058 }
16059
16060 impl ToolLifecycleRecordingHooks {
16061 fn new() -> Self {
16063 Self {
16064 events: parking_lot::Mutex::new(Vec::new()),
16065 records: parking_lot::Mutex::new(Vec::new()),
16066 }
16067 }
16068
16069 fn events(&self) -> Vec<String> {
16071 self.events.lock().clone()
16072 }
16073
16074 fn records(&self) -> Vec<ToolExecutionRecord> {
16076 self.records.lock().clone()
16077 }
16078 }
16079
16080 struct ContextEchoTool;
16082
16083 #[async_trait]
16084 impl LLMProvider for RootTurnProbeProvider {
16085 async fn complete(
16086 &self,
16087 _messages: &[ChatMessage],
16088 _config: Option<&LLMConfig>,
16089 ) -> std::result::Result<LLMResponse, LLMError> {
16090 let _ = self.complete_entered.send(());
16091 Ok(LLMResponse::new("blocking complete", FinishReason::Stop))
16092 }
16093
16094 async fn complete_stream(
16095 &self,
16096 _messages: &[ChatMessage],
16097 _config: Option<&LLMConfig>,
16098 ) -> std::result::Result<
16099 Box<dyn Stream<Item = std::result::Result<LLMChunk, LLMError>> + Unpin + Send>,
16100 LLMError,
16101 > {
16102 Ok(Box::new(futures::stream::iter(vec![Ok(
16103 LLMChunk::final_chunk("stream complete", FinishReason::Stop, None),
16104 )])))
16105 }
16106
16107 fn provider_name(&self) -> &str {
16108 "root-turn-probe"
16109 }
16110
16111 fn supports(&self, feature: LLMFeature) -> bool {
16112 matches!(feature, LLMFeature::Streaming)
16113 }
16114 }
16115
16116 #[async_trait]
16117 impl ai_agents_core::Tool for ContextEchoTool {
16118 fn id(&self) -> &str {
16119 "context_echo"
16120 }
16121
16122 fn name(&self) -> &str {
16123 "Context Echo"
16124 }
16125
16126 fn description(&self) -> &str {
16127 "Returns selected execution context fields."
16128 }
16129
16130 fn input_schema(&self) -> Value {
16131 serde_json::json!({"type": "object"})
16132 }
16133
16134 fn policy_bindings(&self) -> ai_agents_core::ToolPolicyBindings {
16135 ai_agents_core::ToolPolicyBindings {
16136 path_fields: vec![ai_agents_core::PathPolicyBinding::read("path")],
16137 result_limit_fields: vec![ai_agents_core::ResultLimitBinding::new(
16138 "max_results",
16139 ai_agents_core::ResultLimitKind::MaxResults,
16140 )],
16141 ..Default::default()
16142 }
16143 }
16144
16145 async fn execute(
16146 &self,
16147 _args: Value,
16148 ctx: ai_agents_core::ToolExecutionContext,
16149 ) -> ToolResult {
16150 ToolResult::ok(
16151 serde_json::json!({
16152 "requested_name": ctx.requested_name,
16153 "canonical_id": ctx.canonical_id,
16154 "display_name": ctx.display_name,
16155 "max_results": ctx.limits.max_results,
16156 "custom_config": ctx.custom_config,
16157 })
16158 .to_string(),
16159 )
16160 }
16161 }
16162
16163 #[async_trait]
16164 impl ai_agents_core::Tool for RetryDeadlineTool {
16165 fn id(&self) -> &str {
16166 "retry_deadline"
16167 }
16168
16169 fn name(&self) -> &str {
16170 "Retry Deadline"
16171 }
16172
16173 fn description(&self) -> &str {
16174 "Records one deadline per retry invocation."
16175 }
16176
16177 fn input_schema(&self) -> Value {
16178 serde_json::json!({"type": "object"})
16179 }
16180
16181 fn safety_metadata(&self) -> ai_agents_core::ToolSafetyMetadata {
16182 ai_agents_core::ToolSafetyMetadata::compute()
16183 }
16184
16185 fn classify_call(&self, _args: &Value) -> ai_agents_core::ToolCallClassification {
16186 let mut classification =
16187 ai_agents_core::ToolCallClassification::from_metadata(&self.safety_metadata());
16188 classification.timeout_ms = Some(1_000);
16189 classification.safely_retryable = true;
16190 classification
16191 }
16192
16193 async fn execute(
16195 &self,
16196 _args: Value,
16197 ctx: ai_agents_core::ToolExecutionContext,
16198 ) -> ToolResult {
16199 let deadline = ctx
16200 .deadline
16201 .expect("each invocation must receive a deadline");
16202 self.remaining_ms.lock().push(
16203 deadline
16204 .signed_duration_since(chrono::Utc::now())
16205 .num_milliseconds(),
16206 );
16207 self.deadlines.lock().push(deadline);
16208 let call = self.calls.fetch_add(1, Ordering::SeqCst);
16209 if call == 0 {
16210 tokio::time::sleep(std::time::Duration::from_millis(20)).await;
16211 ToolResult::error("retry")
16212 } else {
16213 ToolResult::ok("done")
16214 }
16215 }
16216 }
16217
16218 struct ClassifiedTimeoutTool {
16220 id: &'static str,
16221 calls: Arc<std::sync::atomic::AtomicUsize>,
16222 timeout_ms: u64,
16223 sleep_ms: u64,
16224 requires_approval: bool,
16225 remaining_ms: Arc<parking_lot::Mutex<Vec<i64>>>,
16226 }
16227
16228 struct ApprovalModifiedTimeoutTool {
16230 calls: Arc<std::sync::atomic::AtomicUsize>,
16231 }
16232
16233 struct SlowTool;
16235
16236 struct FlakyWriteTool {
16238 calls: Arc<std::sync::atomic::AtomicUsize>,
16239 }
16240
16241 struct LockedWriteTool {
16243 active: Arc<std::sync::atomic::AtomicUsize>,
16244 max_active: Arc<std::sync::atomic::AtomicUsize>,
16245 }
16246
16247 struct MultiResourceWriteTool {
16248 active: Arc<std::sync::atomic::AtomicUsize>,
16249 max_active: Arc<std::sync::atomic::AtomicUsize>,
16250 }
16251
16252 #[derive(Clone)]
16253 struct PathMutationGate {
16254 entered: Arc<AtomicBool>,
16255 entered_notify: Arc<tokio::sync::Notify>,
16256 release: Arc<tokio::sync::Notify>,
16257 }
16258
16259 impl PathMutationGate {
16260 fn new() -> Self {
16261 Self {
16262 entered: Arc::new(AtomicBool::new(false)),
16263 entered_notify: Arc::new(tokio::sync::Notify::new()),
16264 release: Arc::new(tokio::sync::Notify::new()),
16265 }
16266 }
16267
16268 async fn wait_until_entered(&self) {
16269 if !self.entered.load(Ordering::SeqCst) {
16270 self.entered_notify.notified().await;
16271 }
16272 }
16273
16274 fn release(&self) {
16275 self.release.notify_one();
16276 }
16277 }
16278
16279 struct BlockingPathMutationTool {
16280 id: &'static str,
16281 path_fields: Vec<ai_agents_core::PathPolicyBinding>,
16282 gate: PathMutationGate,
16283 }
16284
16285 struct NoBindingWriteTool {
16286 active: Arc<std::sync::atomic::AtomicUsize>,
16287 max_active: Arc<std::sync::atomic::AtomicUsize>,
16288 }
16289
16290 struct RecoveryTestTool {
16291 id: String,
16292 succeeds: bool,
16293 calls: Arc<std::sync::atomic::AtomicUsize>,
16294 max_output_chars: Option<usize>,
16295 }
16296
16297 struct BlockingApprovalHandler {
16298 entered: Arc<tokio::sync::Barrier>,
16299 release: Arc<tokio::sync::Notify>,
16300 result: ApprovalResult,
16301 }
16302
16303 struct CountingApprovalHandler {
16304 calls: Arc<std::sync::atomic::AtomicUsize>,
16305 }
16306
16307 struct DriftingFallbackProvider {
16309 refreshed: AtomicBool,
16310 primary_calls: Arc<std::sync::atomic::AtomicUsize>,
16311 secondary_calls: Arc<std::sync::atomic::AtomicUsize>,
16312 }
16313
16314 struct RefreshFallbackProviderHooks {
16316 agent: parking_lot::Mutex<Option<Weak<RuntimeAgent>>>,
16317 lifecycle: Arc<ToolLifecycleRecordingHooks>,
16318 }
16319
16320 struct RuntimeWebFetchTransport {
16321 calls: Arc<std::sync::atomic::AtomicUsize>,
16322 }
16323
16324 struct RuntimeWebFetchResolver;
16325
16326 struct ReentrantToolHooks {
16327 agent: parking_lot::Mutex<Option<Weak<RuntimeAgent>>>,
16328 invoked: AtomicBool,
16329 nested_success: AtomicBool,
16330 }
16331
16332 #[async_trait]
16333 impl ai_agents_core::Tool for ClassifiedTimeoutTool {
16334 fn id(&self) -> &str {
16336 self.id
16337 }
16338
16339 fn name(&self) -> &str {
16341 "Classified Timeout"
16342 }
16343
16344 fn description(&self) -> &str {
16346 "Records and waits under one call-level timeout."
16347 }
16348
16349 fn input_schema(&self) -> Value {
16351 serde_json::json!({"type": "object"})
16352 }
16353
16354 fn classify_call(&self, _args: &Value) -> ai_agents_core::ToolCallClassification {
16356 let mut classification =
16357 ai_agents_core::ToolCallClassification::from_metadata(&self.safety_metadata());
16358 classification.timeout_ms = Some(self.timeout_ms);
16359 classification.requires_approval = self.requires_approval;
16360 classification
16361 }
16362
16363 async fn execute(
16365 &self,
16366 _args: Value,
16367 ctx: ai_agents_core::ToolExecutionContext,
16368 ) -> ToolResult {
16369 self.calls.fetch_add(1, Ordering::SeqCst);
16370 let deadline = ctx
16371 .deadline
16372 .expect("each invocation must receive a deadline");
16373 self.remaining_ms.lock().push(
16374 deadline
16375 .signed_duration_since(chrono::Utc::now())
16376 .num_milliseconds(),
16377 );
16378 tokio::time::sleep(Duration::from_millis(self.sleep_ms)).await;
16379 ToolResult::ok("done")
16380 }
16381 }
16382
16383 #[async_trait]
16384 impl ai_agents_core::Tool for ApprovalModifiedTimeoutTool {
16385 fn id(&self) -> &str {
16387 "approval_modified_timeout"
16388 }
16389
16390 fn name(&self) -> &str {
16392 "Approval Modified Timeout"
16393 }
16394
16395 fn description(&self) -> &str {
16397 "Becomes invalid only after approval modifies its arguments."
16398 }
16399
16400 fn input_schema(&self) -> Value {
16402 serde_json::json!({"type": "object"})
16403 }
16404
16405 fn policy_bindings(&self) -> ai_agents_core::ToolPolicyBindings {
16407 ai_agents_core::ToolPolicyBindings {
16408 path_fields: vec![ai_agents_core::PathPolicyBinding::write("path")],
16409 ..Default::default()
16410 }
16411 }
16412
16413 fn safety_metadata(&self) -> ai_agents_core::ToolSafetyMetadata {
16415 ai_agents_core::ToolSafetyMetadata {
16416 read_only: false,
16417 concurrency_safe: false,
16418 operation: ai_agents_core::ToolOperationKind::Write,
16419 side_effect_level: ai_agents_core::ToolSideEffectLevel::LocalWrite,
16420 requires_network: false,
16421 destructive: false,
16422 open_world: false,
16423 host_dependent: false,
16424 requires_user_interaction: false,
16425 supports_cancellation: true,
16426 default_requires_approval: true,
16427 should_defer_schema: false,
16428 max_output_chars: Some(1024),
16429 max_result_size_chars: Some(1024),
16430 }
16431 }
16432
16433 fn classify_call(&self, args: &Value) -> ai_agents_core::ToolCallClassification {
16435 let mut classification =
16436 ai_agents_core::ToolCallClassification::from_metadata(&self.safety_metadata());
16437 classification.timeout_ms = Some(if args["invalid_timeout"].as_bool() == Some(true) {
16438 u64::MAX
16439 } else {
16440 1_000
16441 });
16442 classification
16443 }
16444
16445 async fn execute(
16447 &self,
16448 _args: Value,
16449 _ctx: ai_agents_core::ToolExecutionContext,
16450 ) -> ToolResult {
16451 self.calls.fetch_add(1, Ordering::SeqCst);
16452 ToolResult::ok("unexpected")
16453 }
16454 }
16455
16456 #[async_trait]
16457 impl ai_agents_core::Tool for SlowTool {
16458 fn id(&self) -> &str {
16459 "slow"
16460 }
16461
16462 fn name(&self) -> &str {
16463 "Slow"
16464 }
16465
16466 fn description(&self) -> &str {
16467 "Waits until cancelled or timed out."
16468 }
16469
16470 fn input_schema(&self) -> Value {
16471 serde_json::json!({"type": "object"})
16472 }
16473
16474 async fn execute(
16475 &self,
16476 _args: Value,
16477 _ctx: ai_agents_core::ToolExecutionContext,
16478 ) -> ToolResult {
16479 tokio::time::sleep(std::time::Duration::from_secs(5)).await;
16480 ToolResult::ok("done")
16481 }
16482 }
16483
16484 #[async_trait]
16485 impl ai_agents_core::Tool for FlakyWriteTool {
16486 fn id(&self) -> &str {
16487 "flaky_write"
16488 }
16489
16490 fn name(&self) -> &str {
16491 "Flaky Write"
16492 }
16493
16494 fn description(&self) -> &str {
16495 "Fails on the first write attempt."
16496 }
16497
16498 fn input_schema(&self) -> Value {
16499 serde_json::json!({"type": "object", "properties": {"path": {"type": "string"}}})
16500 }
16501
16502 fn policy_bindings(&self) -> ai_agents_core::ToolPolicyBindings {
16503 ai_agents_core::ToolPolicyBindings {
16504 path_fields: vec![ai_agents_core::PathPolicyBinding::write("path")],
16505 ..Default::default()
16506 }
16507 }
16508
16509 fn safety_metadata(&self) -> ai_agents_core::ToolSafetyMetadata {
16510 ai_agents_core::ToolSafetyMetadata {
16511 read_only: false,
16512 concurrency_safe: false,
16513 operation: ai_agents_core::ToolOperationKind::Write,
16514 side_effect_level: ai_agents_core::ToolSideEffectLevel::LocalWrite,
16515 requires_network: false,
16516 destructive: false,
16517 open_world: false,
16518 host_dependent: false,
16519 requires_user_interaction: false,
16520 supports_cancellation: true,
16521 default_requires_approval: false,
16522 should_defer_schema: false,
16523 max_output_chars: Some(1024),
16524 max_result_size_chars: Some(1024),
16525 }
16526 }
16527
16528 fn classify_call(&self, _args: &Value) -> ai_agents_core::ToolCallClassification {
16529 let mut classification =
16530 ai_agents_core::ToolCallClassification::from_metadata(&self.safety_metadata());
16531 classification.safely_retryable = false;
16532 classification
16533 }
16534
16535 async fn execute(
16536 &self,
16537 _args: Value,
16538 _ctx: ai_agents_core::ToolExecutionContext,
16539 ) -> ToolResult {
16540 let call = self.calls.fetch_add(1, Ordering::SeqCst);
16541 if call == 0 {
16542 ToolResult::error("first failure")
16543 } else {
16544 ToolResult::ok("second success")
16545 }
16546 }
16547 }
16548
16549 #[async_trait]
16550 impl ai_agents_core::Tool for LockedWriteTool {
16551 fn id(&self) -> &str {
16552 "locked_write"
16553 }
16554
16555 fn name(&self) -> &str {
16556 "Locked Write"
16557 }
16558
16559 fn description(&self) -> &str {
16560 "Tracks concurrent execution on one resource."
16561 }
16562
16563 fn input_schema(&self) -> Value {
16564 serde_json::json!({"type": "object", "properties": {"path": {"type": "string"}}})
16565 }
16566
16567 fn policy_bindings(&self) -> ai_agents_core::ToolPolicyBindings {
16568 ai_agents_core::ToolPolicyBindings {
16569 path_fields: vec![ai_agents_core::PathPolicyBinding::write("path")],
16570 ..Default::default()
16571 }
16572 }
16573
16574 fn safety_metadata(&self) -> ai_agents_core::ToolSafetyMetadata {
16575 ai_agents_core::ToolSafetyMetadata {
16576 read_only: false,
16577 concurrency_safe: false,
16578 operation: ai_agents_core::ToolOperationKind::Write,
16579 side_effect_level: ai_agents_core::ToolSideEffectLevel::LocalWrite,
16580 requires_network: false,
16581 destructive: false,
16582 open_world: false,
16583 host_dependent: false,
16584 requires_user_interaction: false,
16585 supports_cancellation: true,
16586 default_requires_approval: false,
16587 should_defer_schema: false,
16588 max_output_chars: Some(1024),
16589 max_result_size_chars: Some(1024),
16590 }
16591 }
16592
16593 async fn execute(
16594 &self,
16595 _args: Value,
16596 _ctx: ai_agents_core::ToolExecutionContext,
16597 ) -> ToolResult {
16598 let active = self.active.fetch_add(1, Ordering::SeqCst) + 1;
16599 loop {
16600 let current_max = self.max_active.load(Ordering::SeqCst);
16601 if active <= current_max {
16602 break;
16603 }
16604 if self
16605 .max_active
16606 .compare_exchange(current_max, active, Ordering::SeqCst, Ordering::SeqCst)
16607 .is_ok()
16608 {
16609 break;
16610 }
16611 }
16612 tokio::time::sleep(std::time::Duration::from_millis(100)).await;
16613 self.active.fetch_sub(1, Ordering::SeqCst);
16614 ToolResult::ok("done")
16615 }
16616 }
16617
16618 #[async_trait]
16619 impl ai_agents_core::Tool for MultiResourceWriteTool {
16620 fn id(&self) -> &str {
16621 "multi_resource_write"
16622 }
16623
16624 fn name(&self) -> &str {
16625 "Multi Resource Write"
16626 }
16627
16628 fn description(&self) -> &str {
16629 "Tracks concurrent execution across source and destination resources."
16630 }
16631
16632 fn input_schema(&self) -> Value {
16633 serde_json::json!({"type": "object"})
16634 }
16635
16636 fn policy_bindings(&self) -> ai_agents_core::ToolPolicyBindings {
16637 ai_agents_core::ToolPolicyBindings {
16638 path_fields: vec![
16639 ai_agents_core::PathPolicyBinding::read_write("source_path"),
16640 ai_agents_core::PathPolicyBinding::write("destination_path"),
16641 ],
16642 ..Default::default()
16643 }
16644 }
16645
16646 fn safety_metadata(&self) -> ai_agents_core::ToolSafetyMetadata {
16647 LockedWriteTool {
16648 active: Arc::clone(&self.active),
16649 max_active: Arc::clone(&self.max_active),
16650 }
16651 .safety_metadata()
16652 }
16653
16654 async fn execute(
16655 &self,
16656 _args: Value,
16657 _ctx: ai_agents_core::ToolExecutionContext,
16658 ) -> ToolResult {
16659 let active = self.active.fetch_add(1, Ordering::SeqCst) + 1;
16660 self.max_active.fetch_max(active, Ordering::SeqCst);
16661 tokio::time::sleep(std::time::Duration::from_millis(75)).await;
16662 self.active.fetch_sub(1, Ordering::SeqCst);
16663 ToolResult::ok("done")
16664 }
16665 }
16666
16667 #[async_trait]
16668 impl ai_agents_core::Tool for BlockingPathMutationTool {
16669 fn id(&self) -> &str {
16670 self.id
16671 }
16672
16673 fn name(&self) -> &str {
16674 self.id
16675 }
16676
16677 fn description(&self) -> &str {
16678 "Blocks a path mutation until the test releases it."
16679 }
16680
16681 fn input_schema(&self) -> Value {
16682 serde_json::json!({"type": "object"})
16683 }
16684
16685 fn policy_bindings(&self) -> ai_agents_core::ToolPolicyBindings {
16686 ai_agents_core::ToolPolicyBindings {
16687 path_fields: self.path_fields.clone(),
16688 ..Default::default()
16689 }
16690 }
16691
16692 fn safety_metadata(&self) -> ai_agents_core::ToolSafetyMetadata {
16693 ai_agents_core::ToolSafetyMetadata {
16694 read_only: false,
16695 concurrency_safe: false,
16696 operation: ai_agents_core::ToolOperationKind::Write,
16697 side_effect_level: ai_agents_core::ToolSideEffectLevel::LocalWrite,
16698 requires_network: false,
16699 destructive: false,
16700 open_world: false,
16701 host_dependent: false,
16702 requires_user_interaction: false,
16703 supports_cancellation: true,
16704 default_requires_approval: false,
16705 should_defer_schema: false,
16706 max_output_chars: Some(1024),
16707 max_result_size_chars: Some(1024),
16708 }
16709 }
16710
16711 async fn execute(
16712 &self,
16713 _args: Value,
16714 _ctx: ai_agents_core::ToolExecutionContext,
16715 ) -> ToolResult {
16716 self.gate.entered.store(true, Ordering::SeqCst);
16717 self.gate.entered_notify.notify_one();
16718 self.gate.release.notified().await;
16719 ToolResult::ok("done")
16720 }
16721 }
16722
16723 #[async_trait]
16724 impl ai_agents_core::Tool for NoBindingWriteTool {
16725 fn id(&self) -> &str {
16726 "no_binding_write"
16727 }
16728
16729 fn name(&self) -> &str {
16730 "No Binding Write"
16731 }
16732
16733 fn description(&self) -> &str {
16734 "Tracks concurrent execution without resource bindings."
16735 }
16736
16737 fn input_schema(&self) -> Value {
16738 serde_json::json!({"type": "object"})
16739 }
16740
16741 fn safety_metadata(&self) -> ai_agents_core::ToolSafetyMetadata {
16742 LockedWriteTool {
16743 active: Arc::clone(&self.active),
16744 max_active: Arc::clone(&self.max_active),
16745 }
16746 .safety_metadata()
16747 }
16748
16749 async fn execute(
16750 &self,
16751 _args: Value,
16752 _ctx: ai_agents_core::ToolExecutionContext,
16753 ) -> ToolResult {
16754 let active = self.active.fetch_add(1, Ordering::SeqCst) + 1;
16755 self.max_active.fetch_max(active, Ordering::SeqCst);
16756 tokio::time::sleep(std::time::Duration::from_millis(75)).await;
16757 self.active.fetch_sub(1, Ordering::SeqCst);
16758 ToolResult::ok("done")
16759 }
16760 }
16761
16762 #[async_trait]
16763 impl ai_agents_core::Tool for RecoveryTestTool {
16764 fn id(&self) -> &str {
16765 &self.id
16766 }
16767
16768 fn name(&self) -> &str {
16769 &self.id
16770 }
16771
16772 fn description(&self) -> &str {
16773 "Records recovery execution and returns a configured result."
16774 }
16775
16776 fn input_schema(&self) -> Value {
16777 serde_json::json!({"type": "object"})
16778 }
16779
16780 fn policy_bindings(&self) -> ai_agents_core::ToolPolicyBindings {
16781 ai_agents_core::ToolPolicyBindings {
16782 path_fields: vec![ai_agents_core::PathPolicyBinding::write("path")],
16783 ..Default::default()
16784 }
16785 }
16786
16787 fn safety_metadata(&self) -> ai_agents_core::ToolSafetyMetadata {
16789 ai_agents_core::ToolSafetyMetadata {
16790 read_only: false,
16791 concurrency_safe: false,
16792 operation: ai_agents_core::ToolOperationKind::Write,
16793 side_effect_level: ai_agents_core::ToolSideEffectLevel::LocalWrite,
16794 requires_network: false,
16795 destructive: false,
16796 open_world: false,
16797 host_dependent: false,
16798 requires_user_interaction: false,
16799 supports_cancellation: true,
16800 default_requires_approval: false,
16801 should_defer_schema: false,
16802 max_output_chars: Some(self.max_output_chars.unwrap_or(1024)),
16803 max_result_size_chars: Some(1024),
16804 }
16805 }
16806
16807 async fn execute(
16809 &self,
16810 _args: Value,
16811 _ctx: ai_agents_core::ToolExecutionContext,
16812 ) -> ToolResult {
16813 self.calls.fetch_add(1, Ordering::SeqCst);
16814 let mut result = if self.succeeds {
16815 ToolResult::ok(format!("{} succeeded", self.id))
16816 } else {
16817 ToolResult::error(format!("{} failed", self.id))
16818 };
16819 result.metadata = Some(HashMap::from([(
16820 "recovery_test_tool".to_string(),
16821 Value::String(self.id.clone()),
16822 )]));
16823 result
16824 }
16825 }
16826
16827 #[async_trait]
16828 impl WebFetchTransport for RuntimeWebFetchTransport {
16829 async fn send(
16831 &self,
16832 _request: WebFetchTransportRequest,
16833 ) -> std::result::Result<WebFetchTransportResponse, String> {
16834 Err("validated addresses are required".to_string())
16835 }
16836
16837 async fn send_validated(
16839 &self,
16840 _request: WebFetchTransportRequest,
16841 _addresses: &[std::net::SocketAddr],
16842 ) -> std::result::Result<WebFetchTransportResponse, String> {
16843 self.calls.fetch_add(1, Ordering::SeqCst);
16844 Ok(WebFetchTransportResponse {
16845 status: 200,
16846 content_type: Some("text/plain".to_string()),
16847 location: None,
16848 body: b"approved".to_vec(),
16849 })
16850 }
16851 }
16852
16853 #[async_trait]
16854 impl WebFetchResolver for RuntimeWebFetchResolver {
16855 async fn resolve(
16857 &self,
16858 _host: &str,
16859 _port: u16,
16860 ) -> std::result::Result<Vec<std::net::IpAddr>, String> {
16861 Ok(vec![std::net::IpAddr::V4(std::net::Ipv4Addr::new(
16862 93, 184, 216, 34,
16863 ))])
16864 }
16865 }
16866
16867 #[async_trait]
16868 impl ToolProvider for DriftingFallbackProvider {
16869 fn id(&self) -> &str {
16871 "drifting_fallback"
16872 }
16873
16874 fn name(&self) -> &str {
16876 "Drifting Fallback"
16877 }
16878
16879 fn provider_type(&self) -> ToolProviderType {
16881 ToolProviderType::Custom
16882 }
16883
16884 async fn list_tools(&self) -> Vec<ToolDescriptor> {
16886 let alias = ToolAliases::new().with_name("en", "fallback alias");
16887 let mut primary = ToolDescriptor::new(
16888 "primary",
16889 "Primary",
16890 "Fails before fallback.",
16891 serde_json::json!({"type": "object"}),
16892 );
16893 let mut secondary = ToolDescriptor::new(
16894 "secondary",
16895 "Secondary",
16896 "Must not execute after final canonical drift.",
16897 serde_json::json!({"type": "object"}),
16898 );
16899 if self.refreshed.load(Ordering::SeqCst) {
16900 primary = primary.with_aliases(alias);
16901 } else {
16902 secondary = secondary.with_aliases(alias);
16903 }
16904 vec![primary, secondary]
16905 }
16906
16907 async fn get_tool(&self, tool_id: &str) -> Option<Arc<dyn Tool>> {
16909 let calls = match tool_id {
16910 "primary" => Arc::clone(&self.primary_calls),
16911 "secondary" => Arc::clone(&self.secondary_calls),
16912 _ => return None,
16913 };
16914 Some(Arc::new(RecoveryTestTool {
16915 id: tool_id.to_string(),
16916 succeeds: false,
16917 calls,
16918 max_output_chars: None,
16919 }))
16920 }
16921
16922 fn supports_refresh(&self) -> bool {
16924 true
16925 }
16926
16927 async fn refresh(&self) -> std::result::Result<(), ToolProviderError> {
16929 self.refreshed.store(true, Ordering::SeqCst);
16930 Ok(())
16931 }
16932 }
16933
16934 #[async_trait]
16935 impl AgentHooks for RefreshFallbackProviderHooks {
16936 async fn on_tool_start(&self, tool: &str, args: &Value) {
16938 self.lifecycle.on_tool_start(tool, args).await;
16939 if tool != "secondary" {
16940 return;
16941 }
16942 let agent = self.agent.lock().as_ref().and_then(Weak::upgrade);
16943 if let Some(agent) = agent {
16944 agent
16945 .tools
16946 .refresh_provider("drifting_fallback")
16947 .await
16948 .unwrap();
16949 }
16950 }
16951
16952 async fn on_tool_complete(&self, tool: &str, result: &ToolResult, duration_ms: u64) {
16953 self.lifecycle
16954 .on_tool_complete(tool, result, duration_ms)
16955 .await;
16956 }
16957
16958 async fn on_tool_execution_record(&self, record: &ToolExecutionRecord) {
16959 self.lifecycle.on_tool_execution_record(record).await;
16960 }
16961
16962 async fn on_error(&self, error: &AgentError) {
16963 self.lifecycle.on_error(error).await;
16964 }
16965 }
16966
16967 #[async_trait]
16968 impl ApprovalHandler for BlockingApprovalHandler {
16969 async fn request_approval(
16970 &self,
16971 _request: ai_agents_hitl::ApprovalRequest,
16972 ) -> ApprovalResult {
16973 self.entered.wait().await;
16974 self.release.notified().await;
16975 self.result.clone()
16976 }
16977 }
16978
16979 #[async_trait]
16980 impl ApprovalHandler for CountingApprovalHandler {
16981 async fn request_approval(
16982 &self,
16983 _request: ai_agents_hitl::ApprovalRequest,
16984 ) -> ApprovalResult {
16985 self.calls.fetch_add(1, Ordering::SeqCst);
16986 ApprovalResult::Approved
16987 }
16988 }
16989
16990 #[async_trait]
16991 impl AgentHooks for ReentrantToolHooks {
16992 async fn on_tool_complete(&self, tool: &str, _result: &ToolResult, _duration_ms: u64) {
16993 if tool != "reentrant_write" || self.invoked.swap(true, Ordering::SeqCst) {
16994 return;
16995 }
16996 let agent = self.agent.lock().as_ref().and_then(Weak::upgrade);
16997 if let Some(agent) = agent {
16998 let result = agent
16999 .invoke_tool(ToolExecutionRequest::new(
17000 "nested-hook-call",
17001 "reentrant_write",
17002 serde_json::json!({"path": "./hook.txt"}),
17003 ToolCallSource::Manual,
17004 ))
17005 .await;
17006 self.nested_success
17007 .store(result.is_ok_and(|record| record.success), Ordering::SeqCst);
17008 }
17009 }
17010 }
17011
17012 #[async_trait]
17013 impl AgentHooks for ResponseCountingHooks {
17014 async fn on_response(&self, _response: &AgentResponse) {
17015 self.responses.fetch_add(1, Ordering::SeqCst);
17016 }
17017 }
17018
17019 #[async_trait]
17020 impl AgentHooks for ResponseChatHooks {
17021 async fn on_response(&self, _response: &AgentResponse) {
17023 if self.invoked.swap(true, Ordering::SeqCst) {
17024 return;
17025 }
17026 let target = self.target.lock().as_ref().and_then(Weak::upgrade);
17027 let result = if let Some(target) = target {
17028 target
17029 .chat("nested response hook call")
17030 .await
17031 .map(|response| response.content)
17032 .map_err(|error| error.to_string())
17033 } else {
17034 Err("response hook target is unavailable".to_string())
17035 };
17036 *self.nested_result.lock() = Some(result);
17037 }
17038 }
17039
17040 #[async_trait]
17041 impl AgentHooks for ConcurrentResponseHooks {
17042 async fn on_response(&self, _response: &AgentResponse) {
17044 if self.invoked.swap(true, Ordering::SeqCst) {
17045 return;
17046 }
17047 let Some(registry) = self.registry.upgrade() else {
17048 *self.nested_result.lock() =
17049 Some(Err("concurrent registry is unavailable".to_string()));
17050 return;
17051 };
17052 let agents = [ai_agents_state::ConcurrentAgentRef::Id(
17053 self.child_id.clone(),
17054 )];
17055 let aggregation = ai_agents_state::AggregationConfig {
17056 strategy: ai_agents_state::AggregationStrategy::FirstWins,
17057 synthesizer_llm: None,
17058 synthesizer_prompt: None,
17059 vote: None,
17060 };
17061 let result = crate::orchestration::concurrent(
17062 ®istry,
17063 "nested concurrent response hook call",
17064 &agents,
17065 &aggregation,
17066 None,
17067 Some(1),
17068 None,
17069 ai_agents_state::PartialFailureAction::Abort,
17070 None,
17071 )
17072 .await
17073 .map(|result| result.response.content)
17074 .map_err(|error| error.to_string());
17075 *self.nested_result.lock() = Some(result);
17076 }
17077 }
17078
17079 #[async_trait]
17080 impl AgentHooks for ToolLifecycleRecordingHooks {
17081 async fn on_tool_start(&self, tool: &str, _args: &Value) {
17082 self.events.lock().push(format!("start:{tool}"));
17083 }
17084
17085 async fn on_tool_complete(&self, tool: &str, result: &ToolResult, _duration_ms: u64) {
17086 self.events
17087 .lock()
17088 .push(format!("complete:{tool}:{}", result.success));
17089 }
17090
17091 async fn on_tool_execution_record(&self, record: &ToolExecutionRecord) {
17092 self.events.lock().push(format!(
17093 "record:{}:{}",
17094 record.canonical_id, record.executed
17095 ));
17096 self.records.lock().push(record.clone());
17097 }
17098
17099 async fn on_error(&self, _error: &AgentError) {
17101 self.events.lock().push("error".to_string());
17102 }
17103 }
17104
17105 struct ApprovalRecordingHooks {
17106 events: parking_lot::Mutex<Vec<String>>,
17107 }
17108
17109 impl ApprovalRecordingHooks {
17110 fn new() -> Self {
17111 Self {
17112 events: parking_lot::Mutex::new(Vec::new()),
17113 }
17114 }
17115
17116 fn events(&self) -> Vec<String> {
17117 self.events.lock().clone()
17118 }
17119 }
17120
17121 #[async_trait]
17122 impl AgentHooks for ApprovalRecordingHooks {
17123 async fn on_approval_result(&self, request_id: &str, result: &ApprovalResult) {
17124 self.events.lock().push(format!(
17125 "raw:{}:{}",
17126 request_id,
17127 approval_result_name(result)
17128 ));
17129 }
17130
17131 async fn on_approval_resolved(
17132 &self,
17133 request: &ai_agents_hitl::ApprovalRequest,
17134 raw_result: &ApprovalResult,
17135 outcome: &ApprovalResolvedOutcome,
17136 ) {
17137 self.events.lock().push(format!(
17138 "resolved:{}:{}:{}",
17139 request.id,
17140 approval_result_name(raw_result),
17141 approval_outcome_name(outcome)
17142 ));
17143 }
17144 }
17145
17146 fn approval_result_name(result: &ApprovalResult) -> &'static str {
17147 match result {
17148 ApprovalResult::Approved => "approved",
17149 ApprovalResult::Rejected { .. } => "rejected",
17150 ApprovalResult::Modified { .. } => "modified",
17151 ApprovalResult::Timeout => "timeout",
17152 }
17153 }
17154
17155 fn approval_outcome_name(outcome: &ApprovalResolvedOutcome) -> &'static str {
17156 match outcome {
17157 ApprovalResolvedOutcome::Approved => "approved",
17158 ApprovalResolvedOutcome::Rejected { .. } => "rejected",
17159 ApprovalResolvedOutcome::Modified { .. } => "modified",
17160 ApprovalResolvedOutcome::Error { .. } => "error",
17161 }
17162 }
17163
17164 fn assert_correlated_approval_events(
17165 events: &[String],
17166 raw_status: &str,
17167 outcome_status: &str,
17168 ) {
17169 assert_eq!(events.len(), 2);
17170 let raw: Vec<_> = events[0].split(':').collect();
17171 let resolved: Vec<_> = events[1].split(':').collect();
17172 assert_eq!(raw[0], "raw");
17173 assert_eq!(resolved[0], "resolved");
17174 assert_eq!(raw[1], resolved[1]);
17175 assert_eq!(raw[2], raw_status);
17176 assert_eq!(resolved[2], raw_status);
17177 assert_eq!(resolved[3], outcome_status);
17178 }
17179
17180 fn approval_security_config(policy_enabled: bool) -> ToolSecurityConfig {
17181 let mut security = ToolSecurityConfig {
17182 enabled: true,
17183 fail_closed: true,
17184 ..Default::default()
17185 };
17186 let policy = ai_agents_tools::ToolPolicyConfig {
17187 enabled: policy_enabled,
17188 write_paths: vec![".".to_string()],
17189 require_confirmation: true,
17190 ..Default::default()
17191 };
17192 security.tools.insert("locked_write".to_string(), policy);
17193 security
17194 }
17195
17196 struct MutationTestWorkspace {
17197 root: std::path::PathBuf,
17198 }
17199
17200 impl MutationTestWorkspace {
17201 fn new() -> Self {
17202 let root = std::env::temp_dir().join(format!(
17203 "ai-agents-runtime-mutation-{}",
17204 uuid::Uuid::new_v4()
17205 ));
17206 std::fs::create_dir_all(&root).unwrap();
17207 Self { root }
17208 }
17209 }
17210
17211 impl Drop for MutationTestWorkspace {
17212 fn drop(&mut self) {
17213 let _ = std::fs::remove_dir_all(&self.root);
17214 }
17215 }
17216
17217 async fn wait_for_resource_lock_strong_count(locks: &ToolResourceLocks, minimum: usize) {
17218 tokio::time::timeout(std::time::Duration::from_secs(2), async {
17219 loop {
17220 let strong_count = locks
17221 .read()
17222 .get("path-mutation:global")
17223 .map_or(0, |lock| lock.strong_count());
17224 if strong_count >= minimum {
17225 break;
17226 }
17227 tokio::task::yield_now().await;
17228 }
17229 })
17230 .await
17231 .expect("path mutation call did not reach the shared lock");
17232 }
17233
17234 async fn assert_path_mutation_pair_serialized(
17235 first_id: &'static str,
17236 first_fields: Vec<ai_agents_core::PathPolicyBinding>,
17237 first_args: Value,
17238 second_id: &'static str,
17239 second_fields: Vec<ai_agents_core::PathPolicyBinding>,
17240 second_args: Value,
17241 ) {
17242 let locks = new_tool_resource_locks();
17243 let first_gate = PathMutationGate::new();
17244 let second_gate = PathMutationGate::new();
17245 second_gate.release();
17246 let agent = Arc::new(
17247 AgentBuilder::new()
17248 .system_prompt("Test global path mutation locking.")
17249 .llm(Arc::new(mock_with_response("done")))
17250 .tool(Arc::new(BlockingPathMutationTool {
17251 id: first_id,
17252 path_fields: first_fields,
17253 gate: first_gate.clone(),
17254 }))
17255 .tool(Arc::new(BlockingPathMutationTool {
17256 id: second_id,
17257 path_fields: second_fields,
17258 gate: second_gate.clone(),
17259 }))
17260 .build()
17261 .unwrap()
17262 .with_shared_resource_locks(Arc::clone(&locks)),
17263 );
17264
17265 let first = {
17266 let agent = Arc::clone(&agent);
17267 tokio::spawn(async move {
17268 agent
17269 .invoke_tool(ToolExecutionRequest::new(
17270 format!("{}-first", first_id),
17271 first_id,
17272 first_args,
17273 ToolCallSource::Manual,
17274 ))
17275 .await
17276 .unwrap()
17277 })
17278 };
17279 first_gate.wait_until_entered().await;
17280
17281 let second = {
17282 let agent = Arc::clone(&agent);
17283 tokio::spawn(async move {
17284 agent
17285 .invoke_tool(ToolExecutionRequest::new(
17286 format!("{}-second", second_id),
17287 second_id,
17288 second_args,
17289 ToolCallSource::Manual,
17290 ))
17291 .await
17292 .unwrap()
17293 })
17294 };
17295 wait_for_resource_lock_strong_count(&locks, 2).await;
17296 assert!(!second_gate.entered.load(Ordering::SeqCst));
17297 assert!(!second.is_finished());
17298
17299 first_gate.release();
17300 let (first, second) = tokio::time::timeout(std::time::Duration::from_secs(2), async {
17301 tokio::join!(first, second)
17302 })
17303 .await
17304 .expect("serialized path mutation calls did not finish");
17305 assert!(first.unwrap().success);
17306 assert!(second.unwrap().success);
17307 assert!(second_gate.entered.load(Ordering::SeqCst));
17308 assert!(locks.read().is_empty());
17309 }
17310
17311 #[derive(Clone, Copy)]
17312 enum MutationDenial {
17313 Policy,
17314 Approval,
17315 }
17316
17317 fn mutation_denial_security_config(
17318 tool_id: &str,
17319 workspace: &std::path::Path,
17320 denial: MutationDenial,
17321 ) -> ToolSecurityConfig {
17322 let workspace = workspace.to_string_lossy().into_owned();
17323 let mut policy = ai_agents_tools::ToolPolicyConfig {
17324 read_paths: vec![workspace.clone()],
17325 write_paths: vec![workspace.clone()],
17326 ..Default::default()
17327 };
17328 match denial {
17329 MutationDenial::Policy => policy.blocked_paths = vec![workspace],
17330 MutationDenial::Approval => policy.require_confirmation = true,
17331 }
17332
17333 let mut security = ToolSecurityConfig {
17334 enabled: true,
17335 fail_closed: true,
17336 ..Default::default()
17337 };
17338 security.tools.insert(tool_id.to_string(), policy);
17339 security
17340 }
17341
17342 async fn assert_path_mutation_denied(tool: Arc<dyn Tool>, denial: MutationDenial) {
17343 let workspace = MutationTestWorkspace::new();
17344 let tool_id = tool.id().to_string();
17345 let preserved = workspace.root.join(format!("{}-preserved.txt", tool_id));
17346 let destination = workspace.root.join(format!("{}-destination.txt", tool_id));
17347 std::fs::write(&preserved, "preserved").unwrap();
17348 let arguments = match tool_id.as_str() {
17349 "copy_path" | "move_path" => serde_json::json!({
17350 "source_path": preserved.to_string_lossy(),
17351 "destination_path": destination.to_string_lossy(),
17352 "dry_run": false
17353 }),
17354 "delete_path" => serde_json::json!({
17355 "path": preserved.to_string_lossy(),
17356 "recursive": false,
17357 "dry_run": false
17358 }),
17359 _ => panic!("unsupported mutation tool: {}", tool_id),
17360 };
17361 let security = mutation_denial_security_config(&tool_id, &workspace.root, denial);
17362 let builder = AgentBuilder::new()
17363 .system_prompt("Test mutation denial.")
17364 .llm(Arc::new(mock_with_response("done")))
17365 .tool(tool)
17366 .tool_security(ToolSecurityEngine::new(security));
17367 let builder = match denial {
17368 MutationDenial::Policy => builder,
17369 MutationDenial::Approval => builder
17370 .hitl_engine(HITLEngine::new(ai_agents_hitl::HITLConfig::default()))
17371 .approval_handler(Arc::new(RejectAllHandler::new())),
17372 };
17373 let agent = builder.build().unwrap();
17374
17375 let record = agent
17376 .invoke_tool(ToolExecutionRequest::new(
17377 format!("{}-denied", tool_id),
17378 tool_id.clone(),
17379 arguments,
17380 ToolCallSource::Manual,
17381 ))
17382 .await
17383 .unwrap();
17384
17385 assert!(!record.executed, "{} must not be invoked", tool_id);
17386 assert!(!record.success);
17387 match denial {
17388 MutationDenial::Policy => {
17389 assert_eq!(record.policy.outcome, PermissionOutcome::Deny);
17390 assert!(record.approval.as_ref().is_some_and(|approval| matches!(
17391 &approval.status,
17392 ToolApprovalStatus::NotRequired
17393 )));
17394 }
17395 MutationDenial::Approval => {
17396 assert_eq!(record.policy.outcome, PermissionOutcome::RequiresApproval);
17397 assert!(record.approval.as_ref().is_some_and(|approval| matches!(
17398 &approval.status,
17399 ToolApprovalStatus::Rejected
17400 )));
17401 }
17402 }
17403 assert_eq!(std::fs::read_to_string(&preserved).unwrap(), "preserved");
17404 assert!(!destination.exists());
17405 }
17406
17407 fn recovery_manager_with_fallbacks(
17408 fallbacks: impl IntoIterator<Item = (String, String)>,
17409 ) -> RecoveryManager {
17410 use ai_agents_recovery::{ErrorRecoveryConfig, ToolRecoveryConfig, ToolRetryConfig};
17411
17412 let per_tool = fallbacks
17413 .into_iter()
17414 .map(|(tool, fallback_tool)| {
17415 (
17416 tool,
17417 ToolRetryConfig {
17418 max_retries: 0,
17419 timeout_ms: Some(1_000),
17420 on_failure: ToolFailureAction::Fallback { fallback_tool },
17421 },
17422 )
17423 })
17424 .collect();
17425 RecoveryManager::new(ErrorRecoveryConfig {
17426 tools: ToolRecoveryConfig {
17427 per_tool,
17428 ..Default::default()
17429 },
17430 ..Default::default()
17431 })
17432 }
17433
17434 fn approval_check() -> HITLCheckResult {
17435 HITLCheckResult::required(
17436 ApprovalTrigger::tool("test", serde_json::json!({})),
17437 HashMap::new(),
17438 "Approve?",
17439 None,
17440 )
17441 }
17442
17443 fn agent_with_approval_result(
17444 raw_result: ApprovalResult,
17445 timeout_action: TimeoutAction,
17446 hooks: Arc<ApprovalRecordingHooks>,
17447 ) -> RuntimeAgent {
17448 use ai_agents_hitl::{CallbackHandler, HITLConfig};
17449
17450 let config = HITLConfig {
17451 on_timeout: timeout_action,
17452 ..Default::default()
17453 };
17454 let handler = CallbackHandler::new(move |_| raw_result.clone());
17455 AgentBuilder::new()
17456 .system_prompt("Test HITL hooks.")
17457 .llm(Arc::new(mock_with_response("done")))
17458 .build()
17459 .unwrap()
17460 .with_hooks(hooks)
17461 .with_hitl(HITLEngine::new(config), Arc::new(handler))
17462 }
17463
17464 #[tokio::test]
17465 async fn approval_hooks_expose_direct_effective_decisions_after_raw_results() {
17466 let cases = vec![
17467 (ApprovalResult::Approved, "approved"),
17468 (
17469 ApprovalResult::Rejected {
17470 reason: Some("denied".to_string()),
17471 },
17472 "rejected",
17473 ),
17474 (
17475 ApprovalResult::Modified {
17476 changes: HashMap::from([("value".to_string(), serde_json::json!(2))]),
17477 },
17478 "modified",
17479 ),
17480 ];
17481
17482 for (raw_result, expected) in cases {
17483 let hooks = Arc::new(ApprovalRecordingHooks::new());
17484 let agent =
17485 agent_with_approval_result(raw_result, TimeoutAction::Reject, hooks.clone());
17486
17487 let result = agent.request_hitl_approval(approval_check()).await.unwrap();
17488
17489 assert_eq!(approval_result_name(&result), expected);
17490 assert_correlated_approval_events(&hooks.events(), expected, expected);
17491 }
17492 }
17493
17494 #[tokio::test]
17495 async fn approval_hooks_expose_timeout_policy_decisions() {
17496 for (timeout_action, expected) in [
17497 (TimeoutAction::Approve, "approved"),
17498 (TimeoutAction::Reject, "rejected"),
17499 ] {
17500 let hooks = Arc::new(ApprovalRecordingHooks::new());
17501 let agent =
17502 agent_with_approval_result(ApprovalResult::Timeout, timeout_action, hooks.clone());
17503
17504 let result = agent.request_hitl_approval(approval_check()).await.unwrap();
17505
17506 assert_eq!(approval_result_name(&result), expected);
17507 assert_correlated_approval_events(&hooks.events(), "timeout", expected);
17508 }
17509 }
17510
17511 #[tokio::test]
17512 async fn timeout_error_fires_correlated_resolved_error_before_returning() {
17513 let hooks = Arc::new(ApprovalRecordingHooks::new());
17514 let agent = agent_with_approval_result(
17515 ApprovalResult::Timeout,
17516 TimeoutAction::Error,
17517 hooks.clone(),
17518 );
17519
17520 let error = agent
17521 .request_hitl_approval(approval_check())
17522 .await
17523 .unwrap_err();
17524
17525 assert!(error.to_string().contains("HITL approval timeout"));
17526 assert_correlated_approval_events(&hooks.events(), "timeout", "error");
17527 }
17528
17529 #[tokio::test]
17531 async fn test_integration_yaml_to_chat_basic() {
17532 let mock = mock_with_response("Hello! How can I help you?");
17533 let agent = AgentBuilder::new()
17534 .system_prompt("You are a test assistant.")
17535 .llm(Arc::new(mock))
17536 .build()
17537 .unwrap();
17538
17539 let response = agent.chat("Hi").await.unwrap();
17540 assert!(!response.content.is_empty());
17541 assert_eq!(response.content, "Hello! How can I help you?");
17542 }
17543
17544 #[tokio::test]
17545 async fn stream_events_emit_one_authoritative_final_without_legacy_done() {
17546 let agent = AgentBuilder::new()
17547 .system_prompt("You are a test assistant.")
17548 .llm(Arc::new(mock_with_response(
17549 "Hello from the final response.",
17550 )))
17551 .build()
17552 .unwrap();
17553
17554 let mut stream = agent.chat_stream_events("Hi").await.unwrap();
17555 let mut final_responses = Vec::new();
17556 let mut legacy_done = 0;
17557 while let Some(event) = stream.next().await {
17558 match event {
17559 AgentStreamEvent::Chunk(StreamChunk::Done {}) => legacy_done += 1,
17560 AgentStreamEvent::Chunk(StreamChunk::Error { message }) => {
17561 panic!("unexpected stream error: {message}")
17562 }
17563 AgentStreamEvent::Final(response) => final_responses.push(response),
17564 AgentStreamEvent::Chunk(_) => {}
17565 }
17566 }
17567
17568 assert_eq!(legacy_done, 0);
17569 assert_eq!(final_responses.len(), 1);
17570 let response = final_responses.pop().unwrap();
17571 assert_eq!(response.content, "Hello from the final response.");
17572 assert!(
17573 response
17574 .metadata
17575 .as_ref()
17576 .is_some_and(|metadata| { metadata.contains_key("reasoning") })
17577 );
17578 }
17579
17580 #[tokio::test]
17581 async fn stream_final_content_includes_output_processing_after_provisional_chunks() {
17582 let yaml = r#"
17583name: ProcessedStreamAgent
17584system_prompt: "Answer directly."
17585process:
17586 output:
17587 - type: format
17588 config:
17589 template: "{{ response }} [finalized]"
17590streaming:
17591 enabled: true
17592"#;
17593 let agent = AgentBuilder::from_yaml(yaml)
17594 .unwrap()
17595 .llm(Arc::new(mock_with_response("provisional answer")))
17596 .auto_configure_features()
17597 .unwrap()
17598 .build()
17599 .unwrap();
17600
17601 let mut stream = agent.chat_stream_events("Hi").await.unwrap();
17602 let mut provisional = String::new();
17603 let mut final_content = None;
17604 while let Some(event) = stream.next().await {
17605 match event {
17606 AgentStreamEvent::Chunk(StreamChunk::Content { text }) => {
17607 provisional.push_str(&text)
17608 }
17609 AgentStreamEvent::Chunk(StreamChunk::Error { message }) => {
17610 panic!("unexpected stream error: {message}")
17611 }
17612 AgentStreamEvent::Final(response) => final_content = Some(response.content),
17613 AgentStreamEvent::Chunk(_) => {}
17614 }
17615 }
17616
17617 assert_eq!(provisional, "provisional answer");
17618 assert_eq!(
17619 final_content.as_deref(),
17620 Some("provisional answer [finalized]")
17621 );
17622 }
17623
17624 #[tokio::test]
17625 async fn stream_events_preserve_tool_progress_and_final_tool_calls() {
17626 let agent = AgentBuilder::new()
17627 .system_prompt("Use the echo tool once, then answer.")
17628 .llm(Arc::new(mock_with_responses(vec![
17629 r#"{"tool":"echo","arguments":{"message":"hello"}}"#,
17630 "Echo completed.",
17631 ])))
17632 .tool(Arc::new(ai_agents_tools::EchoTool::new()))
17633 .build()
17634 .unwrap();
17635
17636 let mut stream = agent.chat_stream_events("echo hello").await.unwrap();
17637 let mut starts = 0;
17638 let mut results = 0;
17639 let mut ends = 0;
17640 let mut final_response = None;
17641 while let Some(event) = stream.next().await {
17642 match event {
17643 AgentStreamEvent::Chunk(StreamChunk::ToolCallStart { name, .. }) => {
17644 assert_eq!(name, "echo");
17645 starts += 1;
17646 }
17647 AgentStreamEvent::Chunk(StreamChunk::ToolResult { name, success, .. }) => {
17648 assert_eq!(name, "echo");
17649 assert!(success);
17650 results += 1;
17651 }
17652 AgentStreamEvent::Chunk(StreamChunk::ToolCallEnd { .. }) => ends += 1,
17653 AgentStreamEvent::Chunk(StreamChunk::Error { message }) => {
17654 panic!("unexpected stream error: {message}")
17655 }
17656 AgentStreamEvent::Final(response) => final_response = Some(response),
17657 AgentStreamEvent::Chunk(_) => {}
17658 }
17659 }
17660
17661 assert_eq!((starts, results, ends), (1, 1, 1));
17662 let response = final_response.expect("tool stream must finalize");
17663 assert_eq!(response.content, "Echo completed.");
17664 assert_eq!(
17665 response.tool_calls.as_ref().map(|calls| calls
17666 .iter()
17667 .map(|call| call.name.as_str())
17668 .collect::<Vec<_>>()),
17669 Some(vec!["echo"])
17670 );
17671 }
17672
17673 #[tokio::test]
17674 async fn legacy_stream_still_emits_one_done_chunk() {
17675 let agent = AgentBuilder::new()
17676 .system_prompt("You are a test assistant.")
17677 .llm(Arc::new(mock_with_response(
17678 "Hello from the legacy stream.",
17679 )))
17680 .build()
17681 .unwrap();
17682
17683 let mut stream = agent.chat_stream("Hi").await.unwrap();
17684 let mut done = 0;
17685 while let Some(chunk) = stream.next().await {
17686 match chunk {
17687 StreamChunk::Done {} => done += 1,
17688 StreamChunk::Error { message } => panic!("unexpected stream error: {message}"),
17689 _ => {}
17690 }
17691 }
17692
17693 assert_eq!(done, 1);
17694 }
17695
17696 #[tokio::test]
17698 async fn test_integration_multi_turn_conversation() {
17699 let mock = mock_with_responses(vec![
17700 "Hello! I'm your assistant.",
17701 "The weather is sunny today.",
17702 "Goodbye!",
17703 ]);
17704 let agent = AgentBuilder::new()
17705 .system_prompt("You are helpful.")
17706 .llm(Arc::new(mock))
17707 .build()
17708 .unwrap();
17709
17710 let r1 = agent.chat("Hi").await.unwrap();
17711 assert_eq!(r1.content, "Hello! I'm your assistant.");
17712
17713 let r2 = agent.chat("What's the weather?").await.unwrap();
17714 assert_eq!(r2.content, "The weather is sunny today.");
17715
17716 let r3 = agent.chat("Bye").await.unwrap();
17717 assert_eq!(r3.content, "Goodbye!");
17718
17719 let messages = agent.memory.get_messages(None).await.unwrap();
17721 assert_eq!(messages.len(), 6);
17723 }
17724
17725 #[test]
17726 fn later_approval_preserves_modified_evidence() {
17727 let arguments = serde_json::json!({"dry_run": true});
17728 let mut record = Some(ToolApprovalRecord {
17729 status: ToolApprovalStatus::Modified,
17730 reason: None,
17731 modified_arguments: Some(arguments.clone()),
17732 });
17733
17734 merge_approved_record(&mut record);
17735
17736 let record = record.unwrap();
17737 assert!(matches!(record.status, ToolApprovalStatus::Modified));
17738 assert_eq!(record.modified_arguments, Some(arguments));
17739 }
17740
17741 #[test]
17742 fn approval_binding_rejects_replaced_tool_implementation() {
17743 let reviewed_tool: Arc<dyn ai_agents_core::Tool> = Arc::new(ContextEchoTool);
17744 let same_tool = Arc::clone(&reviewed_tool);
17745 let replacement_tool: Arc<dyn ai_agents_core::Tool> = Arc::new(ContextEchoTool);
17746 let arguments = serde_json::json!({"path": "."});
17747 let versions = ToolDecisionVersions {
17748 policy: 2,
17749 registry: 3,
17750 runtime_control: 4,
17751 state: Some(5),
17752 };
17753 let binding = ToolApprovalBinding {
17754 canonical_id: "context_echo".to_string(),
17755 arguments: arguments.clone(),
17756 confirmation_required: true,
17757 policy_version: versions.policy,
17758 runtime_control_version: versions.runtime_control,
17759 state_generation: versions.state,
17760 reviewed_tool,
17761 };
17762
17763 assert!(!binding.is_stale("context_echo", &arguments, true, versions, &same_tool,));
17764 assert!(binding.is_stale(
17765 "context_echo",
17766 &arguments,
17767 true,
17768 versions,
17769 &replacement_tool,
17770 ));
17771 }
17772
17773 #[tokio::test]
17774 async fn approved_mutation_to_dry_run_remains_executable() {
17775 use ai_agents_hitl::CallbackHandler;
17776
17777 let handler = CallbackHandler::new(|_| ApprovalResult::Modified {
17778 changes: HashMap::from([("dry_run".to_string(), serde_json::json!(true))]),
17779 });
17780 let agent = AgentBuilder::new()
17781 .system_prompt("Test safer approval modifications.")
17782 .llm(Arc::new(mock_with_response("done")))
17783 .tool(Arc::new(ai_agents_tools::FileWriteTool::new()))
17784 .hitl_engine(HITLEngine::new(ai_agents_hitl::HITLConfig::default()))
17785 .approval_handler(Arc::new(handler))
17786 .build()
17787 .unwrap();
17788
17789 let record = agent
17790 .invoke_tool(ToolExecutionRequest::new(
17791 "approved-dry-run",
17792 "file_write",
17793 serde_json::json!({
17794 "path": "./approval-dry-run.txt",
17795 "content": "not written"
17796 }),
17797 ToolCallSource::Manual,
17798 ))
17799 .await
17800 .unwrap();
17801
17802 assert!(record.executed);
17803 assert!(record.success);
17804 assert_eq!(record.executed_arguments["dry_run"], true);
17805 assert!(matches!(
17806 record.approval.as_ref().map(|approval| &approval.status),
17807 Some(ToolApprovalStatus::Modified)
17808 ));
17809 let output: Value = serde_json::from_str(&record.output).unwrap();
17810 assert_eq!(output["mutation_performed"], false);
17811 }
17812
17813 #[tokio::test]
17815 async fn shared_executor_approval_reaches_web_fetch_transport() {
17816 use ai_agents_hitl::{CallbackHandler, HITLConfig};
17817 use ai_agents_tools::{DomainPolicyConfig, ToolPolicyConfig};
17818
17819 let calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
17820 let tool = WebFetchTool::with_transport_and_resolver(
17821 Arc::new(RuntimeWebFetchTransport {
17822 calls: Arc::clone(&calls),
17823 }),
17824 Arc::new(RuntimeWebFetchResolver),
17825 );
17826 let mut security = ToolSecurityConfig {
17827 enabled: true,
17828 fail_closed: true,
17829 ..Default::default()
17830 };
17831 security.tools.insert(
17832 "web_fetch".to_string(),
17833 ToolPolicyConfig {
17834 domains: DomainPolicyConfig {
17835 requires_approval: vec!["approval.test".to_string()],
17836 ..Default::default()
17837 },
17838 allowed_schemes: vec!["https".to_string()],
17839 allowed_ports: vec![443],
17840 ..Default::default()
17841 },
17842 );
17843 let handler = CallbackHandler::new(|_| ApprovalResult::Approved);
17844 let agent = AgentBuilder::new()
17845 .system_prompt("Test approved web fetch execution.")
17846 .llm(Arc::new(mock_with_response("done")))
17847 .tool(Arc::new(tool))
17848 .tool_security(ToolSecurityEngine::new(security))
17849 .build()
17850 .unwrap()
17851 .with_hitl(HITLEngine::new(HITLConfig::default()), Arc::new(handler));
17852
17853 let record = agent
17854 .invoke_tool(ToolExecutionRequest::new(
17855 "approved-web-fetch",
17856 "web_fetch",
17857 serde_json::json!({
17858 "url": "https://approval.test/page",
17859 "cache_ttl_seconds": 0
17860 }),
17861 ToolCallSource::Manual,
17862 ))
17863 .await
17864 .unwrap();
17865
17866 assert!(record.success);
17867 assert!(
17868 record
17869 .approval
17870 .as_ref()
17871 .is_some_and(|approval| matches!(approval.status, ToolApprovalStatus::Approved))
17872 );
17873 assert_eq!(calls.load(Ordering::SeqCst), 1);
17874 }
17875
17876 #[tokio::test]
17877 async fn context_preserves_requested_and_canonical_identity() {
17878 let mock = mock_with_response("hello");
17879 let mut tools = ai_agents_tools::ToolRegistry::new();
17880 tools.register(Arc::new(ContextEchoTool)).unwrap();
17881
17882 let mut security = ToolSecurityConfig {
17883 enabled: true,
17884 fail_closed: true,
17885 ..Default::default()
17886 };
17887 let mut policy = ai_agents_tools::ToolPolicyConfig {
17888 read_paths: vec![".".to_string()],
17889 max_results: Some(7),
17890 ..Default::default()
17891 };
17892 policy
17893 .config
17894 .insert("backend".to_string(), serde_json::json!("memory"));
17895 security.tools.insert("context_echo".to_string(), policy);
17896
17897 let agent = AgentBuilder::new()
17898 .system_prompt("You are helpful.")
17899 .llm(Arc::new(mock))
17900 .tools(tools)
17901 .tool_security(ToolSecurityEngine::new(security))
17902 .build()
17903 .unwrap();
17904
17905 let record = agent
17906 .invoke_tool(ToolExecutionRequest::new(
17907 "ctx-call",
17908 "Context Echo",
17909 serde_json::json!({"path": ".", "max_results": 99}),
17910 ToolCallSource::Manual,
17911 ))
17912 .await
17913 .unwrap();
17914
17915 assert!(record.success);
17916 assert!(matches!(&record.source, ToolCallSource::Manual));
17917 assert_eq!(record.requested_name, "Context Echo");
17918 assert_eq!(record.canonical_id, "context_echo");
17919 assert_eq!(record.policy.outcome, PermissionOutcome::Allow);
17920 assert_eq!(record.executed_arguments["max_results"], 7);
17921 let output: Value = serde_json::from_str(&record.output).unwrap();
17922 assert_eq!(output["requested_name"], "Context Echo");
17923 assert_eq!(output["canonical_id"], "context_echo");
17924 assert_eq!(output["max_results"], 7);
17925 assert_eq!(output["custom_config"]["backend"], "memory");
17926 assert!(record.metadata.contains_key("effective_limits"));
17927 assert!(record.metadata.contains_key("policy_snapshot"));
17928 }
17929
17930 #[tokio::test]
17931 async fn test_runtime_control_cancels_active_tool_call() {
17932 let mock = mock_with_response("hello");
17933 let agent = Arc::new(
17934 AgentBuilder::new()
17935 .system_prompt("You are helpful.")
17936 .llm(Arc::new(mock))
17937 .tool(Arc::new(SlowTool))
17938 .build()
17939 .unwrap(),
17940 );
17941 let control = agent.runtime_control();
17942 let running_agent = Arc::clone(&agent);
17943 let handle = tokio::spawn(async move {
17944 running_agent
17945 .invoke_tool(ToolExecutionRequest::new(
17946 "slow-call",
17947 "slow",
17948 serde_json::json!({}),
17949 ToolCallSource::Manual,
17950 ))
17951 .await
17952 .unwrap()
17953 });
17954
17955 tokio::time::sleep(std::time::Duration::from_millis(100)).await;
17956 control.cancel_all();
17957 let record = handle.await.unwrap();
17958
17959 assert!(record.executed);
17960 assert!(record.cancelled);
17961 assert!(!record.success);
17962 assert!(record.cancellation_reason.is_some());
17963 }
17964
17965 #[tokio::test]
17967 async fn cancelled_tool_does_not_enter_fallback() {
17968 let fallback_calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
17969 let agent = Arc::new(
17970 AgentBuilder::new()
17971 .system_prompt("Test cancellation before fallback.")
17972 .llm(Arc::new(mock_with_response("done")))
17973 .tool(Arc::new(SlowTool))
17974 .tool(Arc::new(RecoveryTestTool {
17975 id: "fallback".to_string(),
17976 succeeds: true,
17977 calls: Arc::clone(&fallback_calls),
17978 max_output_chars: None,
17979 }))
17980 .recovery_manager(recovery_manager_with_fallbacks([(
17981 "slow".to_string(),
17982 "fallback".to_string(),
17983 )]))
17984 .build()
17985 .unwrap(),
17986 );
17987 let control = agent.runtime_control();
17988 let running_agent = Arc::clone(&agent);
17989 let handle = tokio::spawn(async move {
17990 running_agent
17991 .invoke_tool(ToolExecutionRequest::new(
17992 "cancelled-fallback-call",
17993 "slow",
17994 serde_json::json!({}),
17995 ToolCallSource::Manual,
17996 ))
17997 .await
17998 .unwrap()
17999 });
18000
18001 tokio::time::sleep(Duration::from_millis(100)).await;
18002 control.cancel_all();
18003 let record = handle.await.unwrap();
18004
18005 assert!(record.executed);
18006 assert!(record.cancelled);
18007 assert!(!record.success);
18008 assert_eq!(record.canonical_id, "slow");
18009 assert_eq!(fallback_calls.load(Ordering::SeqCst), 0);
18010 assert_eq!(agent.tool_call_history().len(), 1);
18011 }
18012
18013 #[tokio::test]
18014 async fn non_idempotent_tool_calls_are_not_retried() {
18015 use ai_agents_recovery::{ErrorRecoveryConfig, ToolRecoveryConfig, ToolRetryConfig};
18016
18017 let mock = mock_with_response("hello");
18018 let calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
18019 let agent = AgentBuilder::new()
18020 .system_prompt("You are helpful.")
18021 .llm(Arc::new(mock))
18022 .tool(Arc::new(FlakyWriteTool {
18023 calls: Arc::clone(&calls),
18024 }))
18025 .recovery_manager(RecoveryManager::new(ErrorRecoveryConfig {
18026 tools: ToolRecoveryConfig {
18027 default: ToolRetryConfig {
18028 max_retries: 2,
18029 ..Default::default()
18030 },
18031 ..Default::default()
18032 },
18033 ..Default::default()
18034 }))
18035 .build()
18036 .unwrap();
18037
18038 let record = agent
18039 .invoke_tool(ToolExecutionRequest::new(
18040 "flaky-call",
18041 "flaky_write",
18042 serde_json::json!({"path": "./tmp.txt"}),
18043 ToolCallSource::Manual,
18044 ))
18045 .await
18046 .unwrap();
18047
18048 assert!(!record.success);
18049 assert_eq!(calls.load(Ordering::SeqCst), 1);
18050 }
18051
18052 #[tokio::test]
18053 async fn safely_retryable_tool_receives_a_fresh_deadline_per_attempt() {
18054 use ai_agents_recovery::{ErrorRecoveryConfig, ToolRecoveryConfig, ToolRetryConfig};
18055
18056 let calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
18057 let deadlines = Arc::new(parking_lot::Mutex::new(Vec::new()));
18058 let remaining_ms = Arc::new(parking_lot::Mutex::new(Vec::new()));
18059 let agent = AgentBuilder::new()
18060 .system_prompt("Test retry deadlines.")
18061 .llm(Arc::new(mock_with_response("done")))
18062 .tool(Arc::new(RetryDeadlineTool {
18063 calls: Arc::clone(&calls),
18064 deadlines: Arc::clone(&deadlines),
18065 remaining_ms: Arc::clone(&remaining_ms),
18066 }))
18067 .recovery_manager(RecoveryManager::new(ErrorRecoveryConfig {
18068 tools: ToolRecoveryConfig {
18069 per_tool: HashMap::from([(
18070 "retry_deadline".to_string(),
18071 ToolRetryConfig {
18072 max_retries: 1,
18073 ..Default::default()
18074 },
18075 )]),
18076 ..Default::default()
18077 },
18078 ..Default::default()
18079 }))
18080 .build()
18081 .unwrap();
18082
18083 let record = agent
18084 .invoke_tool(ToolExecutionRequest::new(
18085 "retry-deadline-call",
18086 "retry_deadline",
18087 serde_json::json!({}),
18088 ToolCallSource::Manual,
18089 ))
18090 .await
18091 .unwrap();
18092
18093 assert!(record.executed);
18094 assert!(record.success);
18095 assert_eq!(calls.load(Ordering::SeqCst), 2);
18096 let deadlines = deadlines.lock();
18097 assert_eq!(deadlines.len(), 2);
18098 assert!(
18099 deadlines[1] > deadlines[0],
18100 "retry inherited the first invocation deadline"
18101 );
18102 let remaining_ms = remaining_ms.lock();
18103 assert_eq!(remaining_ms.len(), 2);
18104 assert!(
18105 remaining_ms
18106 .iter()
18107 .all(|remaining| (800..=1_000).contains(remaining))
18108 );
18109 }
18110
18111 #[tokio::test]
18113 async fn call_classification_timeout_controls_deadline_and_timer() {
18114 let calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
18115 let remaining_ms = Arc::new(parking_lot::Mutex::new(Vec::new()));
18116 let agent = AgentBuilder::new()
18117 .system_prompt("Test call-level timeout.")
18118 .llm(Arc::new(mock_with_response("done")))
18119 .tool(Arc::new(ClassifiedTimeoutTool {
18120 id: "classified_timeout",
18121 calls: Arc::clone(&calls),
18122 timeout_ms: 100,
18123 sleep_ms: 150,
18124 requires_approval: false,
18125 remaining_ms: Arc::clone(&remaining_ms),
18126 }))
18127 .build()
18128 .unwrap();
18129
18130 let started = Instant::now();
18131 let record = agent
18132 .invoke_tool(ToolExecutionRequest::new(
18133 "classified-timeout-call",
18134 "classified_timeout",
18135 serde_json::json!({}),
18136 ToolCallSource::Manual,
18137 ))
18138 .await
18139 .unwrap();
18140
18141 assert!(record.executed);
18142 assert!(record.timed_out);
18143 assert!(!record.success);
18144 assert_eq!(calls.load(Ordering::SeqCst), 1);
18145 assert!(started.elapsed() < Duration::from_secs(1));
18146 let remaining_ms = remaining_ms.lock();
18147 assert_eq!(remaining_ms.len(), 1);
18148 assert!((1..=100).contains(&remaining_ms[0]));
18149 }
18150
18151 #[tokio::test]
18153 async fn recovery_timeout_only_lowers_call_and_policy_timeouts() {
18154 use ai_agents_recovery::{ErrorRecoveryConfig, ToolRecoveryConfig, ToolRetryConfig};
18155
18156 let calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
18157 let remaining_ms = Arc::new(parking_lot::Mutex::new(Vec::new()));
18158 let agent = AgentBuilder::new()
18159 .system_prompt("Test recovery timeout.")
18160 .llm(Arc::new(mock_with_response("done")))
18161 .tool(Arc::new(ClassifiedTimeoutTool {
18162 id: "recovery_timeout",
18163 calls: Arc::clone(&calls),
18164 timeout_ms: 1_000,
18165 sleep_ms: 150,
18166 requires_approval: false,
18167 remaining_ms: Arc::clone(&remaining_ms),
18168 }))
18169 .recovery_manager(RecoveryManager::new(ErrorRecoveryConfig {
18170 tools: ToolRecoveryConfig {
18171 per_tool: HashMap::from([(
18172 "recovery_timeout".to_string(),
18173 ToolRetryConfig {
18174 timeout_ms: Some(100),
18175 ..Default::default()
18176 },
18177 )]),
18178 ..Default::default()
18179 },
18180 ..Default::default()
18181 }))
18182 .build()
18183 .unwrap();
18184
18185 let started = Instant::now();
18186 let record = agent
18187 .invoke_tool(ToolExecutionRequest::new(
18188 "recovery-timeout-call",
18189 "recovery_timeout",
18190 serde_json::json!({}),
18191 ToolCallSource::Manual,
18192 ))
18193 .await
18194 .unwrap();
18195
18196 assert!(record.executed);
18197 assert!(record.timed_out);
18198 assert!(!record.success);
18199 assert_eq!(calls.load(Ordering::SeqCst), 1);
18200 assert!(started.elapsed() < Duration::from_secs(1));
18201 assert_eq!(record.metadata["effective_limits"]["timeout_ms"], 100);
18202 let remaining_ms = remaining_ms.lock();
18203 assert_eq!(remaining_ms.len(), 1);
18204 assert!((1..=100).contains(&remaining_ms[0]));
18205 }
18206
18207 #[tokio::test]
18209 async fn recovery_default_timeout_controls_deadline_and_timer() {
18210 use ai_agents_recovery::{ErrorRecoveryConfig, ToolRecoveryConfig, ToolRetryConfig};
18211
18212 let calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
18213 let remaining_ms = Arc::new(parking_lot::Mutex::new(Vec::new()));
18214 let agent = AgentBuilder::new()
18215 .system_prompt("Test default recovery timeout.")
18216 .llm(Arc::new(mock_with_response("done")))
18217 .tool(Arc::new(ClassifiedTimeoutTool {
18218 id: "default_recovery_timeout",
18219 calls: Arc::clone(&calls),
18220 timeout_ms: 1_000,
18221 sleep_ms: 150,
18222 requires_approval: false,
18223 remaining_ms: Arc::clone(&remaining_ms),
18224 }))
18225 .recovery_manager(RecoveryManager::new(ErrorRecoveryConfig {
18226 tools: ToolRecoveryConfig {
18227 default: ToolRetryConfig {
18228 timeout_ms: Some(100),
18229 ..Default::default()
18230 },
18231 ..Default::default()
18232 },
18233 ..Default::default()
18234 }))
18235 .build()
18236 .unwrap();
18237
18238 let started = Instant::now();
18239 let record = agent
18240 .invoke_tool(ToolExecutionRequest::new(
18241 "default-recovery-timeout-call",
18242 "default_recovery_timeout",
18243 serde_json::json!({}),
18244 ToolCallSource::Manual,
18245 ))
18246 .await
18247 .unwrap();
18248
18249 assert!(record.executed);
18250 assert!(record.timed_out);
18251 assert!(!record.success);
18252 assert_eq!(calls.load(Ordering::SeqCst), 1);
18253 assert!(started.elapsed() < Duration::from_secs(1));
18254 assert_eq!(record.metadata["effective_limits"]["timeout_ms"], 100);
18255 let remaining_ms = remaining_ms.lock();
18256 assert_eq!(remaining_ms.len(), 1);
18257 assert!((1..=100).contains(&remaining_ms[0]));
18258 }
18259
18260 #[test]
18262 fn recovery_timeout_cannot_widen_security_baseline() {
18263 let security_engine = ToolSecurityEngine::new(ToolSecurityConfig {
18264 default_timeout_ms: 100,
18265 ..Default::default()
18266 });
18267 let safety = ToolSafetyMetadata::compute();
18268 let mut classification = ToolCallClassification::from_metadata(&safety);
18269 classification.timeout_ms = Some(500);
18270
18271 let (limits, timeout) = RuntimeAgent::effective_tool_limits(
18272 &security_engine,
18273 "recovery_cannot_widen",
18274 &safety,
18275 &classification,
18276 Some(1_000),
18277 )
18278 .unwrap();
18279
18280 assert_eq!(limits.timeout_ms, Some(100));
18281 assert_eq!(timeout.timer, Duration::from_millis(100));
18282 }
18283
18284 #[tokio::test]
18286 async fn invalid_call_timeout_stops_before_approval_or_tool_invocation() {
18287 let tool_calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
18288 let approval_calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
18289 let remaining_ms = Arc::new(parking_lot::Mutex::new(Vec::new()));
18290 let mut security = ToolSecurityConfig {
18291 enabled: true,
18292 ..Default::default()
18293 };
18294 security.tools.insert(
18295 "invalid_call_timeout".to_string(),
18296 ai_agents_tools::ToolPolicyConfig {
18297 require_confirmation: true,
18298 ..Default::default()
18299 },
18300 );
18301 let agent = AgentBuilder::new()
18302 .system_prompt("Test invalid call timeout.")
18303 .llm(Arc::new(mock_with_response("done")))
18304 .tool(Arc::new(ClassifiedTimeoutTool {
18305 id: "invalid_call_timeout",
18306 calls: Arc::clone(&tool_calls),
18307 timeout_ms: u64::MAX,
18308 sleep_ms: 0,
18309 requires_approval: false,
18310 remaining_ms,
18311 }))
18312 .tool_security(ToolSecurityEngine::new(security))
18313 .hitl_engine(HITLEngine::new(ai_agents_hitl::HITLConfig::default()))
18314 .approval_handler(Arc::new(CountingApprovalHandler {
18315 calls: Arc::clone(&approval_calls),
18316 }))
18317 .build()
18318 .unwrap();
18319
18320 let error = agent
18321 .invoke_tool(ToolExecutionRequest::new(
18322 "invalid-call-timeout",
18323 "invalid_call_timeout",
18324 serde_json::json!({}),
18325 ToolCallSource::Manual,
18326 ))
18327 .await
18328 .unwrap_err();
18329
18330 assert!(error.to_string().contains(
18331 "effective tool timeout_ms must be no greater than 3153600000000000 milliseconds"
18332 ));
18333 assert_eq!(approval_calls.load(Ordering::SeqCst), 0);
18334 assert_eq!(tool_calls.load(Ordering::SeqCst), 0);
18335 }
18336
18337 #[tokio::test]
18339 async fn invalid_modified_call_timeout_stops_before_lock_or_invocation() {
18340 use ai_agents_hitl::CallbackHandler;
18341
18342 let blocker_gate = PathMutationGate::new();
18343 let tool_calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
18344 let hooks = Arc::new(ToolLifecycleRecordingHooks::new());
18345 let handler = CallbackHandler::new(|_| ApprovalResult::Modified {
18346 changes: HashMap::from([("invalid_timeout".to_string(), Value::Bool(true))]),
18347 });
18348 let agent = Arc::new(
18349 AgentBuilder::new()
18350 .system_prompt("Test final call timeout validation.")
18351 .llm(Arc::new(mock_with_response("done")))
18352 .tool(Arc::new(BlockingPathMutationTool {
18353 id: "timeout_lock_blocker",
18354 path_fields: vec![ai_agents_core::PathPolicyBinding::write("path")],
18355 gate: blocker_gate.clone(),
18356 }))
18357 .tool(Arc::new(ApprovalModifiedTimeoutTool {
18358 calls: Arc::clone(&tool_calls),
18359 }))
18360 .hitl_engine(HITLEngine::new(ai_agents_hitl::HITLConfig::default()))
18361 .approval_handler(Arc::new(handler))
18362 .hooks(hooks.clone())
18363 .build()
18364 .unwrap(),
18365 );
18366 let blocking_agent = Arc::clone(&agent);
18367 let blocker = tokio::spawn(async move {
18368 blocking_agent
18369 .invoke_tool(ToolExecutionRequest::new(
18370 "timeout-lock-blocker",
18371 "timeout_lock_blocker",
18372 serde_json::json!({"path": "./shared-timeout.txt"}),
18373 ToolCallSource::Manual,
18374 ))
18375 .await
18376 .unwrap()
18377 });
18378 blocker_gate.wait_until_entered().await;
18379
18380 let record = tokio::time::timeout(
18381 Duration::from_millis(500),
18382 agent.invoke_tool(ToolExecutionRequest::new(
18383 "invalid-modified-timeout",
18384 "approval_modified_timeout",
18385 serde_json::json!({
18386 "path": "./shared-timeout.txt",
18387 "invalid_timeout": false
18388 }),
18389 ToolCallSource::Manual,
18390 )),
18391 )
18392 .await
18393 .expect("final timeout validation must not wait for the held path lock")
18394 .unwrap();
18395
18396 blocker_gate.release();
18397 assert!(blocker.await.unwrap().success);
18398 assert!(!record.executed);
18399 assert!(!record.success);
18400 assert_eq!(record.policy.outcome, PermissionOutcome::Deny);
18401 assert!(record.output.contains(
18402 "effective tool timeout_ms must be no greater than 3153600000000000 milliseconds"
18403 ));
18404 assert_eq!(tool_calls.load(Ordering::SeqCst), 0);
18405 let invalid_request_events = hooks
18406 .events()
18407 .into_iter()
18408 .filter(|event| event.contains("approval_modified_timeout") || event == "error")
18409 .collect::<Vec<_>>();
18410 assert_eq!(
18411 invalid_request_events,
18412 vec![
18413 "start:approval_modified_timeout",
18414 "complete:approval_modified_timeout:false",
18415 "record:approval_modified_timeout:false",
18416 "error"
18417 ]
18418 );
18419 }
18420
18421 #[tokio::test]
18422 async fn side_effecting_tools_are_serialized_per_resource() {
18423 let mock = mock_with_response("hello");
18424 let active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
18425 let max_active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
18426 let agent = Arc::new(
18427 AgentBuilder::new()
18428 .system_prompt("You are helpful.")
18429 .llm(Arc::new(mock))
18430 .tool(Arc::new(LockedWriteTool {
18431 active: Arc::clone(&active),
18432 max_active: Arc::clone(&max_active),
18433 }))
18434 .build()
18435 .unwrap(),
18436 );
18437
18438 let left = {
18439 let agent = Arc::clone(&agent);
18440 tokio::spawn(async move {
18441 agent
18442 .invoke_tool(ToolExecutionRequest::new(
18443 "lock-1",
18444 "locked_write",
18445 serde_json::json!({"path": "./same.txt"}),
18446 ToolCallSource::Manual,
18447 ))
18448 .await
18449 .unwrap()
18450 })
18451 };
18452 let right = {
18453 let agent = Arc::clone(&agent);
18454 tokio::spawn(async move {
18455 agent
18456 .invoke_tool(ToolExecutionRequest::new(
18457 "lock-2",
18458 "locked_write",
18459 serde_json::json!({"path": "./same.txt"}),
18460 ToolCallSource::Manual,
18461 ))
18462 .await
18463 .unwrap()
18464 })
18465 };
18466
18467 let left = left.await.unwrap();
18468 let right = right.await.unwrap();
18469 assert!(left.success);
18470 assert!(right.success);
18471 assert_eq!(max_active.load(Ordering::SeqCst), 1);
18472 }
18473
18474 #[tokio::test]
18475 async fn path_resources_use_shared_global_lock_and_cleanup() {
18476 let active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
18477 let max_active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
18478 let bindings = ai_agents_core::ToolPolicyBindings {
18479 path_fields: vec![
18480 ai_agents_core::PathPolicyBinding::read_write("source_path"),
18481 ai_agents_core::PathPolicyBinding::write("destination_path"),
18482 ],
18483 ..Default::default()
18484 };
18485 let classification = ai_agents_core::ToolCallClassification::from_metadata(
18486 &MultiResourceWriteTool {
18487 active: Arc::clone(&active),
18488 max_active: Arc::clone(&max_active),
18489 }
18490 .safety_metadata(),
18491 );
18492 let left_args = serde_json::json!({
18493 "source_path": "./a/../first.txt",
18494 "destination_path": "./second.txt"
18495 });
18496 let right_args = serde_json::json!({
18497 "source_path": "./second.txt",
18498 "destination_path": "./first.txt"
18499 });
18500 let left_keys = tool_resource_lock_keys(
18501 "multi_resource_write",
18502 &left_args,
18503 &bindings,
18504 &classification,
18505 );
18506 let right_keys = tool_resource_lock_keys(
18507 "multi_resource_write",
18508 &right_args,
18509 &bindings,
18510 &classification,
18511 );
18512 assert_eq!(left_keys, right_keys);
18513 assert_eq!(left_keys, vec!["path-mutation:global".to_string()]);
18514
18515 let locks = new_tool_resource_locks();
18516 let build_agent = || {
18517 AgentBuilder::new()
18518 .system_prompt("Test shared resource locks.")
18519 .llm(Arc::new(mock_with_response("done")))
18520 .tool(Arc::new(MultiResourceWriteTool {
18521 active: Arc::clone(&active),
18522 max_active: Arc::clone(&max_active),
18523 }))
18524 .build()
18525 .unwrap()
18526 .with_shared_resource_locks(Arc::clone(&locks))
18527 };
18528 let left_agent = Arc::new(build_agent());
18529 let right_agent = Arc::new(build_agent());
18530 let left = tokio::spawn(async move {
18531 left_agent
18532 .invoke_tool(ToolExecutionRequest::new(
18533 "multi-left",
18534 "multi_resource_write",
18535 left_args,
18536 ToolCallSource::Manual,
18537 ))
18538 .await
18539 .unwrap()
18540 });
18541 let right = tokio::spawn(async move {
18542 right_agent
18543 .invoke_tool(ToolExecutionRequest::new(
18544 "multi-right",
18545 "multi_resource_write",
18546 right_args,
18547 ToolCallSource::Manual,
18548 ))
18549 .await
18550 .unwrap()
18551 });
18552 let (left, right) = tokio::time::timeout(std::time::Duration::from_secs(2), async {
18553 tokio::join!(left, right)
18554 })
18555 .await
18556 .expect("reversed resource acquisition must not deadlock");
18557
18558 assert!(left.unwrap().success);
18559 assert!(right.unwrap().success);
18560 assert_eq!(max_active.load(Ordering::SeqCst), 1);
18561 assert!(locks.read().is_empty());
18562 }
18563
18564 #[tokio::test]
18565 async fn global_path_lock_serializes_copy_destination_with_file_write() {
18566 assert_path_mutation_pair_serialized(
18567 "copy_path",
18568 CopyPathTool::new().policy_bindings().path_fields,
18569 serde_json::json!({
18570 "source_path": "./source.txt",
18571 "destination_path": "./shared.txt"
18572 }),
18573 "file_write",
18574 FileWriteTool::new().policy_bindings().path_fields,
18575 serde_json::json!({"path": "./shared.txt"}),
18576 )
18577 .await;
18578 }
18579
18580 #[tokio::test]
18581 async fn parent_and_spawned_runtime_share_global_path_lock() {
18582 let workspace = MutationTestWorkspace::new();
18583 let destination = workspace.root.join("spawned.txt");
18584 let parent_gate = PathMutationGate::new();
18585 let parent = Arc::new(
18586 AgentBuilder::from_yaml(
18587 r#"
18588name: LockParent
18589system_prompt: parent
18590llm:
18591 default: default
18592tools:
18593 - parent_path_write
18594spawner:
18595 shared_llms: true
18596"#,
18597 )
18598 .unwrap()
18599 .llm(Arc::new(mock_with_response("done")))
18600 .auto_configure_spawner()
18601 .await
18602 .unwrap()
18603 .tool(Arc::new(BlockingPathMutationTool {
18604 id: "parent_path_write",
18605 path_fields: vec![ai_agents_core::PathPolicyBinding::write("path")],
18606 gate: parent_gate.clone(),
18607 }))
18608 .build()
18609 .unwrap(),
18610 );
18611
18612 let mut child_spec = crate::spec::AgentSpec {
18613 name: "LockChild".to_string(),
18614 system_prompt: "child".to_string(),
18615 tools: Some(vec![crate::spec::ToolEntry::Simple(
18616 "file_write".to_string(),
18617 )]),
18618 ..Default::default()
18619 };
18620 child_spec.tool_security.enabled = true;
18621 child_spec.tool_security.fail_closed = true;
18622 let file_write_policy = ai_agents_tools::ToolPolicyConfig {
18623 write_paths: vec![workspace.root.to_string_lossy().into_owned()],
18624 allow_without_confirmation: true,
18625 ..Default::default()
18626 };
18627 child_spec
18628 .tool_security
18629 .tools
18630 .insert("file_write".to_string(), file_write_policy);
18631 let spawned = parent
18632 .spawner()
18633 .unwrap()
18634 .spawn_from_spec(child_spec)
18635 .await
18636 .unwrap();
18637 assert!(Arc::ptr_eq(
18638 &parent.resource_locks,
18639 &spawned.agent.resource_locks
18640 ));
18641 assert!(!Arc::ptr_eq(
18642 &parent.runtime_control,
18643 &spawned.agent.runtime_control
18644 ));
18645
18646 let parent_call = {
18647 let parent = Arc::clone(&parent);
18648 let destination = destination.clone();
18649 tokio::spawn(async move {
18650 parent
18651 .invoke_tool(ToolExecutionRequest::new(
18652 "parent-lock-holder",
18653 "parent_path_write",
18654 serde_json::json!({"path": destination}),
18655 ToolCallSource::Manual,
18656 ))
18657 .await
18658 .unwrap()
18659 })
18660 };
18661 parent_gate.wait_until_entered().await;
18662
18663 let child_call = {
18664 let child = Arc::clone(&spawned.agent);
18665 let destination = destination.clone();
18666 tokio::spawn(async move {
18667 child
18668 .invoke_tool(ToolExecutionRequest::new(
18669 "spawned-file-write",
18670 "file_write",
18671 serde_json::json!({
18672 "path": destination,
18673 "content": "spawned",
18674 "dry_run": false
18675 }),
18676 ToolCallSource::Manual,
18677 ))
18678 .await
18679 .unwrap()
18680 })
18681 };
18682 wait_for_resource_lock_strong_count(&parent.resource_locks, 2).await;
18683 assert!(!child_call.is_finished());
18684
18685 parent_gate.release();
18686 let (parent_record, child_record) =
18687 tokio::time::timeout(std::time::Duration::from_secs(2), async {
18688 tokio::join!(parent_call, child_call)
18689 })
18690 .await
18691 .expect("parent and spawned path mutations did not finish");
18692 assert!(parent_record.unwrap().success);
18693 assert!(child_record.unwrap().success);
18694 assert_eq!(std::fs::read_to_string(destination).unwrap(), "spawned");
18695 assert!(parent.resource_locks.read().is_empty());
18696 }
18697
18698 #[tokio::test]
18699 async fn cancelled_global_path_lock_waiter_does_not_retain_weak_entry() {
18700 let locks = new_tool_resource_locks();
18701 let holder_gate = PathMutationGate::new();
18702 let waiter_gate = PathMutationGate::new();
18703 waiter_gate.release();
18704 let holder = Arc::new(
18705 AgentBuilder::new()
18706 .system_prompt("Hold the global path lock.")
18707 .llm(Arc::new(mock_with_response("done")))
18708 .tool(Arc::new(BlockingPathMutationTool {
18709 id: "holder_write",
18710 path_fields: vec![ai_agents_core::PathPolicyBinding::write("path")],
18711 gate: holder_gate.clone(),
18712 }))
18713 .build()
18714 .unwrap()
18715 .with_shared_resource_locks(Arc::clone(&locks)),
18716 );
18717 let waiter = Arc::new(
18718 AgentBuilder::new()
18719 .system_prompt("Wait for the global path lock.")
18720 .llm(Arc::new(mock_with_response("done")))
18721 .tool(Arc::new(BlockingPathMutationTool {
18722 id: "waiter_write",
18723 path_fields: vec![ai_agents_core::PathPolicyBinding::write("path")],
18724 gate: waiter_gate.clone(),
18725 }))
18726 .build()
18727 .unwrap()
18728 .with_shared_resource_locks(Arc::clone(&locks)),
18729 );
18730
18731 let holder_call = {
18732 let holder = Arc::clone(&holder);
18733 tokio::spawn(async move {
18734 holder
18735 .invoke_tool(ToolExecutionRequest::new(
18736 "holder-call",
18737 "holder_write",
18738 serde_json::json!({"path": "./shared.txt"}),
18739 ToolCallSource::Manual,
18740 ))
18741 .await
18742 .unwrap()
18743 })
18744 };
18745 holder_gate.wait_until_entered().await;
18746
18747 let waiter_call = {
18748 let waiter = Arc::clone(&waiter);
18749 tokio::spawn(async move {
18750 waiter
18751 .invoke_tool(ToolExecutionRequest::new(
18752 "waiter-call",
18753 "waiter_write",
18754 serde_json::json!({"path": "./shared.txt"}),
18755 ToolCallSource::Manual,
18756 ))
18757 .await
18758 .unwrap()
18759 })
18760 };
18761 wait_for_resource_lock_strong_count(&locks, 2).await;
18762 waiter.runtime_control().cancel_all();
18763
18764 let waiter_record = tokio::time::timeout(std::time::Duration::from_secs(2), waiter_call)
18765 .await
18766 .expect("cancelled lock waiter did not finish")
18767 .unwrap();
18768 assert!(!waiter_record.success);
18769 assert!(!waiter_record.executed);
18770 assert!(waiter_record.cancelled);
18771 assert_eq!(
18772 waiter_record.cancellation_reason.as_deref(),
18773 Some("runtime control cancellation")
18774 );
18775 assert!(!waiter_gate.entered.load(Ordering::SeqCst));
18776 assert_eq!(
18777 locks
18778 .read()
18779 .get("path-mutation:global")
18780 .map_or(0, |lock| lock.strong_count()),
18781 1
18782 );
18783
18784 holder_gate.release();
18785 let holder_record = tokio::time::timeout(std::time::Duration::from_secs(2), holder_call)
18786 .await
18787 .expect("lock holder did not finish")
18788 .unwrap();
18789 assert!(holder_record.success);
18790 assert!(locks.read().is_empty());
18791 }
18792
18793 #[tokio::test]
18794 async fn path_mutation_policy_and_approval_denials_do_not_invoke_tools() {
18795 for denial in [MutationDenial::Policy, MutationDenial::Approval] {
18796 let tools: [Arc<dyn Tool>; 3] = [
18797 Arc::new(CopyPathTool::new()),
18798 Arc::new(MovePathTool::new()),
18799 Arc::new(DeletePathTool::new()),
18800 ];
18801 for tool in tools {
18802 assert_path_mutation_denied(tool, denial).await;
18803 }
18804 }
18805 }
18806
18807 #[tokio::test]
18808 async fn policy_denial_keeps_executor_hook_lifecycle_and_record_authority() {
18809 let workspace = MutationTestWorkspace::new();
18810 let target = workspace.root.join("denied.txt");
18811 let hooks = Arc::new(ToolLifecycleRecordingHooks::new());
18812 let agent = AgentBuilder::new()
18813 .system_prompt("Test denied tool hooks.")
18814 .llm(Arc::new(mock_with_response("done")))
18815 .tool(Arc::new(FileWriteTool::new()))
18816 .tool_security(ToolSecurityEngine::new(mutation_denial_security_config(
18817 "file_write",
18818 &workspace.root,
18819 MutationDenial::Policy,
18820 )))
18821 .hooks(hooks.clone())
18822 .build()
18823 .unwrap();
18824
18825 let record = agent
18826 .invoke_tool(ToolExecutionRequest::new(
18827 "denied-hook-call",
18828 "file_write",
18829 serde_json::json!({
18830 "path": target.to_string_lossy(),
18831 "content": "blocked"
18832 }),
18833 ToolCallSource::Manual,
18834 ))
18835 .await
18836 .unwrap();
18837
18838 assert!(!record.executed);
18839 assert!(!record.success);
18840 assert_eq!(record.policy.outcome, PermissionOutcome::Deny);
18841 assert_eq!(
18842 hooks.events(),
18843 vec![
18844 "start:file_write",
18845 "complete:file_write:false",
18846 "record:file_write:false",
18847 "error"
18848 ]
18849 );
18850 assert!(!target.exists());
18851 }
18852
18853 #[tokio::test]
18854 async fn approval_argument_changes_are_rechecked_against_final_scope() {
18855 let active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
18856 let max_active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
18857 let entered = Arc::new(tokio::sync::Barrier::new(2));
18858 let release = Arc::new(tokio::sync::Notify::new());
18859 let handler = Arc::new(BlockingApprovalHandler {
18860 entered: Arc::clone(&entered),
18861 release: Arc::clone(&release),
18862 result: ApprovalResult::Modified {
18863 changes: HashMap::from([(
18864 "path".to_string(),
18865 Value::String("./after-approval.txt".to_string()),
18866 )]),
18867 },
18868 });
18869 let agent = Arc::new(
18870 AgentBuilder::new()
18871 .system_prompt("Test final scope validation.")
18872 .llm(Arc::new(mock_with_response("done")))
18873 .tool(Arc::new(LockedWriteTool {
18874 active: Arc::clone(&active),
18875 max_active: Arc::clone(&max_active),
18876 }))
18877 .tool_security(ToolSecurityEngine::new(approval_security_config(true)))
18878 .hitl_engine(HITLEngine::new(ai_agents_hitl::HITLConfig::default()))
18879 .approval_handler(handler)
18880 .build()
18881 .unwrap(),
18882 );
18883 let control = agent.runtime_control();
18884 let running = Arc::clone(&agent);
18885 let call = tokio::spawn(async move {
18886 running
18887 .invoke_tool(ToolExecutionRequest::new(
18888 "approval-scope",
18889 "locked_write",
18890 serde_json::json!({"path": "./before-approval.txt"}),
18891 ToolCallSource::Manual,
18892 ))
18893 .await
18894 .unwrap()
18895 });
18896 entered.wait().await;
18897 let expected_version = control.set_tool_scope(Vec::new());
18898 release.notify_one();
18899 let record = call.await.unwrap();
18900
18901 assert!(!record.executed);
18902 assert!(!record.success);
18903 assert_eq!(record.runtime_config_version, expected_version);
18904 assert_eq!(record.executed_arguments["path"], "./after-approval.txt");
18905 assert_eq!(max_active.load(Ordering::SeqCst), 0);
18906 assert_eq!(
18907 record.metadata["runtime_scope_snapshot"],
18908 serde_json::json!([])
18909 );
18910 }
18911
18912 #[tokio::test]
18913 async fn approval_is_rechecked_against_final_policy_snapshot() {
18914 let active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
18915 let max_active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
18916 let entered = Arc::new(tokio::sync::Barrier::new(2));
18917 let release = Arc::new(tokio::sync::Notify::new());
18918 let handler = Arc::new(BlockingApprovalHandler {
18919 entered: Arc::clone(&entered),
18920 release: Arc::clone(&release),
18921 result: ApprovalResult::Approved,
18922 });
18923 let agent = Arc::new(
18924 AgentBuilder::new()
18925 .system_prompt("Test final policy validation.")
18926 .llm(Arc::new(mock_with_response("done")))
18927 .tool(Arc::new(LockedWriteTool {
18928 active: Arc::clone(&active),
18929 max_active: Arc::clone(&max_active),
18930 }))
18931 .tool_security(ToolSecurityEngine::new(approval_security_config(true)))
18932 .hitl_engine(HITLEngine::new(ai_agents_hitl::HITLConfig::default()))
18933 .approval_handler(handler)
18934 .build()
18935 .unwrap(),
18936 );
18937 let control = agent.runtime_control();
18938 let running = Arc::clone(&agent);
18939 let call = tokio::spawn(async move {
18940 running
18941 .invoke_tool(ToolExecutionRequest::new(
18942 "approval-policy",
18943 "locked_write",
18944 serde_json::json!({"path": "./policy.txt"}),
18945 ToolCallSource::Manual,
18946 ))
18947 .await
18948 .unwrap()
18949 });
18950 entered.wait().await;
18951 let expected_version = control.set_tool_security(approval_security_config(false));
18952 release.notify_one();
18953 let record = call.await.unwrap();
18954
18955 assert!(!record.executed);
18956 assert!(!record.success);
18957 assert_eq!(record.runtime_config_version, expected_version);
18958 assert_eq!(record.policy.outcome, PermissionOutcome::Unavailable);
18959 assert_eq!(max_active.load(Ordering::SeqCst), 0);
18960 assert!(record.metadata.contains_key("policy_snapshot"));
18961 }
18962
18963 #[test]
18964 fn invalid_live_policy_does_not_replace_snapshot_or_generation() {
18965 let agent = AgentBuilder::new()
18966 .system_prompt("Test runtime policy validation.")
18967 .llm(Arc::new(mock_with_response("done")))
18968 .build()
18969 .unwrap();
18970 let control = agent.runtime_control();
18971 let mut valid = ToolSecurityConfig::default();
18972 valid.tools.insert(
18973 "web_search".to_string(),
18974 ai_agents_tools::ToolPolicyConfig {
18975 max_results: Some(5),
18976 ..Default::default()
18977 },
18978 );
18979 let generation = control.try_set_tool_security(valid).unwrap();
18980
18981 let mut invalid = ToolSecurityConfig::default();
18982 invalid.tools.insert(
18983 "web_search".to_string(),
18984 ai_agents_tools::ToolPolicyConfig {
18985 max_results: Some(0),
18986 ..Default::default()
18987 },
18988 );
18989 let error = control.try_set_tool_security(invalid).unwrap_err();
18990
18991 assert!(
18992 error
18993 .to_string()
18994 .contains("max_results must be greater than 0")
18995 );
18996 assert_eq!(control.version(), generation);
18997 assert_eq!(
18998 control
18999 .state
19000 .tool_security_override
19001 .read()
19002 .as_ref()
19003 .unwrap()
19004 .config()
19005 .tools["web_search"]
19006 .max_results,
19007 Some(5)
19008 );
19009 }
19010
19011 #[test]
19013 fn invalid_timeout_config_stops_before_approval_or_tool_invocation() {
19014 let tool_calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19015 let approval_calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19016 let spec = crate::spec::AgentSpec {
19017 tool_security: ToolSecurityConfig {
19018 enabled: true,
19019 default_timeout_ms: u64::MAX,
19020 ..Default::default()
19021 },
19022 ..Default::default()
19023 };
19024
19025 let result = AgentBuilder::from_spec(spec)
19026 .llm(Arc::new(mock_with_response("done")))
19027 .tool(Arc::new(FlakyWriteTool {
19028 calls: Arc::clone(&tool_calls),
19029 }))
19030 .hitl_engine(HITLEngine::new(ai_agents_hitl::HITLConfig::default()))
19031 .approval_handler(Arc::new(CountingApprovalHandler {
19032 calls: Arc::clone(&approval_calls),
19033 }))
19034 .build();
19035
19036 assert!(result.is_err());
19037 assert_eq!(approval_calls.load(Ordering::SeqCst), 0);
19038 assert_eq!(tool_calls.load(Ordering::SeqCst), 0);
19039 }
19040
19041 #[test]
19043 fn invalid_recovery_timeout_config_stops_before_approval_or_tool_invocation() {
19044 use ai_agents_recovery::{ErrorRecoveryConfig, ToolRecoveryConfig, ToolRetryConfig};
19045
19046 let tool_calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19047 let approval_calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19048 let spec = crate::spec::AgentSpec {
19049 error_recovery: ErrorRecoveryConfig {
19050 tools: ToolRecoveryConfig {
19051 default: ToolRetryConfig {
19052 timeout_ms: Some(u64::MAX),
19053 ..Default::default()
19054 },
19055 ..Default::default()
19056 },
19057 ..Default::default()
19058 },
19059 ..Default::default()
19060 };
19061
19062 let result = AgentBuilder::from_spec(spec)
19063 .llm(Arc::new(mock_with_response("done")))
19064 .tool(Arc::new(FlakyWriteTool {
19065 calls: Arc::clone(&tool_calls),
19066 }))
19067 .hitl_engine(HITLEngine::new(ai_agents_hitl::HITLConfig::default()))
19068 .approval_handler(Arc::new(CountingApprovalHandler {
19069 calls: Arc::clone(&approval_calls),
19070 }))
19071 .build();
19072
19073 assert!(result.is_err());
19074 assert_eq!(approval_calls.load(Ordering::SeqCst), 0);
19075 assert_eq!(tool_calls.load(Ordering::SeqCst), 0);
19076 }
19077
19078 #[test]
19080 fn invalid_timeout_policy_does_not_replace_snapshot_or_generation() {
19081 let agent = AgentBuilder::new()
19082 .system_prompt("Test runtime timeout policy validation.")
19083 .llm(Arc::new(mock_with_response("done")))
19084 .build()
19085 .unwrap();
19086 let control = agent.runtime_control();
19087 let valid = ToolSecurityConfig {
19088 default_timeout_ms: 5_000,
19089 ..Default::default()
19090 };
19091 let generation = control.try_set_tool_security(valid).unwrap();
19092
19093 let invalid = ToolSecurityConfig {
19094 default_timeout_ms: MAX_TOOL_TIMEOUT_MS + 1,
19095 ..Default::default()
19096 };
19097 let error = control.try_set_tool_security(invalid).unwrap_err();
19098
19099 assert!(error.to_string().contains(&format!(
19100 "tool_security.default_timeout_ms must be no greater than {MAX_TOOL_TIMEOUT_MS} milliseconds"
19101 )));
19102 assert_eq!(control.version(), generation);
19103 assert_eq!(
19104 control
19105 .state
19106 .tool_security_override
19107 .read()
19108 .as_ref()
19109 .unwrap()
19110 .config()
19111 .default_timeout_ms,
19112 5_000
19113 );
19114 }
19115
19116 #[test]
19118 fn runtime_tool_timeout_conversion_enforces_the_stable_boundary() {
19119 let timeout = RuntimeAgent::validated_tool_timeout(MAX_TOOL_TIMEOUT_MS).unwrap();
19120 assert_eq!(timeout.timer, Duration::from_millis(MAX_TOOL_TIMEOUT_MS));
19121 assert_eq!(
19122 timeout.deadline_delta,
19123 chrono::Duration::milliseconds(MAX_TOOL_TIMEOUT_MS as i64)
19124 );
19125
19126 for timeout_ms in [MAX_TOOL_TIMEOUT_MS + 1, u64::MAX] {
19127 let error = RuntimeAgent::validated_tool_timeout(timeout_ms).unwrap_err();
19128 assert!(error.to_string().contains(&format!(
19129 "effective tool timeout_ms must be no greater than {MAX_TOOL_TIMEOUT_MS} milliseconds"
19130 )));
19131 }
19132 }
19133
19134 #[tokio::test]
19135 async fn persistent_override_preserves_rate_history_within_generation() {
19136 let calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19137 let agent = AgentBuilder::new()
19138 .system_prompt("Test persistent policy overrides.")
19139 .llm(Arc::new(mock_with_response("done")))
19140 .tool(Arc::new(RecoveryTestTool {
19141 id: "limited_override".to_string(),
19142 succeeds: true,
19143 calls: Arc::clone(&calls),
19144 max_output_chars: None,
19145 }))
19146 .build()
19147 .unwrap();
19148 let mut security = ToolSecurityConfig {
19149 enabled: true,
19150 fail_closed: true,
19151 ..Default::default()
19152 };
19153 let policy = ai_agents_tools::ToolPolicyConfig {
19154 write_paths: vec![".".to_string()],
19155 rate_limit: Some(1),
19156 ..Default::default()
19157 };
19158 security
19159 .tools
19160 .insert("limited_override".to_string(), policy);
19161 let generation = agent.runtime_control().set_tool_security(security);
19162
19163 let first = agent
19164 .invoke_tool(ToolExecutionRequest::new(
19165 "limited-first",
19166 "limited_override",
19167 serde_json::json!({"path": "./limited.txt"}),
19168 ToolCallSource::Manual,
19169 ))
19170 .await
19171 .unwrap();
19172 let second = agent
19173 .invoke_tool(ToolExecutionRequest::new(
19174 "limited-second",
19175 "limited_override",
19176 serde_json::json!({"path": "./limited.txt"}),
19177 ToolCallSource::Manual,
19178 ))
19179 .await
19180 .unwrap();
19181
19182 assert!(first.success);
19183 assert_eq!(first.policy_version, generation);
19184 assert!(!second.executed);
19185 assert!(second.output.contains("Rate limit exceeded"));
19186 assert_eq!(second.policy_version, generation);
19187 assert_eq!(calls.load(Ordering::SeqCst), 1);
19188 }
19189
19190 #[tokio::test]
19191 async fn concurrent_rate_admission_consumes_capacity_atomically() {
19192 let calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19193 let tool = Arc::new(RecoveryTestTool {
19194 id: "atomic_rate".to_string(),
19195 succeeds: true,
19196 calls: Arc::clone(&calls),
19197 max_output_chars: None,
19198 });
19199 let arguments = serde_json::json!({"path": "./atomic-rate.txt"});
19200 let bindings = tool.policy_bindings();
19201 let classification = tool.classify_call(&arguments);
19202 let resource_keys =
19203 tool_resource_lock_keys(tool.id(), &arguments, &bindings, &classification);
19204 let mut security = ToolSecurityConfig {
19205 enabled: true,
19206 fail_closed: true,
19207 ..Default::default()
19208 };
19209 let policy = ai_agents_tools::ToolPolicyConfig {
19210 write_paths: vec![".".to_string()],
19211 rate_limit: Some(1),
19212 ..Default::default()
19213 };
19214 security.tools.insert(tool.id().to_string(), policy);
19215 let agent = Arc::new(
19216 AgentBuilder::new()
19217 .system_prompt("Test atomic rate admission.")
19218 .llm(Arc::new(mock_with_response("done")))
19219 .tool(tool)
19220 .tool_security(ToolSecurityEngine::new(security))
19221 .build()
19222 .unwrap(),
19223 );
19224 let held = agent
19225 .acquire_tool_resource_locks(&resource_keys)
19226 .await
19227 .unwrap();
19228 let left = {
19229 let agent = Arc::clone(&agent);
19230 let arguments = arguments.clone();
19231 tokio::spawn(async move {
19232 agent
19233 .invoke_tool(ToolExecutionRequest::new(
19234 "atomic-rate-left",
19235 "atomic_rate",
19236 arguments,
19237 ToolCallSource::Manual,
19238 ))
19239 .await
19240 .unwrap()
19241 })
19242 };
19243 let right = {
19244 let agent = Arc::clone(&agent);
19245 tokio::spawn(async move {
19246 agent
19247 .invoke_tool(ToolExecutionRequest::new(
19248 "atomic-rate-right",
19249 "atomic_rate",
19250 arguments,
19251 ToolCallSource::Manual,
19252 ))
19253 .await
19254 .unwrap()
19255 })
19256 };
19257 tokio::time::sleep(std::time::Duration::from_millis(25)).await;
19258 drop(held);
19259 let (left, right) = tokio::join!(left, right);
19260 let records = [left.unwrap(), right.unwrap()];
19261
19262 assert_eq!(records.iter().filter(|record| record.success).count(), 1);
19263 assert_eq!(records.iter().filter(|record| record.executed).count(), 1);
19264 assert!(
19265 records.iter().any(|record| {
19266 !record.executed && record.output.contains("Rate limit exceeded")
19267 })
19268 );
19269 assert_eq!(calls.load(Ordering::SeqCst), 1);
19270 }
19271
19272 #[tokio::test]
19273 async fn changed_policy_generation_invalidates_pending_approval() {
19274 let active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19275 let max_active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19276 let entered = Arc::new(tokio::sync::Barrier::new(2));
19277 let release = Arc::new(tokio::sync::Notify::new());
19278 let handler = Arc::new(BlockingApprovalHandler {
19279 entered: Arc::clone(&entered),
19280 release: Arc::clone(&release),
19281 result: ApprovalResult::Approved,
19282 });
19283 let agent = Arc::new(
19284 AgentBuilder::new()
19285 .system_prompt("Test stale approval denial.")
19286 .llm(Arc::new(mock_with_response("done")))
19287 .tool(Arc::new(LockedWriteTool {
19288 active: Arc::clone(&active),
19289 max_active: Arc::clone(&max_active),
19290 }))
19291 .tool_security(ToolSecurityEngine::new(approval_security_config(true)))
19292 .hitl_engine(HITLEngine::new(ai_agents_hitl::HITLConfig::default()))
19293 .approval_handler(handler)
19294 .build()
19295 .unwrap(),
19296 );
19297 let running = Arc::clone(&agent);
19298 let call = tokio::spawn(async move {
19299 running
19300 .invoke_tool(ToolExecutionRequest::new(
19301 "stale-approval",
19302 "locked_write",
19303 serde_json::json!({"path": "./stale.txt"}),
19304 ToolCallSource::Manual,
19305 ))
19306 .await
19307 .unwrap()
19308 });
19309 entered.wait().await;
19310 let generation = agent
19311 .runtime_control()
19312 .set_tool_security(approval_security_config(true));
19313 release.notify_one();
19314 let record = call.await.unwrap();
19315
19316 assert!(!record.executed);
19317 assert!(record.output.contains("Approval became stale"));
19318 assert_eq!(record.policy_version, generation);
19319 assert_eq!(max_active.load(Ordering::SeqCst), 0);
19320 }
19321
19322 #[tokio::test]
19323 async fn final_policy_reapplies_argument_caps_after_approval_changes() {
19324 use ai_agents_hitl::CallbackHandler;
19325
19326 let mut security = ToolSecurityConfig {
19327 enabled: true,
19328 fail_closed: true,
19329 ..Default::default()
19330 };
19331 let policy = ai_agents_tools::ToolPolicyConfig {
19332 read_paths: vec![".".to_string()],
19333 max_results: Some(5),
19334 require_confirmation: true,
19335 ..Default::default()
19336 };
19337 security.tools.insert("context_echo".to_string(), policy);
19338 let handler = CallbackHandler::new(|_| ApprovalResult::Modified {
19339 changes: HashMap::from([("max_results".to_string(), serde_json::json!(99))]),
19340 });
19341 let agent = AgentBuilder::new()
19342 .system_prompt("Test final argument caps.")
19343 .llm(Arc::new(mock_with_response("done")))
19344 .tool(Arc::new(ContextEchoTool))
19345 .tool_security(ToolSecurityEngine::new(security))
19346 .hitl_engine(HITLEngine::new(ai_agents_hitl::HITLConfig::default()))
19347 .approval_handler(Arc::new(handler))
19348 .build()
19349 .unwrap();
19350
19351 let record = agent
19352 .invoke_tool(ToolExecutionRequest::new(
19353 "final-cap",
19354 "context_echo",
19355 serde_json::json!({"path": ".", "max_results": 1}),
19356 ToolCallSource::Manual,
19357 ))
19358 .await
19359 .unwrap();
19360
19361 assert!(record.success);
19362 assert_eq!(record.executed_arguments["max_results"], 5);
19363 assert_eq!(
19364 record.approval.unwrap().modified_arguments.unwrap()["max_results"],
19365 5
19366 );
19367 }
19368
19369 #[tokio::test]
19370 async fn no_binding_writes_use_canonical_fallback_lock() {
19371 let active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19372 let max_active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19373 let agent = Arc::new(
19374 AgentBuilder::new()
19375 .system_prompt("Test fallback resource locks.")
19376 .llm(Arc::new(mock_with_response("done")))
19377 .tool(Arc::new(NoBindingWriteTool {
19378 active: Arc::clone(&active),
19379 max_active: Arc::clone(&max_active),
19380 }))
19381 .build()
19382 .unwrap(),
19383 );
19384 let left = {
19385 let agent = Arc::clone(&agent);
19386 tokio::spawn(async move {
19387 agent
19388 .invoke_tool(ToolExecutionRequest::new(
19389 "no-binding-left",
19390 "no_binding_write",
19391 serde_json::json!({}),
19392 ToolCallSource::Manual,
19393 ))
19394 .await
19395 .unwrap()
19396 })
19397 };
19398 let right = {
19399 let agent = Arc::clone(&agent);
19400 tokio::spawn(async move {
19401 agent
19402 .invoke_tool(ToolExecutionRequest::new(
19403 "no-binding-right",
19404 "no_binding_write",
19405 serde_json::json!({}),
19406 ToolCallSource::Manual,
19407 ))
19408 .await
19409 .unwrap()
19410 })
19411 };
19412 let (left, right) = tokio::join!(left, right);
19413
19414 assert!(left.unwrap().success);
19415 assert!(right.unwrap().success);
19416 assert_eq!(max_active.load(Ordering::SeqCst), 1);
19417 }
19418
19419 #[tokio::test]
19420 async fn parent_and_child_paths_share_a_resource_lock() {
19421 let active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19422 let max_active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19423 let agent = Arc::new(
19424 AgentBuilder::new()
19425 .system_prompt("Test parent child resource locks.")
19426 .llm(Arc::new(mock_with_response("done")))
19427 .tool(Arc::new(LockedWriteTool {
19428 active: Arc::clone(&active),
19429 max_active: Arc::clone(&max_active),
19430 }))
19431 .build()
19432 .unwrap(),
19433 );
19434 let parent = format!("./lock-parent-{}", uuid::Uuid::new_v4());
19435 let child = format!("{}/child.txt", parent);
19436 let left = {
19437 let agent = Arc::clone(&agent);
19438 tokio::spawn(async move {
19439 agent
19440 .invoke_tool(ToolExecutionRequest::new(
19441 "parent-lock",
19442 "locked_write",
19443 serde_json::json!({"path": parent}),
19444 ToolCallSource::Manual,
19445 ))
19446 .await
19447 .unwrap()
19448 })
19449 };
19450 let right = {
19451 let agent = Arc::clone(&agent);
19452 tokio::spawn(async move {
19453 agent
19454 .invoke_tool(ToolExecutionRequest::new(
19455 "child-lock",
19456 "locked_write",
19457 serde_json::json!({"path": child}),
19458 ToolCallSource::Manual,
19459 ))
19460 .await
19461 .unwrap()
19462 })
19463 };
19464 let (left, right) = tokio::join!(left, right);
19465
19466 assert!(left.unwrap().success);
19467 assert!(right.unwrap().success);
19468 assert_eq!(max_active.load(Ordering::SeqCst), 1);
19469 }
19470
19471 #[tokio::test]
19472 async fn tool_hooks_can_reenter_after_resource_guards_are_dropped() {
19473 let calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19474 let hooks = Arc::new(ReentrantToolHooks {
19475 agent: parking_lot::Mutex::new(None),
19476 invoked: AtomicBool::new(false),
19477 nested_success: AtomicBool::new(false),
19478 });
19479 let agent = Arc::new(
19480 AgentBuilder::new()
19481 .system_prompt("Test hook reentrancy.")
19482 .llm(Arc::new(mock_with_response("done")))
19483 .tool(Arc::new(RecoveryTestTool {
19484 id: "reentrant_write".to_string(),
19485 succeeds: true,
19486 calls: Arc::clone(&calls),
19487 max_output_chars: None,
19488 }))
19489 .hooks(hooks.clone())
19490 .build()
19491 .unwrap(),
19492 );
19493 *hooks.agent.lock() = Some(Arc::downgrade(&agent));
19494 let record = tokio::time::timeout(
19495 std::time::Duration::from_secs(2),
19496 agent.invoke_tool(ToolExecutionRequest::new(
19497 "outer-hook-call",
19498 "reentrant_write",
19499 serde_json::json!({"path": "./hook.txt"}),
19500 ToolCallSource::Manual,
19501 )),
19502 )
19503 .await
19504 .expect("tool completion hook must not retain resource guards")
19505 .unwrap();
19506
19507 assert!(record.success);
19508 assert!(hooks.nested_success.load(Ordering::SeqCst));
19509 assert_eq!(calls.load(Ordering::SeqCst), 2);
19510 }
19511
19512 #[tokio::test]
19514 async fn fallback_finalizes_original_record_before_shared_execution() {
19515 let primary_calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19516 let fallback_calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19517 let hooks = Arc::new(ToolLifecycleRecordingHooks::new());
19518 let agent = AgentBuilder::new()
19519 .system_prompt("Test fallback execution.")
19520 .llm(Arc::new(mock_with_response("done")))
19521 .tool(Arc::new(RecoveryTestTool {
19522 id: "primary".to_string(),
19523 succeeds: false,
19524 calls: Arc::clone(&primary_calls),
19525 max_output_chars: None,
19526 }))
19527 .tool(Arc::new(RecoveryTestTool {
19528 id: "fallback".to_string(),
19529 succeeds: true,
19530 calls: Arc::clone(&fallback_calls),
19531 max_output_chars: None,
19532 }))
19533 .recovery_manager(recovery_manager_with_fallbacks([(
19534 "primary".to_string(),
19535 "fallback".to_string(),
19536 )]))
19537 .hooks(hooks.clone())
19538 .build()
19539 .unwrap();
19540 let record = tokio::time::timeout(
19541 std::time::Duration::from_secs(2),
19542 agent.invoke_tool(ToolExecutionRequest::new(
19543 "fallback-call",
19544 "primary",
19545 serde_json::json!({"path": "./shared.txt"}),
19546 ToolCallSource::Manual,
19547 )),
19548 )
19549 .await
19550 .expect("fallback must not retain the primary resource guard")
19551 .unwrap();
19552
19553 assert_eq!(
19554 hooks.events(),
19555 vec![
19556 "start:primary",
19557 "complete:primary:false",
19558 "record:primary:true",
19559 "error",
19560 "start:fallback",
19561 "complete:fallback:true",
19562 "record:fallback:true",
19563 ]
19564 );
19565 let records = hooks.records();
19566 assert_eq!(records.len(), 2);
19567 let original = &records[0];
19568 assert_eq!(original.canonical_id, "primary");
19569 assert!(matches!(original.source, ToolCallSource::Manual));
19570 assert!(original.executed);
19571 assert!(!original.success);
19572
19573 let fallback = &records[1];
19574 assert_eq!(fallback.canonical_id, "fallback");
19575 assert_eq!(fallback.call_id, "fallback-call");
19576 assert!(matches!(
19577 &fallback.source,
19578 ToolCallSource::Fallback { original_tool } if original_tool == "primary"
19579 ));
19580 assert!(fallback.executed);
19581 assert!(fallback.success);
19582 assert_eq!(record.canonical_id, fallback.canonical_id);
19583 assert_eq!(record.output, fallback.output);
19584
19585 let history = agent.tool_call_history();
19586 assert_eq!(
19587 history
19588 .iter()
19589 .map(|entry| entry.tool_id.as_str())
19590 .collect::<Vec<_>>(),
19591 vec!["primary", "fallback"]
19592 );
19593 assert_eq!(history[0].result.get("success"), Some(&Value::Bool(false)));
19594 assert_eq!(primary_calls.load(Ordering::SeqCst), 1);
19595 assert_eq!(fallback_calls.load(Ordering::SeqCst), 1);
19596 }
19597
19598 #[tokio::test]
19600 async fn self_fallback_cycle_is_denied_before_reinvocation() {
19601 let calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19602 let hooks = Arc::new(ToolLifecycleRecordingHooks::new());
19603 let agent = AgentBuilder::new()
19604 .system_prompt("Test self-fallback cycle admission.")
19605 .llm(Arc::new(mock_with_response("done")))
19606 .tool(Arc::new(RecoveryTestTool {
19607 id: "primary".to_string(),
19608 succeeds: false,
19609 calls: Arc::clone(&calls),
19610 max_output_chars: None,
19611 }))
19612 .recovery_manager(recovery_manager_with_fallbacks([(
19613 "primary".to_string(),
19614 "primary".to_string(),
19615 )]))
19616 .hooks(hooks.clone())
19617 .build()
19618 .unwrap();
19619
19620 let record = tokio::time::timeout(
19621 std::time::Duration::from_secs(2),
19622 agent.invoke_tool(ToolExecutionRequest::new(
19623 "self-fallback-call",
19624 "primary",
19625 serde_json::json!({"path": "./shared.txt"}),
19626 ToolCallSource::Manual,
19627 )),
19628 )
19629 .await
19630 .expect("self fallback must terminate without recursive execution")
19631 .unwrap();
19632
19633 assert_eq!(calls.load(Ordering::SeqCst), 1);
19634 assert_eq!(record.canonical_id, "primary");
19635 assert!(!record.executed);
19636 assert!(!record.success);
19637 assert_eq!(record.policy.outcome, PermissionOutcome::Deny);
19638 assert!(record.output.contains("fallback cycle"));
19639 assert!(matches!(
19640 record.source,
19641 ToolCallSource::Fallback { ref original_tool } if original_tool == "primary"
19642 ));
19643 assert_eq!(
19644 record.metadata.get("fallback_chain"),
19645 Some(&serde_json::json!(["primary"]))
19646 );
19647 assert_eq!(
19648 hooks.events(),
19649 vec![
19650 "start:primary",
19651 "complete:primary:false",
19652 "record:primary:true",
19653 "error",
19654 "complete:primary:false",
19655 "record:primary:false",
19656 "error",
19657 ]
19658 );
19659 assert_eq!(agent.tool_call_history().len(), 2);
19660 }
19661
19662 #[tokio::test]
19664 async fn alias_mediated_fallback_cycle_is_denied_canonically() {
19665 let primary_calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19666 let secondary_calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19667 let hooks = Arc::new(ToolLifecycleRecordingHooks::new());
19668 let agent = AgentBuilder::new()
19669 .system_prompt("Test canonical fallback cycle admission.")
19670 .llm(Arc::new(mock_with_response("done")))
19671 .tool(Arc::new(RecoveryTestTool {
19672 id: "primary".to_string(),
19673 succeeds: false,
19674 calls: Arc::clone(&primary_calls),
19675 max_output_chars: None,
19676 }))
19677 .tool(Arc::new(RecoveryTestTool {
19678 id: "secondary".to_string(),
19679 succeeds: false,
19680 calls: Arc::clone(&secondary_calls),
19681 max_output_chars: None,
19682 }))
19683 .recovery_manager(recovery_manager_with_fallbacks([
19684 ("primary".to_string(), "secondary".to_string()),
19685 ("secondary".to_string(), "primary alias".to_string()),
19686 ]))
19687 .hooks(hooks.clone())
19688 .build()
19689 .unwrap();
19690 agent.tools.set_tool_aliases(
19691 "primary",
19692 ToolAliases::new().with_name("en", "primary alias"),
19693 );
19694
19695 let record = tokio::time::timeout(
19696 std::time::Duration::from_secs(2),
19697 agent.invoke_tool(ToolExecutionRequest::new(
19698 "alias-fallback-call",
19699 "primary",
19700 serde_json::json!({"path": "./shared.txt"}),
19701 ToolCallSource::Manual,
19702 )),
19703 )
19704 .await
19705 .expect("alias-mediated fallback cycle must terminate")
19706 .unwrap();
19707
19708 assert_eq!(primary_calls.load(Ordering::SeqCst), 1);
19709 assert_eq!(secondary_calls.load(Ordering::SeqCst), 1);
19710 assert_eq!(record.requested_name, "primary alias");
19711 assert_eq!(record.canonical_id, "primary");
19712 assert!(!record.executed);
19713 assert!(record.output.contains("fallback cycle"));
19714 assert_eq!(
19715 record.metadata.get("fallback_chain"),
19716 Some(&serde_json::json!(["primary", "secondary"]))
19717 );
19718 assert_eq!(hooks.records().len(), 3);
19719 assert_eq!(agent.tool_call_history().len(), 3);
19720 }
19721
19722 #[tokio::test]
19724 async fn final_canonical_drift_cannot_bypass_fallback_ancestry() {
19725 let primary_calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19726 let secondary_calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19727 let provider = Arc::new(DriftingFallbackProvider {
19728 refreshed: AtomicBool::new(false),
19729 primary_calls: Arc::clone(&primary_calls),
19730 secondary_calls: Arc::clone(&secondary_calls),
19731 });
19732 let registry = ToolRegistry::new();
19733 registry.register_provider(provider).await.unwrap();
19734 let lifecycle = Arc::new(ToolLifecycleRecordingHooks::new());
19735 let hooks = Arc::new(RefreshFallbackProviderHooks {
19736 agent: parking_lot::Mutex::new(None),
19737 lifecycle: Arc::clone(&lifecycle),
19738 });
19739 let agent = Arc::new(
19740 AgentBuilder::new()
19741 .system_prompt("Test final canonical fallback admission.")
19742 .llm(Arc::new(mock_with_response("done")))
19743 .tools(registry)
19744 .recovery_manager(recovery_manager_with_fallbacks([(
19745 "primary".to_string(),
19746 "fallback alias".to_string(),
19747 )]))
19748 .hooks(hooks.clone())
19749 .build()
19750 .unwrap(),
19751 );
19752 *hooks.agent.lock() = Some(Arc::downgrade(&agent));
19753
19754 let record = agent
19755 .invoke_tool(ToolExecutionRequest::new(
19756 "drifting-fallback-call",
19757 "primary",
19758 serde_json::json!({"path": "./shared.txt"}),
19759 ToolCallSource::Manual,
19760 ))
19761 .await
19762 .unwrap();
19763
19764 assert_eq!(primary_calls.load(Ordering::SeqCst), 1);
19765 assert_eq!(secondary_calls.load(Ordering::SeqCst), 0);
19766 assert_eq!(record.canonical_id, "secondary");
19767 assert!(!record.executed);
19768 assert_eq!(record.policy.outcome, PermissionOutcome::Deny);
19769 assert!(record.output.contains("fallback cycle"));
19770 assert_eq!(
19771 record.metadata.get("fallback_chain"),
19772 Some(&serde_json::json!(["primary", "secondary"]))
19773 );
19774 assert_eq!(
19775 record.metadata.get("final_resolved_canonical_id"),
19776 Some(&serde_json::json!("primary"))
19777 );
19778 assert_eq!(
19779 lifecycle.events(),
19780 vec![
19781 "start:primary",
19782 "complete:primary:false",
19783 "record:primary:true",
19784 "error",
19785 "start:secondary",
19786 "complete:secondary:false",
19787 "record:secondary:false",
19788 "error",
19789 ]
19790 );
19791 let records = lifecycle.records();
19792 assert_eq!(records.len(), 2);
19793 assert_eq!(records[1].canonical_id, "secondary");
19794 assert_eq!(
19795 records[1].metadata.get("final_resolved_canonical_id"),
19796 Some(&serde_json::json!("primary"))
19797 );
19798 let history = agent.tool_call_history();
19799 assert_eq!(
19800 history
19801 .iter()
19802 .map(|entry| entry.tool_id.as_str())
19803 .collect::<Vec<_>>(),
19804 vec!["primary", "secondary"]
19805 );
19806 }
19807
19808 #[tokio::test]
19810 async fn acyclic_fallback_chain_is_denied_after_the_hop_limit() {
19811 let tool_count = MAX_TOOL_FALLBACK_HOPS + 2;
19812 let calls = (0..tool_count)
19813 .map(|_| Arc::new(std::sync::atomic::AtomicUsize::new(0)))
19814 .collect::<Vec<_>>();
19815 let mut builder = AgentBuilder::new()
19816 .system_prompt("Test bounded acyclic fallback admission.")
19817 .llm(Arc::new(mock_with_response("done")));
19818 for (index, counter) in calls.iter().enumerate() {
19819 builder = builder.tool(Arc::new(RecoveryTestTool {
19820 id: format!("fallback_{index}"),
19821 succeeds: false,
19822 calls: Arc::clone(counter),
19823 max_output_chars: None,
19824 }));
19825 }
19826 let fallbacks = (0..tool_count - 1).map(|index| {
19827 (
19828 format!("fallback_{index}"),
19829 format!("fallback_{}", index + 1),
19830 )
19831 });
19832 let agent = builder
19833 .recovery_manager(recovery_manager_with_fallbacks(fallbacks))
19834 .build()
19835 .unwrap();
19836
19837 let record = tokio::time::timeout(
19838 std::time::Duration::from_secs(2),
19839 agent.invoke_tool(ToolExecutionRequest::new(
19840 "bounded-fallback-call",
19841 "fallback_0",
19842 serde_json::json!({"path": "./shared.txt"}),
19843 ToolCallSource::Manual,
19844 )),
19845 )
19846 .await
19847 .expect("bounded fallback chain must terminate")
19848 .unwrap();
19849
19850 for counter in calls.iter().take(MAX_TOOL_FALLBACK_HOPS + 1) {
19851 assert_eq!(counter.load(Ordering::SeqCst), 1);
19852 }
19853 assert_eq!(calls[MAX_TOOL_FALLBACK_HOPS + 1].load(Ordering::SeqCst), 0);
19854 assert_eq!(
19855 record.canonical_id,
19856 format!("fallback_{}", MAX_TOOL_FALLBACK_HOPS + 1)
19857 );
19858 assert!(!record.executed);
19859 assert!(record.output.contains("maximum of 16 hops"));
19860 assert_eq!(agent.tool_call_history().len(), tool_count);
19861 }
19862
19863 #[tokio::test]
19864 async fn diagnostics_without_provider_records_unavailable_without_execution() {
19865 let mock = mock_with_response("hello");
19866 let yaml = r#"
19867name: DiagnosticsNoProviderAgent
19868system_prompt: "Review diagnostics."
19869tools: [diagnostics]
19870"#;
19871 let agent = AgentBuilder::from_yaml(yaml)
19872 .unwrap()
19873 .llm(Arc::new(mock))
19874 .auto_configure_features()
19875 .unwrap()
19876 .build()
19877 .unwrap();
19878
19879 let record = agent
19880 .invoke_tool(ToolExecutionRequest::new(
19881 "diagnostics-call",
19882 "diagnostics",
19883 serde_json::json!({}),
19884 ToolCallSource::Manual,
19885 ))
19886 .await
19887 .unwrap();
19888
19889 assert!(!record.executed);
19890 assert!(!record.success);
19891 assert_eq!(record.policy.outcome, PermissionOutcome::Unavailable);
19892 }
19893
19894 #[tokio::test]
19895 async fn web_search_without_provider_records_unavailable_without_execution() {
19896 let mock = mock_with_response("hello");
19897 let yaml = r#"
19898name: WebSearchNoProviderAgent
19899system_prompt: "You search the web."
19900tools: [web_search]
19901"#;
19902 let agent = AgentBuilder::from_yaml(yaml)
19903 .unwrap()
19904 .llm(Arc::new(mock))
19905 .auto_configure_features()
19906 .unwrap()
19907 .build()
19908 .unwrap();
19909
19910 let record = agent
19911 .invoke_tool(ToolExecutionRequest::new(
19912 "web-search-call",
19913 "web_search",
19914 serde_json::json!({"query": "rust async"}),
19915 ToolCallSource::Manual,
19916 ))
19917 .await
19918 .unwrap();
19919
19920 assert!(!record.executed);
19921 assert!(!record.success);
19922 assert_eq!(record.policy.outcome, PermissionOutcome::Unavailable);
19923 }
19924
19925 #[tokio::test]
19926 async fn unavailable_host_tool_does_not_request_approval() {
19927 let approvals = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19928 let handler = Arc::new(CountingApprovalHandler {
19929 calls: Arc::clone(&approvals),
19930 });
19931 let mut security = ToolSecurityConfig {
19932 enabled: true,
19933 fail_closed: true,
19934 ..Default::default()
19935 };
19936 security.tools.insert(
19937 "web_search".to_string(),
19938 ai_agents_tools::ToolPolicyConfig {
19939 enabled: true,
19940 require_confirmation: true,
19941 ..Default::default()
19942 },
19943 );
19944 let yaml = r#"
19945name: UnavailableApprovalAgent
19946system_prompt: "Search only with approval."
19947tools: [web_search]
19948"#;
19949 let agent = AgentBuilder::from_yaml(yaml)
19950 .unwrap()
19951 .llm(Arc::new(mock_with_response("done")))
19952 .auto_configure_features()
19953 .unwrap()
19954 .tool_security(ToolSecurityEngine::new(security))
19955 .hitl_engine(HITLEngine::new(ai_agents_hitl::HITLConfig::default()))
19956 .approval_handler(handler)
19957 .build()
19958 .unwrap();
19959
19960 let record = agent
19961 .invoke_tool(ToolExecutionRequest::new(
19962 "unavailable-before-approval",
19963 "web_search",
19964 serde_json::json!({"query": "rust async"}),
19965 ToolCallSource::Manual,
19966 ))
19967 .await
19968 .unwrap();
19969
19970 assert_eq!(approvals.load(Ordering::SeqCst), 0);
19971 assert!(!record.executed);
19972 assert!(!record.success);
19973 assert_eq!(record.policy.outcome, PermissionOutcome::Unavailable);
19974 assert!(
19975 record
19976 .approval
19977 .as_ref()
19978 .is_some_and(|approval| matches!(approval.status, ToolApprovalStatus::Unavailable))
19979 );
19980 }
19981
19982 #[tokio::test]
19983 async fn test_spawner_section_does_not_grant_core_tools_when_top_level_tools_omitted() {
19984 let mock = mock_with_response("hello");
19985 let yaml = r#"
19986name: SpawnerNoGrantAgent
19987system_prompt: "You manage agents."
19988spawner:
19989 max_agents: 2
19990"#;
19991 let agent = AgentBuilder::from_yaml(yaml)
19992 .unwrap()
19993 .llm(Arc::new(mock))
19994 .auto_configure_features()
19995 .unwrap()
19996 .auto_configure_spawner()
19997 .await
19998 .unwrap()
19999 .build()
20000 .unwrap();
20001
20002 let available = agent.get_available_tool_ids().await.unwrap();
20003 assert!(available.is_empty());
20004 }
20005
20006 #[tokio::test]
20007 async fn test_spawner_section_does_not_grant_core_tools_when_top_level_tools_empty() {
20008 let mock = mock_with_response("hello");
20009 let yaml = r#"
20010name: EmptySpawnerNoGrantAgent
20011system_prompt: "You manage agents."
20012tools: []
20013spawner:
20014 max_agents: 2
20015"#;
20016 let agent = AgentBuilder::from_yaml(yaml)
20017 .unwrap()
20018 .llm(Arc::new(mock))
20019 .auto_configure_features()
20020 .unwrap()
20021 .auto_configure_spawner()
20022 .await
20023 .unwrap()
20024 .build()
20025 .unwrap();
20026
20027 let available = agent.get_available_tool_ids().await.unwrap();
20028 assert!(available.is_empty());
20029 }
20030
20031 #[tokio::test]
20032 async fn test_management_tools_flag_grants_core_tools_when_top_level_tools_empty() {
20033 let mock = mock_with_response("hello");
20034 let yaml = r#"
20035name: ManagementGrantAgent
20036system_prompt: "You manage agents."
20037tools: []
20038spawner:
20039 management_tools: true
20040"#;
20041 let agent = AgentBuilder::from_yaml(yaml)
20042 .unwrap()
20043 .llm(Arc::new(mock))
20044 .auto_configure_features()
20045 .unwrap()
20046 .auto_configure_spawner()
20047 .await
20048 .unwrap()
20049 .build()
20050 .unwrap();
20051
20052 let available = agent.get_available_tool_ids().await.unwrap();
20053 assert_eq!(available.len(), 4);
20054 assert!(available.contains(&"spawn_agent".to_string()));
20055 assert!(available.contains(&"send_agent_message".to_string()));
20056 assert!(available.contains(&"list_agents".to_string()));
20057 assert!(available.contains(&"remove_agent".to_string()));
20058 }
20059
20060 #[tokio::test]
20061 async fn test_management_tools_flag_grants_core_tools_when_top_level_tools_omitted() {
20062 let mock = mock_with_response("hello");
20063 let yaml = r#"
20064name: ManagementOmittedToolsGrantAgent
20065system_prompt: "You manage agents."
20066spawner:
20067 management_tools: true
20068"#;
20069 let agent = AgentBuilder::from_yaml(yaml)
20070 .unwrap()
20071 .llm(Arc::new(mock))
20072 .auto_configure_features()
20073 .unwrap()
20074 .auto_configure_spawner()
20075 .await
20076 .unwrap()
20077 .build()
20078 .unwrap();
20079
20080 let available = agent.get_available_tool_ids().await.unwrap();
20081 assert_eq!(available.len(), 4);
20082 assert!(available.contains(&"spawn_agent".to_string()));
20083 assert!(available.contains(&"send_agent_message".to_string()));
20084 assert!(available.contains(&"list_agents".to_string()));
20085 assert!(available.contains(&"remove_agent".to_string()));
20086 }
20087
20088 #[tokio::test]
20089 async fn test_management_tools_selected_grants_only_selected_tools() {
20090 let mock = mock_with_response("hello");
20091 let yaml = r#"
20092name: ManagementSelectedGrantAgent
20093system_prompt: "You manage agents."
20094tools: []
20095spawner:
20096 management_tools:
20097 - spawn_agent
20098 - send_agent_message
20099 - list_agents
20100"#;
20101 let agent = AgentBuilder::from_yaml(yaml)
20102 .unwrap()
20103 .llm(Arc::new(mock))
20104 .auto_configure_features()
20105 .unwrap()
20106 .auto_configure_spawner()
20107 .await
20108 .unwrap()
20109 .build()
20110 .unwrap();
20111
20112 let available = agent.get_available_tool_ids().await.unwrap();
20113 assert_eq!(available.len(), 3);
20114 assert!(available.contains(&"spawn_agent".to_string()));
20115 assert!(available.contains(&"send_agent_message".to_string()));
20116 assert!(available.contains(&"list_agents".to_string()));
20117 assert!(!available.contains(&"remove_agent".to_string()));
20118 }
20119
20120 #[tokio::test]
20121 async fn test_orchestration_tools_flag_grants_tools_when_top_level_tools_empty() {
20122 let mock = mock_with_response("hello");
20123 let yaml = r#"
20124name: OrchestrationGrantAgent
20125system_prompt: "You coordinate agents."
20126llms:
20127 default:
20128 provider: openai
20129 model: gpt-4
20130 router:
20131 provider: openai
20132 model: gpt-4
20133llm:
20134 default: default
20135 router: router
20136tools: []
20137spawner:
20138 orchestration_tools: true
20139"#;
20140 let agent = AgentBuilder::from_yaml(yaml)
20141 .unwrap()
20142 .llm(Arc::new(mock))
20143 .auto_configure_features()
20144 .unwrap()
20145 .auto_configure_spawner()
20146 .await
20147 .unwrap()
20148 .build()
20149 .unwrap();
20150
20151 let available = agent.get_available_tool_ids().await.unwrap();
20152 assert_eq!(available.len(), 5);
20153 assert!(available.contains(&"route_to_agent".to_string()));
20154 assert!(available.contains(&"pipeline_process".to_string()));
20155 assert!(available.contains(&"concurrent_ask".to_string()));
20156 assert!(available.contains(&"group_discussion".to_string()));
20157 assert!(available.contains(&"handoff_conversation".to_string()));
20158 }
20159
20160 #[tokio::test]
20161 async fn test_persona_evolve_flag_grants_tool_when_top_level_tools_empty() {
20162 let mock = mock_with_response("hello");
20163 let yaml = r#"
20164name: PersonaGrantAgent
20165system_prompt: "You can evolve persona."
20166llm:
20167 provider: openai
20168 model: gpt-4
20169tools: []
20170persona:
20171 identity:
20172 name: "Guide"
20173 role: "Helper"
20174 evolution:
20175 enabled: true
20176 allow_llm_evolve: true
20177 mutable_fields:
20178 - traits.personality
20179"#;
20180 let agent = AgentBuilder::from_yaml(yaml)
20181 .unwrap()
20182 .llm(Arc::new(mock))
20183 .build()
20184 .unwrap();
20185
20186 let available = agent.get_available_tool_ids().await.unwrap();
20187 assert_eq!(available, vec!["persona_evolve".to_string()]);
20188 }
20189
20190 #[tokio::test]
20191 async fn test_persona_evolve_flag_grants_tool_when_top_level_tools_omitted() {
20192 let mock = mock_with_response("hello");
20193 let yaml = r#"
20194name: PersonaOmittedToolsGrantAgent
20195system_prompt: "You can evolve persona."
20196llm:
20197 provider: openai
20198 model: gpt-4
20199persona:
20200 identity:
20201 name: "Guide"
20202 role: "Helper"
20203 evolution:
20204 enabled: true
20205 allow_llm_evolve: true
20206 mutable_fields:
20207 - traits.personality
20208"#;
20209 let agent = AgentBuilder::from_yaml(yaml)
20210 .unwrap()
20211 .llm(Arc::new(mock))
20212 .build()
20213 .unwrap();
20214
20215 let available = agent.get_available_tool_ids().await.unwrap();
20216 assert_eq!(available, vec!["persona_evolve".to_string()]);
20217 }
20218
20219 #[tokio::test]
20220 async fn test_omitted_yaml_tools_exposes_no_tools() {
20221 let mock = mock_with_response("hello");
20222 let yaml = r#"
20223name: NoToolsAgent
20224system_prompt: "You are helpful."
20225"#;
20226 let agent = AgentBuilder::from_yaml(yaml)
20227 .unwrap()
20228 .llm(Arc::new(mock))
20229 .auto_configure_features()
20230 .unwrap()
20231 .build()
20232 .unwrap();
20233
20234 let available = agent.get_available_tool_ids().await.unwrap();
20235 assert!(available.is_empty());
20236 }
20237
20238 #[tokio::test]
20239 async fn runtime_scope_cannot_widen_omitted_or_empty_yaml_grants() {
20240 for tools in ["", "tools: []"] {
20241 let yaml = format!(
20242 r#"
20243name: RuntimeScopeNoGrantAgent
20244system_prompt: "No ordinary tools are granted."
20245{tools}
20246"#
20247 );
20248 let agent = AgentBuilder::from_yaml(&yaml)
20249 .unwrap()
20250 .llm(Arc::new(mock_with_response("done")))
20251 .auto_configure_features()
20252 .unwrap()
20253 .build()
20254 .unwrap();
20255
20256 agent
20257 .runtime_control()
20258 .set_tool_scope(vec!["calculator".to_string()]);
20259
20260 assert!(agent.get_available_tool_ids().await.unwrap().is_empty());
20261 }
20262 }
20263
20264 #[tokio::test]
20265 async fn runtime_scope_widening_attempt_keeps_only_declared_tools() {
20266 let yaml = r#"
20267name: RuntimeScopeWideningAgent
20268system_prompt: "Runtime scope cannot add authority."
20269tools: [calculator]
20270"#;
20271 let agent = AgentBuilder::from_yaml(yaml)
20272 .unwrap()
20273 .llm(Arc::new(mock_with_response("done")))
20274 .auto_configure_features()
20275 .unwrap()
20276 .build()
20277 .unwrap();
20278
20279 agent
20280 .runtime_control()
20281 .set_tool_scope(vec!["calculator".to_string(), "datetime".to_string()]);
20282
20283 assert_eq!(
20284 agent.get_available_tool_ids().await.unwrap(),
20285 vec!["calculator".to_string()]
20286 );
20287 }
20288
20289 #[tokio::test]
20290 async fn runtime_scope_is_canonical_unique_ordered_and_clear_restores_declared_grant() {
20291 let yaml = r#"
20292name: RuntimeScopeIntersectionAgent
20293system_prompt: "Use only declared tools."
20294tools: [calculator, datetime]
20295"#;
20296 let agent = AgentBuilder::from_yaml(yaml)
20297 .unwrap()
20298 .llm(Arc::new(mock_with_response("done")))
20299 .auto_configure_features()
20300 .unwrap()
20301 .build()
20302 .unwrap();
20303 let mut aliases = ai_agents_tools::ToolAliases::default();
20304 aliases
20305 .names
20306 .insert("en".to_string(), "calculate_alias".to_string());
20307 agent.tools.set_tool_aliases("calculator", aliases);
20308 let control = agent.runtime_control();
20309
20310 control.set_tool_scope(vec![
20311 "datetime".to_string(),
20312 "calculate_alias".to_string(),
20313 "calculator".to_string(),
20314 "unknown".to_string(),
20315 "datetime".to_string(),
20316 ]);
20317 assert_eq!(
20318 agent.get_available_tool_ids().await.unwrap(),
20319 vec!["calculator".to_string(), "datetime".to_string()]
20320 );
20321
20322 control.set_tool_scope(vec!["datetime".to_string()]);
20323 assert_eq!(
20324 agent.get_available_tool_ids().await.unwrap(),
20325 vec!["datetime".to_string()]
20326 );
20327
20328 control.clear_tool_scope_override();
20329 assert_eq!(
20330 agent.get_available_tool_ids().await.unwrap(),
20331 vec!["calculator".to_string(), "datetime".to_string()]
20332 );
20333 }
20334
20335 #[tokio::test]
20336 async fn runtime_scope_preserves_programmatic_registration_as_declared_grant() {
20337 let agent = AgentBuilder::new()
20338 .system_prompt("Use registered tools.")
20339 .llm(Arc::new(mock_with_response("done")))
20340 .tool(Arc::new(ContextEchoTool))
20341 .tool(Arc::new(SlowTool))
20342 .build()
20343 .unwrap();
20344
20345 agent.runtime_control().set_tool_scope(vec![
20346 "Context Echo".to_string(),
20347 "context_echo".to_string(),
20348 "unknown".to_string(),
20349 ]);
20350
20351 assert_eq!(
20352 agent.get_available_tool_ids().await.unwrap(),
20353 vec!["context_echo".to_string()]
20354 );
20355 }
20356
20357 #[tokio::test]
20358 async fn nested_state_scopes_intersect_every_ancestor_with_aliases() {
20359 let yaml = r#"
20360name: NestedStateScopeAgent
20361system_prompt: "Honor every state scope."
20362tools: [calculator, datetime, echo]
20363states:
20364 initial: root
20365 states:
20366 root:
20367 tools: [calculate_alias, datetime]
20368 initial: middle
20369 states:
20370 middle:
20371 initial: leaf
20372 states:
20373 leaf:
20374 tools: [datetime_alias, echo]
20375"#;
20376 let agent = AgentBuilder::from_yaml(yaml)
20377 .unwrap()
20378 .llm(Arc::new(mock_with_response("done")))
20379 .auto_configure_features()
20380 .unwrap()
20381 .build()
20382 .unwrap();
20383 let mut calculator_aliases = ai_agents_tools::ToolAliases::default();
20384 calculator_aliases
20385 .names
20386 .insert("en".to_string(), "calculate_alias".to_string());
20387 agent
20388 .tools
20389 .set_tool_aliases("calculator", calculator_aliases);
20390 let mut datetime_aliases = ai_agents_tools::ToolAliases::default();
20391 datetime_aliases
20392 .names
20393 .insert("en".to_string(), "datetime_alias".to_string());
20394 agent.tools.set_tool_aliases("datetime", datetime_aliases);
20395 agent.runtime_control().set_tool_scope(vec![
20396 "unknown".to_string(),
20397 "datetime_alias".to_string(),
20398 "calculate_alias".to_string(),
20399 "datetime".to_string(),
20400 ]);
20401
20402 assert_eq!(agent.current_state().as_deref(), Some("root.middle.leaf"));
20403 assert_eq!(
20404 agent.get_available_tool_ids().await.unwrap(),
20405 vec!["datetime".to_string()]
20406 );
20407 }
20408
20409 #[tokio::test]
20410 async fn ancestor_empty_state_scope_denies_omitted_descendants() {
20411 let yaml = r#"
20412name: NestedEmptyStateScopeAgent
20413system_prompt: "An empty ancestor scope denies all tools."
20414tools: [calculator]
20415states:
20416 initial: root
20417 states:
20418 root:
20419 tools: []
20420 initial: middle
20421 states:
20422 middle:
20423 initial: leaf
20424 states:
20425 leaf: {}
20426"#;
20427 let agent = AgentBuilder::from_yaml(yaml)
20428 .unwrap()
20429 .llm(Arc::new(mock_with_response("done")))
20430 .auto_configure_features()
20431 .unwrap()
20432 .build()
20433 .unwrap();
20434
20435 assert!(agent.get_available_tool_ids().await.unwrap().is_empty());
20436 }
20437
20438 #[tokio::test]
20439 async fn state_change_during_approval_invalidates_the_reviewed_authority() {
20440 let active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
20441 let max_active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
20442 let entered = Arc::new(tokio::sync::Barrier::new(2));
20443 let release = Arc::new(tokio::sync::Notify::new());
20444 let handler = Arc::new(BlockingApprovalHandler {
20445 entered: Arc::clone(&entered),
20446 release: Arc::clone(&release),
20447 result: ApprovalResult::Approved,
20448 });
20449 let yaml = r#"
20450name: ApprovalStateGenerationAgent
20451system_prompt: "State authority may change during approval."
20452tools: [locked_write]
20453states:
20454 initial: first
20455 states:
20456 first:
20457 tools: [locked_write]
20458 second:
20459 tools: [locked_write]
20460"#;
20461 let agent = Arc::new(
20462 AgentBuilder::from_yaml(yaml)
20463 .unwrap()
20464 .llm(Arc::new(mock_with_response("done")))
20465 .tool(Arc::new(LockedWriteTool {
20466 active: Arc::clone(&active),
20467 max_active: Arc::clone(&max_active),
20468 }))
20469 .tool_security(ToolSecurityEngine::new(approval_security_config(true)))
20470 .hitl_engine(HITLEngine::new(ai_agents_hitl::HITLConfig::default()))
20471 .approval_handler(handler)
20472 .build()
20473 .unwrap(),
20474 );
20475 let running = Arc::clone(&agent);
20476 let call = tokio::spawn(async move {
20477 running
20478 .invoke_tool(ToolExecutionRequest::new(
20479 "approval-state-generation",
20480 "locked_write",
20481 serde_json::json!({"path": "./state-generation.txt"}),
20482 ToolCallSource::Manual,
20483 ))
20484 .await
20485 .unwrap()
20486 });
20487
20488 entered.wait().await;
20489 agent.transition_to("second").await.unwrap();
20490 release.notify_one();
20491 let record = call.await.unwrap();
20492
20493 assert!(!record.executed);
20494 assert!(record.output.contains("Approval became stale"));
20495 assert_eq!(max_active.load(Ordering::SeqCst), 0);
20496 }
20497
20498 #[tokio::test]
20499 async fn state_change_while_waiting_for_resource_lock_fails_final_admission() {
20500 let holder_gate = PathMutationGate::new();
20501 let waiter_gate = PathMutationGate::new();
20502 let yaml = r#"
20503name: LockedStateGenerationAgent
20504system_prompt: "State authority must remain stable through admission."
20505tools: [state_lock_holder, state_lock_waiter]
20506states:
20507 initial: first
20508 states:
20509 first:
20510 tools: [state_lock_holder, state_lock_waiter]
20511 second:
20512 tools: [state_lock_holder, state_lock_waiter]
20513"#;
20514 let agent = Arc::new(
20515 AgentBuilder::from_yaml(yaml)
20516 .unwrap()
20517 .llm(Arc::new(mock_with_response("done")))
20518 .tool(Arc::new(BlockingPathMutationTool {
20519 id: "state_lock_holder",
20520 path_fields: vec![ai_agents_core::PathPolicyBinding::write("path")],
20521 gate: holder_gate.clone(),
20522 }))
20523 .tool(Arc::new(BlockingPathMutationTool {
20524 id: "state_lock_waiter",
20525 path_fields: vec![ai_agents_core::PathPolicyBinding::write("path")],
20526 gate: waiter_gate.clone(),
20527 }))
20528 .build()
20529 .unwrap(),
20530 );
20531 let holder_call = {
20532 let agent = Arc::clone(&agent);
20533 tokio::spawn(async move {
20534 agent
20535 .invoke_tool(ToolExecutionRequest::new(
20536 "state-lock-holder",
20537 "state_lock_holder",
20538 serde_json::json!({"path": "./shared-state-path.txt"}),
20539 ToolCallSource::Manual,
20540 ))
20541 .await
20542 .unwrap()
20543 })
20544 };
20545 holder_gate.wait_until_entered().await;
20546 let waiter_call = {
20547 let agent = Arc::clone(&agent);
20548 tokio::spawn(async move {
20549 agent
20550 .invoke_tool(ToolExecutionRequest::new(
20551 "state-lock-waiter",
20552 "state_lock_waiter",
20553 serde_json::json!({"path": "./shared-state-path.txt"}),
20554 ToolCallSource::Manual,
20555 ))
20556 .await
20557 .unwrap()
20558 })
20559 };
20560
20561 wait_for_resource_lock_strong_count(&agent.resource_locks, 2).await;
20562 agent.transition_to("second").await.unwrap();
20563 holder_gate.release();
20564 let holder_record = holder_call.await.unwrap();
20565 let waiter_record = waiter_call.await.unwrap();
20566
20567 assert!(holder_record.success);
20568 assert!(!waiter_record.executed);
20569 assert!(
20570 waiter_record
20571 .output
20572 .contains("state scope changed before admission")
20573 );
20574 assert!(!waiter_gate.entered.load(Ordering::SeqCst));
20575 }
20576
20577 #[tokio::test]
20578 async fn test_state_tools_cannot_widen_top_level_grant() {
20579 let mock = mock_with_response("hello");
20580 let yaml = r#"
20581name: NarrowToolsAgent
20582system_prompt: "You are helpful."
20583tools:
20584 - calculator
20585states:
20586 initial: current
20587 states:
20588 current:
20589 tools: [datetime]
20590"#;
20591 let agent = AgentBuilder::from_yaml(yaml)
20592 .unwrap()
20593 .llm(Arc::new(mock))
20594 .auto_configure_features()
20595 .unwrap()
20596 .build()
20597 .unwrap();
20598
20599 let available = agent.get_available_tool_ids().await.unwrap();
20600 assert!(available.is_empty());
20601 }
20602
20603 #[tokio::test]
20605 async fn test_integration_tool_execution() {
20606 let mock = mock_with_responses(vec![
20608 r#"I'll calculate that for you.
20610{"tool": "calculator", "arguments": {"expression": "2+2"}}"#,
20611 "The answer is 4.",
20613 ]);
20614 let observed = mock.clone();
20615 let mut tools = ai_agents_tools::ToolRegistry::new();
20616 tools
20617 .register(Arc::new(ai_agents_tools::CalculatorTool))
20618 .unwrap();
20619
20620 let agent = AgentBuilder::new()
20621 .system_prompt("You are a calculator assistant.")
20622 .llm(Arc::new(mock))
20623 .tools(tools)
20624 .build()
20625 .unwrap();
20626
20627 let response = agent.chat("What is 2+2?").await.unwrap();
20628
20629 assert_eq!(response.content, "The answer is 4.");
20630 assert_eq!(response.tool_calls.as_ref().map(Vec::len), Some(1));
20631 assert_eq!(
20632 observed.call_count(),
20633 2,
20634 "tool result must trigger a second LLM call"
20635 );
20636 let history = agent.tool_call_history();
20637 assert_eq!(history.len(), 1);
20638 assert_eq!(history[0].tool_id, "calculator");
20639 assert_eq!(
20640 history[0].result.get("result"),
20641 Some(&serde_json::json!(4.0)),
20642 "{:?}",
20643 history[0].result
20644 );
20645 }
20646
20647 #[test]
20650 fn legacy_tool_call_marker_is_plain_text() {
20651 let agent = AgentBuilder::new()
20652 .system_prompt("x")
20653 .llm(Arc::new(mock_with_response("x")))
20654 .build()
20655 .unwrap();
20656 let parsed = agent
20657 .parse_tool_calls(
20658 r#"[TOOL_CALL: {"name": "calculator", "arguments": {"expression": "2+2"}}]"#,
20659 )
20660 .unwrap();
20661 assert!(parsed.is_none());
20662 }
20663
20664 #[tokio::test]
20665 async fn test_tool_hitl_rejection_finalizes_blocking_turn() {
20666 let responses = Arc::new(std::sync::atomic::AtomicUsize::new(0));
20667 let hooks = Arc::new(ResponseCountingHooks {
20668 responses: Arc::clone(&responses),
20669 });
20670 let mock = mock_with_response(r#"{"tool":"echo","arguments":{"message":"hello"}}"#);
20671 let yaml = r#"
20672name: ToolRejectAgent
20673system_prompt: "You use tools when requested."
20674tools:
20675 - echo
20676hitl:
20677 tools:
20678 echo:
20679 require_approval: true
20680 approval_message: "Approve echo?"
20681"#;
20682 let agent = AgentBuilder::from_yaml(yaml)
20683 .unwrap()
20684 .llm(Arc::new(mock))
20685 .auto_configure_features()
20686 .unwrap()
20687 .hooks(hooks)
20688 .build()
20689 .unwrap();
20690
20691 let response = agent.chat("echo hello").await.unwrap();
20692
20693 assert!(
20694 response.content.contains("Operation cancelled"),
20695 "unexpected response: {}",
20696 response.content
20697 );
20698 assert_eq!(responses.load(Ordering::SeqCst), 1);
20699 let messages = agent.memory.get_messages(None).await.unwrap();
20700 assert_eq!(messages.len(), 3);
20701 assert_eq!(messages[0].content, "echo hello");
20702 assert!(messages[1].content.contains("\"tool\":\"echo\""));
20703 assert!(messages[2].content.contains("rejected by the approver"));
20704 }
20705
20706 #[tokio::test]
20707 async fn test_tool_hitl_rejection_finalizes_streaming_turn() {
20708 use futures::StreamExt;
20709
20710 let responses = Arc::new(std::sync::atomic::AtomicUsize::new(0));
20711 let hooks = Arc::new(ResponseCountingHooks {
20712 responses: Arc::clone(&responses),
20713 });
20714 let mock = mock_with_response(r#"{"tool":"echo","arguments":{"message":"hello"}}"#);
20715 let yaml = r#"
20716name: ToolRejectStreamingAgent
20717system_prompt: "You use tools when requested."
20718tools:
20719 - echo
20720streaming:
20721 enabled: true
20722hitl:
20723 tools:
20724 echo:
20725 require_approval: true
20726 approval_message: "Approve echo?"
20727"#;
20728 let agent = AgentBuilder::from_yaml(yaml)
20729 .unwrap()
20730 .llm(Arc::new(mock))
20731 .auto_configure_features()
20732 .unwrap()
20733 .hooks(hooks)
20734 .build()
20735 .unwrap();
20736
20737 let mut stream = agent.chat_stream("echo hello").await.unwrap();
20738 let mut terminal_error = String::new();
20739 let mut done = false;
20740 while let Some(chunk) = stream.next().await {
20741 match chunk {
20742 StreamChunk::Error { message } => terminal_error = message,
20743 StreamChunk::Done {} => {
20744 done = true;
20745 break;
20746 }
20747 _ => {}
20748 }
20749 }
20750
20751 assert!(done);
20752 assert!(
20753 terminal_error.contains("Operation cancelled"),
20754 "unexpected terminal error: {}",
20755 terminal_error
20756 );
20757 assert_eq!(responses.load(Ordering::SeqCst), 1);
20758 let messages = agent.memory.get_messages(None).await.unwrap();
20759 assert_eq!(messages.len(), 3);
20760 assert_eq!(messages[0].content, "echo hello");
20761 assert!(messages[1].content.contains("\"tool\":\"echo\""));
20762 assert!(messages[2].content.contains("rejected by the approver"));
20763 }
20764
20765 #[tokio::test]
20766 async fn tool_hitl_rejection_preserves_legacy_error_but_finalizes_event_stream() {
20767 let mock = mock_with_response(r#"{"tool":"echo","arguments":{"message":"hello"}}"#);
20768 let yaml = r#"
20769name: ToolRejectEventAgent
20770system_prompt: "You use tools when requested."
20771tools:
20772 - echo
20773streaming:
20774 enabled: true
20775hitl:
20776 tools:
20777 echo:
20778 require_approval: true
20779 approval_message: "Approve echo?"
20780"#;
20781 let agent = AgentBuilder::from_yaml(yaml)
20782 .unwrap()
20783 .llm(Arc::new(mock))
20784 .auto_configure_features()
20785 .unwrap()
20786 .build()
20787 .unwrap();
20788
20789 let mut stream = agent.chat_stream_events("echo hello").await.unwrap();
20790 let mut error_seen = false;
20791 let mut final_response = None;
20792 while let Some(event) = stream.next().await {
20793 match event {
20794 AgentStreamEvent::Chunk(StreamChunk::Error { .. }) => error_seen = true,
20795 AgentStreamEvent::Final(response) => final_response = Some(response),
20796 AgentStreamEvent::Chunk(_) => {}
20797 }
20798 }
20799
20800 assert!(!error_seen);
20801 assert!(
20802 final_response
20803 .is_some_and(|response| { response.content.contains("Operation cancelled") })
20804 );
20805 }
20806
20807 #[tokio::test]
20808 async fn test_pre_response_guard_transition_skips_old_state_llm() {
20809 let mock = mock_with_response("Billing state response");
20810 let call_counter = mock.clone();
20811 let yaml = r#"
20812name: OptimizedStateAgent
20813system_prompt: "You route before answering."
20814runtime:
20815 optimization:
20816 enabled: true
20817 pre_response_deterministic_transitions: true
20818states:
20819 initial: greeting
20820 states:
20821 greeting:
20822 prompt: "Old state prompt that should be skipped."
20823 transitions:
20824 - to: billing
20825 guard:
20826 context:
20827 topic:
20828 eq: billing
20829 timing: pre_response
20830 billing:
20831 prompt: "Answer from the billing state."
20832"#;
20833 let agent = AgentBuilder::from_yaml(yaml)
20834 .unwrap()
20835 .llm(Arc::new(mock))
20836 .build()
20837 .unwrap();
20838 agent
20839 .set_context("topic", serde_json::json!("billing"))
20840 .unwrap();
20841
20842 let response = agent.chat("I need billing help").await.unwrap();
20843
20844 assert_eq!(agent.current_state().as_deref(), Some("billing"));
20845 assert_eq!(response.content, "Billing state response");
20846 assert_eq!(call_counter.call_count(), 1);
20847 assert_eq!(agent.actor_facts().len(), 0);
20848 }
20849
20850 #[tokio::test]
20851 async fn test_set_context_supports_dotted_paths_for_pre_response_guards() {
20852 let mock = mock_with_response("Billing state response");
20853 let call_counter = mock.clone();
20854 let yaml = r#"
20855name: OptimizedStateAgent
20856system_prompt: "You route before answering."
20857runtime:
20858 optimization:
20859 enabled: true
20860 pre_response_deterministic_transitions: true
20861context:
20862 request:
20863 type: runtime
20864 default:
20865 topic: general
20866states:
20867 initial: greeting
20868 states:
20869 greeting:
20870 prompt: "Old state prompt that should be skipped."
20871 transitions:
20872 - to: billing
20873 guard:
20874 context:
20875 request.topic:
20876 eq: billing
20877 timing: pre_response
20878 billing:
20879 prompt: "Answer from the billing state."
20880"#;
20881 let agent = AgentBuilder::from_yaml(yaml)
20882 .unwrap()
20883 .llm(Arc::new(mock))
20884 .build()
20885 .unwrap();
20886 agent
20887 .set_context("request.topic", serde_json::json!("billing"))
20888 .unwrap();
20889
20890 let response = agent.chat("I need billing help").await.unwrap();
20891
20892 assert_eq!(agent.current_state().as_deref(), Some("billing"));
20893 assert_eq!(response.content, "Billing state response");
20894 assert_eq!(call_counter.call_count(), 1);
20895 assert_eq!(
20896 agent.get_context().get("request"),
20897 Some(&serde_json::json!({"topic": "billing"}))
20898 );
20899 }
20900
20901 #[tokio::test]
20902 async fn test_pre_response_rejection_does_not_commit_staged_context_or_user() {
20903 let mock = mock_with_response("billing");
20904 let yaml = r#"
20905name: OptimizedStateAgent
20906system_prompt: "You route before answering."
20907runtime:
20908 optimization:
20909 enabled: true
20910 pre_response_deterministic_transitions: true
20911hitl:
20912 states:
20913 billing:
20914 on_enter: require_approval
20915 approval_message: "Approve billing route?"
20916states:
20917 initial: greeting
20918 states:
20919 greeting:
20920 prompt: "Old state prompt."
20921 extract:
20922 - key: topic
20923 description: "Support topic"
20924 transitions:
20925 - to: billing
20926 guard:
20927 context:
20928 topic:
20929 eq: billing
20930 timing: pre_response
20931 run_extractors: true
20932 billing:
20933 prompt: "Billing state."
20934"#;
20935 let agent = AgentBuilder::from_yaml(yaml)
20936 .unwrap()
20937 .llm(Arc::new(mock))
20938 .build()
20939 .unwrap();
20940
20941 let response = agent
20942 .try_pre_response_transition("billing please")
20943 .await
20944 .unwrap();
20945
20946 assert!(response.is_none());
20947 assert_eq!(agent.current_state().as_deref(), Some("greeting"));
20948 assert!(!agent.get_context().contains_key("topic"));
20949 assert_eq!(agent.memory.get_messages(None).await.unwrap().len(), 0);
20950 }
20951
20952 #[tokio::test]
20953 async fn test_pre_response_extractor_commits_context_on_winning_path() {
20954 let mock = mock_with_responses(vec!["billing", "Billing response"]);
20955 let yaml = r#"
20956name: OptimizedStateAgent
20957system_prompt: "You route before answering."
20958runtime:
20959 optimization:
20960 enabled: true
20961 pre_response_deterministic_transitions: true
20962states:
20963 initial: greeting
20964 states:
20965 greeting:
20966 prompt: "Old state prompt."
20967 extract:
20968 - key: topic
20969 description: "Support topic"
20970 transitions:
20971 - to: billing
20972 guard:
20973 context:
20974 topic:
20975 eq: billing
20976 timing: pre_response
20977 run_extractors: true
20978 billing:
20979 prompt: "Billing state."
20980"#;
20981 let agent = AgentBuilder::from_yaml(yaml)
20982 .unwrap()
20983 .llm(Arc::new(mock))
20984 .build()
20985 .unwrap();
20986
20987 let response = agent.chat("billing please").await.unwrap();
20988
20989 assert_eq!(agent.current_state().as_deref(), Some("billing"));
20990 assert_eq!(response.content, "Billing response");
20991 assert_eq!(
20992 agent.get_context().get("topic"),
20993 Some(&serde_json::json!("billing"))
20994 );
20995 }
20996
20997 #[tokio::test]
20998 async fn test_pre_response_extractor_miss_does_not_mutate_context() {
20999 let mock = mock_with_response("__NONE__");
21000 let yaml = r#"
21001name: OptimizedStateAgent
21002system_prompt: "You route before answering."
21003runtime:
21004 optimization:
21005 enabled: true
21006 pre_response_deterministic_transitions: true
21007states:
21008 initial: greeting
21009 states:
21010 greeting:
21011 prompt: "Old state prompt."
21012 extract:
21013 - key: topic
21014 description: "Support topic"
21015 transitions:
21016 - to: billing
21017 guard:
21018 context:
21019 topic:
21020 eq: billing
21021 timing: pre_response
21022 run_extractors: true
21023 billing:
21024 prompt: "Billing state."
21025"#;
21026 let agent = AgentBuilder::from_yaml(yaml)
21027 .unwrap()
21028 .llm(Arc::new(mock))
21029 .build()
21030 .unwrap();
21031
21032 let response = agent.try_pre_response_transition("hello").await.unwrap();
21033
21034 assert!(response.is_none());
21035 assert_eq!(agent.current_state().as_deref(), Some("greeting"));
21036 assert!(!agent.get_context().contains_key("topic"));
21037 }
21038
21039 #[tokio::test]
21040 async fn test_default_guard_transition_stays_post_response() {
21041 let mock = mock_with_responses(vec!["Greeting response", "Billing response"]);
21042 let call_counter = mock.clone();
21043 let yaml = r#"
21044name: TimingAgent
21045system_prompt: "You route carefully."
21046runtime:
21047 optimization:
21048 enabled: true
21049 pre_response_deterministic_transitions: true
21050states:
21051 initial: greeting
21052 states:
21053 greeting:
21054 prompt: "Old state prompt."
21055 transitions:
21056 - to: billing
21057 guard:
21058 context:
21059 topic:
21060 eq: billing
21061 billing:
21062 prompt: "Billing state."
21063"#;
21064 let agent = AgentBuilder::from_yaml(yaml)
21065 .unwrap()
21066 .llm(Arc::new(mock))
21067 .build()
21068 .unwrap();
21069 agent
21070 .set_context("topic", serde_json::json!("billing"))
21071 .unwrap();
21072
21073 let response = agent.chat("billing please").await.unwrap();
21074
21075 assert_eq!(agent.current_state().as_deref(), Some("billing"));
21076 assert_eq!(response.content, "Billing response");
21077 assert_eq!(call_counter.call_count(), 2);
21078 }
21079
21080 #[tokio::test]
21081 async fn test_explicit_post_response_guard_transition_stays_post_response() {
21082 let mock = mock_with_responses(vec!["Greeting response", "Billing response"]);
21083 let call_counter = mock.clone();
21084 let yaml = r#"
21085name: TimingAgent
21086system_prompt: "You route carefully."
21087runtime:
21088 optimization:
21089 enabled: true
21090 pre_response_deterministic_transitions: true
21091states:
21092 initial: greeting
21093 states:
21094 greeting:
21095 prompt: "Old state prompt."
21096 transitions:
21097 - to: billing
21098 guard:
21099 context:
21100 topic:
21101 eq: billing
21102 timing: post_response
21103 billing:
21104 prompt: "Billing state."
21105"#;
21106 let agent = AgentBuilder::from_yaml(yaml)
21107 .unwrap()
21108 .llm(Arc::new(mock))
21109 .build()
21110 .unwrap();
21111 agent
21112 .set_context("topic", serde_json::json!("billing"))
21113 .unwrap();
21114
21115 let response = agent.chat("billing please").await.unwrap();
21116
21117 assert_eq!(agent.current_state().as_deref(), Some("billing"));
21118 assert_eq!(response.content, "Billing response");
21119 assert_eq!(call_counter.call_count(), 2);
21120 }
21121
21122 #[tokio::test]
21123 async fn test_pre_response_extractors_are_transition_scoped() {
21124 let mock = mock_with_responses(vec!["billing", "Billing response"]);
21125 let yaml = r#"
21126name: ScopedExtractorAgent
21127system_prompt: "You route carefully."
21128runtime:
21129 optimization:
21130 enabled: true
21131 pre_response_deterministic_transitions: true
21132states:
21133 initial: greeting
21134 states:
21135 greeting:
21136 prompt: "Old state prompt."
21137 extract:
21138 - key: topic
21139 description: "Support topic"
21140 transitions:
21141 - to: wrong
21142 guard:
21143 context:
21144 topic:
21145 eq: billing
21146 timing: pre_response
21147 - to: billing
21148 guard:
21149 context:
21150 topic:
21151 eq: billing
21152 timing: pre_response
21153 run_extractors: true
21154 wrong:
21155 prompt: "Wrong state."
21156 billing:
21157 prompt: "Billing state."
21158"#;
21159 let agent = AgentBuilder::from_yaml(yaml)
21160 .unwrap()
21161 .llm(Arc::new(mock))
21162 .build()
21163 .unwrap();
21164
21165 let response = agent.chat("billing please").await.unwrap();
21166
21167 assert_eq!(agent.current_state().as_deref(), Some("billing"));
21168 assert_eq!(response.content, "Billing response");
21169 }
21170
21171 #[tokio::test]
21172 async fn test_pre_response_resolved_intent_routes_early() {
21173 let mock = mock_with_response("Billing response");
21174 let yaml = r#"
21175name: IntentAgent
21176system_prompt: "You route carefully."
21177runtime:
21178 optimization:
21179 enabled: true
21180 pre_response_deterministic_transitions: true
21181states:
21182 initial: greeting
21183 states:
21184 greeting:
21185 prompt: "Old state prompt."
21186 transitions:
21187 - to: billing
21188 intent: billing
21189 timing: pre_response
21190 billing:
21191 prompt: "Billing state."
21192"#;
21193 let agent = AgentBuilder::from_yaml(yaml)
21194 .unwrap()
21195 .llm(Arc::new(mock))
21196 .build()
21197 .unwrap();
21198 agent
21199 .set_context("resolved_intent", serde_json::json!("billing"))
21200 .unwrap();
21201
21202 let response = agent
21203 .try_pre_response_transition("I need billing help")
21204 .await
21205 .unwrap()
21206 .unwrap();
21207
21208 assert_eq!(agent.current_state().as_deref(), Some("billing"));
21209 assert_eq!(response.content, "Billing response");
21210 }
21211
21212 #[tokio::test]
21213 async fn test_background_overflow_error_surfaces() {
21214 let mut config = RuntimeConfig::default();
21215 config.optimization.enabled = true;
21216 config.optimization.post_turn.max_background_tasks = 1;
21217 config.optimization.post_turn.on_background_overflow = BackgroundOverflowPolicy::Error;
21218 let policy = crate::optimization::MaintenanceTaskPolicy {
21219 mode: MaintenanceMode::Background,
21220 await_before_next_turn: AwaitBeforeNextTurn::Always,
21221 };
21222 let agent = AgentBuilder::new()
21223 .system_prompt("You are helpful.")
21224 .llm(Arc::new(mock_with_response("ok")))
21225 .build()
21226 .unwrap()
21227 .with_runtime_config(config);
21228 agent
21229 .background_maintenance
21230 .spawn(None, async { std::future::pending::<Result<()>>().await })
21231 .unwrap();
21232
21233 let result = agent
21234 .spawn_or_handle_background(None, async { Ok(()) }, "facts", &policy)
21235 .await;
21236
21237 assert!(result.is_err());
21238 }
21239
21240 #[tokio::test]
21241 async fn test_speculative_reasoning_low_cap_uses_serial_reasoning() {
21242 let default_mock = mock_with_response("Plain draft response");
21243 let router_mock = mock_with_response("cot");
21244 let router_counter = router_mock.clone();
21245 let yaml = r#"
21246name: ReasoningReservationAgent
21247system_prompt: "You answer plainly unless reasoning wins."
21248llm:
21249 default: default
21250 router: router
21251observability:
21252 enabled: true
21253 export:
21254 write_raw_events: true
21255reasoning:
21256 mode: auto
21257 judge_llm: router
21258runtime:
21259 optimization:
21260 enabled: true
21261 max_speculative_llm_calls_per_turn: 1
21262 speculative_reasoning_auto: true
21263 max_parallel_runtime_tasks: 2
21264"#;
21265 let agent = AgentBuilder::from_yaml(yaml)
21266 .unwrap()
21267 .llm_alias("default", Arc::new(default_mock))
21268 .llm_alias("router", Arc::new(router_mock))
21269 .build()
21270 .unwrap();
21271
21272 let response = agent.chat("hello").await.unwrap();
21273
21274 assert_eq!(response.content, "Plain draft response");
21275 assert_eq!(router_counter.call_count(), 1);
21276 let events = agent.observability().unwrap().raw_events();
21277 assert!(!events.iter().any(|event| {
21278 event.dimensions.get("commit_behavior") == Some(&"reasoning_decision".to_string())
21279 }));
21280 }
21281
21282 #[tokio::test]
21283 async fn test_forced_reasoning_skips_plain_speculative_draft() {
21284 let mock = mock_with_response("Reasoned response");
21285 let yaml = r#"
21286name: ForcedReasoningAgent
21287system_prompt: "You reason before answering."
21288observability:
21289 enabled: true
21290 export:
21291 write_raw_events: true
21292reasoning:
21293 mode: cot
21294runtime:
21295 optimization:
21296 enabled: true
21297 max_speculative_llm_calls_per_turn: 2
21298 speculative_state_transitions: true
21299 max_parallel_runtime_tasks: 2
21300states:
21301 initial: triage
21302 states:
21303 triage:
21304 prompt: "Answer from triage."
21305 transitions:
21306 - to: billing
21307 guard:
21308 context:
21309 route:
21310 eq: billing
21311 timing: parallel
21312 billing:
21313 prompt: "Billing state."
21314"#;
21315 let agent = AgentBuilder::from_yaml(yaml)
21316 .unwrap()
21317 .llm(Arc::new(mock))
21318 .build()
21319 .unwrap();
21320
21321 let response = agent.chat("hello").await.unwrap();
21322
21323 assert_eq!(response.content, "Reasoned response");
21324 let events = agent.observability().unwrap().raw_events();
21325 assert!(
21326 !events
21327 .iter()
21328 .any(|event| event.dimensions.contains_key("branch_status"))
21329 );
21330 }
21331
21332 #[tokio::test]
21333 async fn test_speculative_skill_low_cap_uses_serial_skill_route() {
21334 let default_mock = mock_with_response("Skill committed response");
21335 let router_mock = mock_with_response("helper");
21336 let router_counter = router_mock.clone();
21337 let yaml = r#"
21338name: SkillReservationAgent
21339system_prompt: "Use skills when they match."
21340llm:
21341 default: default
21342 router: router
21343observability:
21344 enabled: true
21345 export:
21346 write_raw_events: true
21347runtime:
21348 optimization:
21349 enabled: true
21350 max_speculative_llm_calls_per_turn: 1
21351 speculative_skill_routing: true
21352 max_parallel_runtime_tasks: 2
21353skills:
21354 - id: helper
21355 description: "Answer helper requests"
21356 trigger: "User asks for helper"
21357 steps:
21358 - prompt: "Answer the helper request: {{ user_input }}"
21359"#;
21360 let agent = AgentBuilder::from_yaml(yaml)
21361 .unwrap()
21362 .llm_alias("default", Arc::new(default_mock))
21363 .llm_alias("router", Arc::new(router_mock))
21364 .build()
21365 .unwrap();
21366
21367 let response = agent.chat("please use helper").await.unwrap();
21368
21369 assert_eq!(response.content, "Skill committed response");
21370 assert_eq!(router_counter.call_count(), 1);
21371 let events = agent.observability().unwrap().raw_events();
21372 assert!(
21373 !events
21374 .iter()
21375 .any(|event| event.dimensions.contains_key("branch_status"))
21376 );
21377 }
21378
21379 #[tokio::test]
21380 async fn test_parallel_transition_low_cap_allows_deterministic_route() {
21381 let mock = mock_with_response("unused");
21382 let call_counter = mock.clone();
21383 let yaml = r#"
21384name: ParallelTransitionLowCapAgent
21385system_prompt: "Route before stale responses when safe."
21386runtime:
21387 optimization:
21388 enabled: true
21389 max_speculative_llm_calls_per_turn: 1
21390 speculative_state_transitions: true
21391 max_parallel_runtime_tasks: 2
21392states:
21393 initial: triage
21394 states:
21395 triage:
21396 prompt: "Triage state."
21397 transitions:
21398 - to: billing
21399 guard:
21400 context:
21401 route:
21402 eq: billing
21403 timing: parallel
21404 billing:
21405 prompt: "Billing state."
21406"#;
21407 let agent = AgentBuilder::from_yaml(yaml)
21408 .unwrap()
21409 .llm(Arc::new(mock))
21410 .build()
21411 .unwrap();
21412 agent
21413 .set_context("route", serde_json::json!("billing"))
21414 .unwrap();
21415 agent.update_active_turn_context("billing help", HashMap::new());
21416 assert!(
21417 agent.reserve_active_speculative_llm_call(
21418 RuntimeOptimizationKind::ParallelStateTransition
21419 )
21420 );
21421
21422 let selection = agent
21423 .select_parallel_transition_candidate("billing help")
21424 .await
21425 .unwrap();
21426 agent.end_root_turn();
21427
21428 match selection {
21429 ParallelTransitionSelection::Candidate(candidate) => {
21430 assert_eq!(candidate.target(), "billing");
21431 }
21432 ParallelTransitionSelection::NoMatch => panic!("deterministic route did not match"),
21433 ParallelTransitionSelection::ReservationExhausted => {
21434 panic!("deterministic route consumed LLM budget")
21435 }
21436 }
21437 assert_eq!(call_counter.call_count(), 0);
21438 }
21439
21440 #[tokio::test]
21441 async fn speculative_transition_drops_loser_before_state_actions() {
21442 let lock = Arc::new(tokio::sync::Mutex::new(()));
21443 let first_started = Arc::new(tokio::sync::Notify::new());
21444 let first_dropped = Arc::new(AtomicBool::new(false));
21445 let committed_after_drop = Arc::new(AtomicBool::new(false));
21446 let default = Arc::new(FirstCallLockingProvider {
21447 lock,
21448 first_started: Arc::clone(&first_started),
21449 first_dropped: Arc::clone(&first_dropped),
21450 committed_after_drop: Arc::clone(&committed_after_drop),
21451 calls: AtomicU64::new(0),
21452 });
21453 let router = Arc::new(RoutingAfterProviderStart {
21454 provider_started: first_started,
21455 });
21456 let yaml = r#"
21457name: SpeculativeCancellationAgent
21458system_prompt: "Route before committed work."
21459llm:
21460 default: default
21461 router: router
21462runtime:
21463 optimization:
21464 enabled: true
21465 max_speculative_llm_calls_per_turn: 2
21466 speculative_state_transitions: true
21467 max_parallel_runtime_tasks: 2
21468states:
21469 initial: triage
21470 states:
21471 triage:
21472 prompt: "Triage state."
21473 transitions:
21474 - to: technical
21475 when: "The request needs technical support"
21476 timing: parallel
21477 technical:
21478 prompt: "Technical state."
21479 on_enter:
21480 - prompt: "Prepare technical context."
21481 llm: default
21482 store_as: preparation
21483"#;
21484 let agent = AgentBuilder::from_yaml(yaml)
21485 .unwrap()
21486 .llm_alias("default", default)
21487 .llm_alias("router", router)
21488 .build()
21489 .unwrap();
21490
21491 let response = tokio::time::timeout(
21492 std::time::Duration::from_secs(2),
21493 agent.chat("I cannot log in because of AUTH-17."),
21494 )
21495 .await
21496 .expect("committed work must not wait on the losing provider future")
21497 .unwrap();
21498
21499 assert_eq!(response.content, "Committed technical response.");
21500 assert_eq!(agent.current_state().as_deref(), Some("technical"));
21501 assert!(first_dropped.load(Ordering::SeqCst));
21502 assert!(committed_after_drop.load(Ordering::SeqCst));
21503 }
21504
21505 #[tokio::test]
21506 async fn buffered_transition_drops_stale_stream_before_redispatch() {
21507 use futures::StreamExt;
21508
21509 let lock = Arc::new(tokio::sync::Mutex::new(()));
21510 let stream_started = Arc::new(tokio::sync::Notify::new());
21511 let stream_dropped = Arc::new(AtomicBool::new(false));
21512 let committed_after_drop = Arc::new(AtomicBool::new(false));
21513 let default = Arc::new(BufferedLockingProvider {
21514 lock,
21515 stream_started: Arc::clone(&stream_started),
21516 stream_dropped: Arc::clone(&stream_dropped),
21517 committed_after_drop: Arc::clone(&committed_after_drop),
21518 });
21519 let router = Arc::new(RoutingAfterProviderStart {
21520 provider_started: stream_started,
21521 });
21522 let yaml = r#"
21523name: BufferedCancellationAgent
21524system_prompt: "Hide stale streamed output."
21525llm:
21526 default: default
21527 router: router
21528streaming:
21529 enabled: true
21530 buffer_size: 8
21531runtime:
21532 optimization:
21533 enabled: true
21534 max_speculative_llm_calls_per_turn: 2
21535 speculative_state_transitions: true
21536 streaming_policy: buffer_until_routing_done
21537 max_parallel_runtime_tasks: 2
21538states:
21539 initial: triage
21540 states:
21541 triage:
21542 prompt: "Triage state."
21543 transitions:
21544 - to: technical
21545 when: "The request needs technical support"
21546 timing: parallel
21547 technical:
21548 prompt: "Technical state."
21549"#;
21550 let agent = AgentBuilder::from_yaml(yaml)
21551 .unwrap()
21552 .llm_alias("default", default)
21553 .llm_alias("router", router)
21554 .build()
21555 .unwrap();
21556
21557 let content = tokio::time::timeout(std::time::Duration::from_secs(2), async {
21558 let mut stream = agent
21559 .chat_stream("AUTH-17 needs technical help.")
21560 .await
21561 .unwrap();
21562 let mut content = String::new();
21563 while let Some(chunk) = stream.next().await {
21564 match chunk {
21565 StreamChunk::Content { text } => content.push_str(&text),
21566 StreamChunk::Done {} => break,
21567 StreamChunk::Error { message } => panic!("unexpected stream error: {message}"),
21568 _ => {}
21569 }
21570 }
21571 content
21572 })
21573 .await
21574 .expect("redispatch must not wait on the stale streaming future");
21575
21576 assert_eq!(content, "Committed technical response.");
21577 assert_eq!(agent.current_state().as_deref(), Some("technical"));
21578 assert!(stream_dropped.load(Ordering::SeqCst));
21579 assert!(committed_after_drop.load(Ordering::SeqCst));
21580 }
21581
21582 #[tokio::test]
21583 async fn buffered_transition_drops_established_stream_before_redispatch() {
21584 use futures::StreamExt;
21585
21586 let stream_started = Arc::new(tokio::sync::Notify::new());
21587 let stream_dropped = Arc::new(AtomicBool::new(false));
21588 let stream_dropped_notify = Arc::new(tokio::sync::Notify::new());
21589 let committed_after_drop = Arc::new(AtomicBool::new(false));
21590 let default = Arc::new(EstablishedStreamProvider {
21591 stream_started: Arc::clone(&stream_started),
21592 stream_dropped: Arc::clone(&stream_dropped),
21593 stream_dropped_notify,
21594 committed_after_drop: Arc::clone(&committed_after_drop),
21595 });
21596 let router = Arc::new(RoutingAfterProviderStart {
21597 provider_started: stream_started,
21598 });
21599 let yaml = r#"
21600name: EstablishedStreamCancellationAgent
21601system_prompt: "Hide stale streamed output."
21602llm:
21603 default: default
21604 router: router
21605streaming:
21606 enabled: true
21607 buffer_size: 8
21608runtime:
21609 optimization:
21610 enabled: true
21611 max_speculative_llm_calls_per_turn: 2
21612 speculative_state_transitions: true
21613 streaming_policy: buffer_until_routing_done
21614 max_parallel_runtime_tasks: 2
21615states:
21616 initial: triage
21617 states:
21618 triage:
21619 prompt: "Triage state."
21620 transitions:
21621 - to: technical
21622 when: "The request needs technical support"
21623 timing: parallel
21624 technical:
21625 prompt: "Technical state."
21626"#;
21627 let agent = AgentBuilder::from_yaml(yaml)
21628 .unwrap()
21629 .llm_alias("default", default)
21630 .llm_alias("router", router)
21631 .build()
21632 .unwrap();
21633
21634 let content = tokio::time::timeout(std::time::Duration::from_secs(2), async {
21635 let mut stream = agent
21636 .chat_stream("AUTH-17 needs technical help.")
21637 .await
21638 .unwrap();
21639 let mut content = String::new();
21640 while let Some(chunk) = stream.next().await {
21641 match chunk {
21642 StreamChunk::Content { text } => content.push_str(&text),
21643 StreamChunk::Done {} => break,
21644 StreamChunk::Error { message } => panic!("unexpected stream error: {message}"),
21645 _ => {}
21646 }
21647 }
21648 content
21649 })
21650 .await
21651 .expect("redispatch must wait for the established stale stream to be dropped");
21652
21653 assert_eq!(content, "Committed technical response.");
21654 assert_eq!(agent.current_state().as_deref(), Some("technical"));
21655 assert!(stream_dropped.load(Ordering::SeqCst));
21656 assert!(committed_after_drop.load(Ordering::SeqCst));
21657 }
21658
21659 #[tokio::test]
21660 async fn test_buffered_streaming_transition_reservation_falls_back() {
21661 use futures::StreamExt;
21662
21663 let mock = mock_with_responses(vec![
21664 "Serial streaming response",
21665 "Serial streaming response",
21666 ]);
21667 let router_mock = mock_with_response("1");
21668 let router_counter = router_mock.clone();
21669 let yaml = r#"
21670name: BufferedReservationFallbackAgent
21671system_prompt: "Stream normally if speculative routing cannot be evaluated."
21672llm:
21673 default: default
21674 router: router
21675observability:
21676 enabled: true
21677 export:
21678 write_raw_events: true
21679streaming:
21680 enabled: true
21681 buffer_size: 8
21682runtime:
21683 optimization:
21684 enabled: true
21685 max_speculative_llm_calls_per_turn: 1
21686 speculative_state_transitions: true
21687 streaming_policy: buffer_until_routing_done
21688 max_parallel_runtime_tasks: 2
21689states:
21690 initial: triage
21691 states:
21692 triage:
21693 prompt: "Triage state."
21694 transitions:
21695 - to: billing
21696 guard:
21697 context:
21698 route:
21699 eq: billing
21700 when: "User asks about billing"
21701 timing: parallel
21702 billing:
21703 prompt: "Billing state."
21704"#;
21705 let agent = AgentBuilder::from_yaml(yaml)
21706 .unwrap()
21707 .llm_alias("default", Arc::new(mock))
21708 .llm_alias("router", Arc::new(router_mock))
21709 .build()
21710 .unwrap();
21711
21712 let mut stream = agent.chat_stream("hello").await.unwrap();
21713 let mut content = String::new();
21714 let mut error = None;
21715 while let Some(chunk) = stream.next().await {
21716 match chunk {
21717 StreamChunk::Content { text } => content.push_str(&text),
21718 StreamChunk::Error { message } => error = Some(message),
21719 StreamChunk::Done {} => break,
21720 _ => {}
21721 }
21722 }
21723
21724 assert_eq!(error, None);
21725 assert_eq!(content, "Serial streaming response");
21726 assert_eq!(router_counter.call_count(), 0);
21727 let events = agent.observability().unwrap().raw_events();
21728 assert!(events.iter().any(|event| {
21729 event.dimensions.get("branch_status") == Some(&"cancelled".to_string())
21730 && event.dimensions.get("commit_behavior")
21731 == Some(&"transition_decision".to_string())
21732 }));
21733 }
21734
21735 #[tokio::test]
21736 async fn test_blocking_error_cleanup_resets_root_turn_for_next_chat() {
21737 let mut mock = mock_with_response("Recovered response");
21738 mock.set_error("boom");
21739 let mut handle = mock.clone();
21740 let agent = AgentBuilder::new()
21741 .system_prompt("You are helpful.")
21742 .llm(Arc::new(mock))
21743 .build()
21744 .unwrap();
21745
21746 assert!(agent.chat("first").await.is_err());
21747 handle.clear_error();
21748 let response = agent.chat("second").await.unwrap();
21749
21750 assert_eq!(response.content, "Recovered response");
21751 let messages = agent.memory.get_messages(None).await.unwrap();
21752 let user_count = messages
21753 .iter()
21754 .filter(|message| message.role == ai_agents_core::Role::User)
21755 .count();
21756 assert_eq!(user_count, 2);
21757 }
21758
21759 #[tokio::test]
21760 async fn test_streaming_error_cleanup_resets_root_turn_for_next_chat() {
21761 use futures::StreamExt;
21762
21763 let mut mock = mock_with_response("Recovered response");
21764 mock.set_error("stream boom");
21765 let mut handle = mock.clone();
21766 let agent = AgentBuilder::new()
21767 .system_prompt("You are helpful.")
21768 .llm(Arc::new(mock))
21769 .build()
21770 .unwrap();
21771
21772 let mut stream = agent.chat_stream("first").await.unwrap();
21773 let mut saw_error = false;
21774 while let Some(chunk) = stream.next().await {
21775 if matches!(chunk, StreamChunk::Error { .. }) {
21776 saw_error = true;
21777 }
21778 }
21779 assert!(saw_error);
21780
21781 handle.clear_error();
21782 let response = agent.chat("second").await.unwrap();
21783
21784 assert_eq!(response.content, "Recovered response");
21785 let messages = agent.memory.get_messages(None).await.unwrap();
21786 let user_count = messages
21787 .iter()
21788 .filter(|message| message.role == ai_agents_core::Role::User)
21789 .count();
21790 assert_eq!(user_count, 2);
21791 }
21792
21793 #[tokio::test]
21794 async fn test_buffered_streaming_route_miss_releases_buffer_limit() {
21795 use futures::StreamExt;
21796
21797 let mut mock = mock_with_response("one two three");
21798 mock.set_latency(10);
21799 let yaml = r#"
21800name: BufferedMissAgent
21801system_prompt: "You stream safely."
21802llm:
21803 default: default
21804streaming:
21805 enabled: true
21806 buffer_size: 1
21807runtime:
21808 optimization:
21809 enabled: true
21810 max_speculative_llm_calls_per_turn: 2
21811 speculative_state_transitions: true
21812 streaming_policy: buffer_until_routing_done
21813 max_parallel_runtime_tasks: 2
21814states:
21815 initial: triage
21816 states:
21817 triage:
21818 prompt: "Answer from triage."
21819 transitions:
21820 - to: billing
21821 guard:
21822 context:
21823 route:
21824 eq: billing
21825 timing: parallel
21826 billing:
21827 prompt: "Billing state."
21828"#;
21829 let agent = AgentBuilder::from_yaml(yaml)
21830 .unwrap()
21831 .llm_alias("default", Arc::new(mock))
21832 .build()
21833 .unwrap();
21834
21835 let mut stream = agent.chat_stream("hello").await.unwrap();
21836 let mut content = String::new();
21837 let mut error = None;
21838 while let Some(chunk) = stream.next().await {
21839 match chunk {
21840 StreamChunk::Content { text } => content.push_str(&text),
21841 StreamChunk::Error { message } => error = Some(message),
21842 StreamChunk::Done {} => break,
21843 _ => {}
21844 }
21845 }
21846
21847 assert_eq!(error, None);
21848 assert_eq!(content, "one two three");
21849 }
21850
21851 #[tokio::test]
21852 async fn test_buffered_streaming_main_failure_finalizes_branch() {
21853 use futures::StreamExt;
21854
21855 let mock = mock_with_response("one two");
21856 let mut router_mock = mock_with_response("0");
21857 router_mock.set_latency(50);
21858 let yaml = r#"
21859name: BufferedFailureAgent
21860system_prompt: "You stream safely."
21861llm:
21862 default: default
21863 router: router
21864observability:
21865 enabled: true
21866 export:
21867 write_raw_events: true
21868streaming:
21869 enabled: true
21870 buffer_size: 1
21871runtime:
21872 optimization:
21873 enabled: true
21874 max_speculative_llm_calls_per_turn: 2
21875 speculative_state_transitions: true
21876 streaming_policy: buffer_until_routing_done
21877 max_parallel_runtime_tasks: 2
21878states:
21879 initial: triage
21880 states:
21881 triage:
21882 prompt: "Ask for the category."
21883 transitions:
21884 - to: billing
21885 when: "User asks about billing"
21886 timing: parallel
21887 billing:
21888 prompt: "Billing state."
21889"#;
21890 let agent = AgentBuilder::from_yaml(yaml)
21891 .unwrap()
21892 .llm_alias("default", Arc::new(mock))
21893 .llm_alias("router", Arc::new(router_mock))
21894 .build()
21895 .unwrap();
21896
21897 let mut stream = agent.chat_stream("hello").await.unwrap();
21898 let mut error = String::new();
21899 while let Some(chunk) = stream.next().await {
21900 if let StreamChunk::Error { message } = chunk {
21901 error = message;
21902 }
21903 }
21904
21905 assert!(
21906 error.contains("stream buffer filled"),
21907 "unexpected stream error: {}",
21908 error
21909 );
21910 let events = agent.observability().unwrap().raw_events();
21911 assert!(events.iter().any(|event| {
21912 event.dimensions.get("branch_status") == Some(&"failed".to_string())
21913 && event.dimensions.get("commit_behavior") == Some(&"final_response".to_string())
21914 && event.dimensions.get("optimization")
21915 == Some(&"buffered_streaming_routing".to_string())
21916 }));
21917 }
21918
21919 #[tokio::test]
21920 async fn test_streaming_preflight_does_not_emit_old_state_content() {
21921 use futures::StreamExt;
21922
21923 let mock = mock_with_response("Billing streamed response");
21924 let yaml = r#"
21925name: StreamingOptimizedAgent
21926system_prompt: "You route before streaming."
21927runtime:
21928 optimization:
21929 enabled: true
21930 pre_response_deterministic_transitions: true
21931streaming:
21932 enabled: true
21933states:
21934 initial: greeting
21935 states:
21936 greeting:
21937 prompt: "OLD_STATE_SENTINEL"
21938 transitions:
21939 - to: billing
21940 guard:
21941 context:
21942 topic:
21943 eq: billing
21944 timing: pre_response
21945 billing:
21946 prompt: "Billing state."
21947"#;
21948 let agent = AgentBuilder::from_yaml(yaml)
21949 .unwrap()
21950 .llm(Arc::new(mock))
21951 .build()
21952 .unwrap();
21953 agent
21954 .set_context("topic", serde_json::json!("billing"))
21955 .unwrap();
21956
21957 let mut stream = agent.chat_stream("billing please").await.unwrap();
21958 let mut content = String::new();
21959 while let Some(chunk) = stream.next().await {
21960 match chunk {
21961 StreamChunk::Content { text } => content.push_str(&text),
21962 StreamChunk::Error { message } => panic!("stream error: {}", message),
21963 StreamChunk::Done {} => break,
21964 _ => {}
21965 }
21966 }
21967
21968 assert_eq!(agent.current_state().as_deref(), Some("billing"));
21969 assert!(content.contains("Billing streamed response"));
21970 assert!(!content.contains("OLD_STATE_SENTINEL"));
21971 }
21972
21973 #[tokio::test]
21975 async fn test_integration_state_machine_basic() {
21976 let yaml = r#"
21977name: StateAgent
21978system_prompt: "You are a support agent."
21979states:
21980 initial: greeting
21981 states:
21982 greeting:
21983 prompt: "Welcome the user warmly."
21984 transitions:
21985 - to: support
21986 when: "User needs help"
21987 auto: true
21988 support:
21989 prompt: "Help solve the user's problem."
21990"#;
21991 let mock = mock_with_responses(vec![
21992 "Welcome! How can I help?", "1", "I'll help you with that.", ]);
21996 let builder = AgentBuilder::from_yaml(yaml).unwrap();
21997 let agent = builder.llm(Arc::new(mock)).build().unwrap();
21998
21999 assert_eq!(agent.current_state(), Some("greeting".to_string()));
22000 let _ = agent.chat("I need help").await.unwrap();
22001 }
22004
22005 #[tokio::test]
22007 async fn test_integration_state_on_enter_set_context() {
22008 let yaml = r#"
22009name: ActionAgent
22010system_prompt: "You are helpful."
22011states:
22012 initial: step1
22013 states:
22014 step1:
22015 prompt: "Step 1"
22016 on_exit:
22017 - set_context:
22018 step1_exited: true
22019 transitions:
22020 - to: step2
22021 when: "always"
22022 auto: true
22023 step2:
22024 prompt: "Step 2"
22025 on_enter:
22026 - set_context:
22027 step2_entered: true
22028"#;
22029 let mock = mock_with_responses(vec![
22031 "Processing step 1.",
22032 "0", ]);
22034 let builder = AgentBuilder::from_yaml(yaml).unwrap();
22035 let agent = builder.llm(Arc::new(mock)).build().unwrap();
22036
22037 assert_eq!(agent.current_state(), Some("step1".to_string()));
22038
22039 agent.transition_to("step2").await.unwrap();
22041
22042 assert_eq!(agent.current_state(), Some("step2".to_string()));
22043
22044 let ctx = agent.get_context();
22046 assert_eq!(ctx.get("step1_exited"), Some(&serde_json::json!(true)));
22047 assert_eq!(ctx.get("step2_entered"), Some(&serde_json::json!(true)));
22048 }
22049
22050 #[tokio::test]
22051 async fn state_action_tool_preserves_source_in_stored_record() {
22052 let yaml = r#"
22053name: StateActionToolAgent
22054system_prompt: "You are helpful."
22055tools:
22056 - context_echo
22057states:
22058 initial: idle
22059 states:
22060 idle:
22061 prompt: "Idle"
22062 active:
22063 prompt: "Active"
22064 on_enter:
22065 - set_context:
22066 action_started: true
22067 - tool: context_echo
22068 args: {}
22069"#;
22070 let agent = AgentBuilder::from_yaml(yaml)
22071 .unwrap()
22072 .llm(Arc::new(mock_with_response("unused")))
22073 .tool(Arc::new(ContextEchoTool))
22074 .build()
22075 .unwrap();
22076
22077 agent.transition_to("active").await.unwrap();
22078
22079 let record: ToolExecutionRecord = serde_json::from_value(
22080 agent
22081 .get_context()
22082 .get("last_tool_record")
22083 .cloned()
22084 .expect("successful state action must store its execution record"),
22085 )
22086 .unwrap();
22087 assert!(record.executed);
22088 assert!(record.success);
22089 assert_eq!(record.canonical_id, "context_echo");
22090 assert!(matches!(
22091 &record.source,
22092 ToolCallSource::StateAction {
22093 state: Some(state),
22094 action_index: 1,
22095 } if state == "active"
22096 ));
22097 }
22098
22099 #[tokio::test]
22100 async fn test_ordinary_transition_uses_on_enter_then_on_reenter() {
22101 let yaml = r#"
22102name: OrdinaryLifecycleAgent
22103system_prompt: "You are helpful."
22104states:
22105 initial: intake
22106 regenerate_on_transition: false
22107 states:
22108 intake:
22109 prompt: "Intake"
22110 transitions:
22111 - to: drafting
22112 guard:
22113 context:
22114 route:
22115 eq: drafting
22116 drafting:
22117 prompt: "Drafting"
22118 on_enter:
22119 - set_context:
22120 draft_version: 1
22121 on_reenter:
22122 - set_context:
22123 draft_version: 2
22124 transitions:
22125 - to: review
22126 guard:
22127 context:
22128 route:
22129 eq: review
22130 review:
22131 prompt: "Review"
22132 on_enter:
22133 - set_context:
22134 review_entry: first
22135 transitions:
22136 - to: drafting
22137 guard:
22138 context:
22139 route:
22140 eq: drafting
22141"#;
22142 let agent = AgentBuilder::from_yaml(yaml)
22143 .unwrap()
22144 .llm(Arc::new(mock_with_responses(vec![
22145 "Intake response",
22146 "Draft response",
22147 "Review response",
22148 ])))
22149 .build()
22150 .unwrap();
22151
22152 agent
22153 .set_context("route", serde_json::json!("drafting"))
22154 .unwrap();
22155 agent.chat("Start a draft").await.unwrap();
22156 assert_eq!(agent.current_state().as_deref(), Some("drafting"));
22157 assert_eq!(
22158 agent.get_context().get("draft_version"),
22159 Some(&serde_json::json!(1))
22160 );
22161
22162 agent
22163 .set_context("route", serde_json::json!("review"))
22164 .unwrap();
22165 agent.chat("Review this").await.unwrap();
22166 assert_eq!(agent.current_state().as_deref(), Some("review"));
22167 assert_eq!(
22168 agent.get_context().get("review_entry"),
22169 Some(&serde_json::json!("first"))
22170 );
22171
22172 agent
22173 .set_context("route", serde_json::json!("drafting"))
22174 .unwrap();
22175 agent.chat("Revise this").await.unwrap();
22176 assert_eq!(agent.current_state().as_deref(), Some("drafting"));
22177 assert_eq!(
22178 agent.get_context().get("draft_version"),
22179 Some(&serde_json::json!(2))
22180 );
22181 }
22182
22183 #[tokio::test]
22184 async fn test_manual_transition_uses_on_enter_then_on_reenter() {
22185 let yaml = r#"
22186name: ManualLifecycleAgent
22187system_prompt: "You are helpful."
22188states:
22189 initial: intake
22190 states:
22191 intake:
22192 prompt: "Intake"
22193 drafting:
22194 prompt: "Drafting"
22195 on_enter:
22196 - set_context:
22197 draft_version: 1
22198 on_reenter:
22199 - set_context:
22200 draft_version: 2
22201 review:
22202 prompt: "Review"
22203"#;
22204 let agent = AgentBuilder::from_yaml(yaml)
22205 .unwrap()
22206 .llm(Arc::new(mock_with_response("unused")))
22207 .build()
22208 .unwrap();
22209
22210 assert!(!agent.get_context().contains_key("draft_version"));
22211 agent.transition_to("drafting").await.unwrap();
22212 assert_eq!(agent.current_state().as_deref(), Some("drafting"));
22213 assert_eq!(
22214 agent.get_context().get("draft_version"),
22215 Some(&serde_json::json!(1))
22216 );
22217
22218 agent.transition_to("review").await.unwrap();
22219 agent.transition_to("drafting").await.unwrap();
22220 assert_eq!(agent.current_state().as_deref(), Some("drafting"));
22221 assert_eq!(
22222 agent.get_context().get("draft_version"),
22223 Some(&serde_json::json!(2))
22224 );
22225 }
22226
22227 #[tokio::test]
22228 async fn test_timeout_transition_uses_on_enter_then_on_reenter() {
22229 let yaml = r#"
22230name: TimeoutLifecycleAgent
22231system_prompt: "You are helpful."
22232states:
22233 initial: intake
22234 regenerate_on_transition: false
22235 states:
22236 intake:
22237 prompt: "Intake"
22238 max_turns: 1
22239 timeout_to: drafting
22240 drafting:
22241 prompt: "Drafting"
22242 max_turns: 1
22243 timeout_to: review
22244 on_enter:
22245 - set_context:
22246 draft_version: 1
22247 on_reenter:
22248 - set_context:
22249 draft_version: 2
22250 review:
22251 prompt: "Review"
22252 max_turns: 1
22253 timeout_to: drafting
22254 on_enter:
22255 - set_context:
22256 review_entry: first
22257"#;
22258 let agent = AgentBuilder::from_yaml(yaml)
22259 .unwrap()
22260 .llm(Arc::new(mock_with_responses(vec![
22261 "Intake",
22262 "First draft",
22263 "Review",
22264 "Revised draft",
22265 ])))
22266 .build()
22267 .unwrap();
22268
22269 agent.chat("First turn").await.unwrap();
22270 assert_eq!(agent.current_state().as_deref(), Some("intake"));
22271 assert!(!agent.get_context().contains_key("draft_version"));
22272
22273 agent.chat("Second turn").await.unwrap();
22274 assert_eq!(agent.current_state().as_deref(), Some("drafting"));
22275 assert_eq!(
22276 agent.get_context().get("draft_version"),
22277 Some(&serde_json::json!(1))
22278 );
22279
22280 agent.chat("Third turn").await.unwrap();
22281 assert_eq!(agent.current_state().as_deref(), Some("review"));
22282 assert_eq!(
22283 agent.get_context().get("review_entry"),
22284 Some(&serde_json::json!("first"))
22285 );
22286
22287 agent.chat("Fourth turn").await.unwrap();
22288 assert_eq!(agent.current_state().as_deref(), Some("drafting"));
22289 assert_eq!(
22290 agent.get_context().get("draft_version"),
22291 Some(&serde_json::json!(2))
22292 );
22293 }
22294
22295 #[tokio::test]
22297 async fn test_integration_process_normalize() {
22298 let yaml = r#"
22299name: ProcessAgent
22300system_prompt: "You are helpful."
22301process:
22302 input:
22303 - type: normalize
22304 config:
22305 trim: true
22306 collapse_whitespace: true
22307"#;
22308 let mock = mock_with_response("Got your message.");
22309 let builder = AgentBuilder::from_yaml(yaml).unwrap();
22310 let agent = builder.llm(Arc::new(mock.clone())).build().unwrap();
22311
22312 let _ = agent.chat(" hello world ").await.unwrap();
22313
22314 let history = mock.call_history();
22316 assert!(!history.is_empty());
22317 let last_call = history.last().unwrap();
22319 let user_msg = last_call
22320 .messages
22321 .iter()
22322 .find(|m| m.role == ai_agents_core::Role::User)
22323 .unwrap();
22324 assert_eq!(user_msg.content, "hello world");
22325 }
22326
22327 #[tokio::test]
22331 async fn test_integration_memory_compression() {
22332 let yaml = r#"
22333name: MemoryAgent
22334system_prompt: "You are helpful."
22335memory:
22336 type: compacting
22337 max_messages: 100
22338 compress_threshold: 5
22339 max_recent_messages: 3
22340 summarize_batch_size: 2
22341"#;
22342 let responses: Vec<&str> = (0..8).map(|_| "Response from assistant.").collect();
22344 let mock = mock_with_responses(responses);
22345 let builder = AgentBuilder::from_yaml(yaml).unwrap();
22346 let agent = builder.llm(Arc::new(mock)).build().unwrap();
22347
22348 for i in 0..6 {
22350 let _ = agent.chat(&format!("Message {}", i)).await.unwrap();
22351 }
22352
22353 let messages = agent.memory.get_messages(None).await.unwrap();
22356 assert!(messages.len() <= 12); }
22360
22361 #[tokio::test]
22363 async fn test_integration_multi_llm_registry() {
22364 let mut mock_default = MockLLMProvider::new("default");
22365 mock_default.set_response("Default LLM response.");
22366 let mut mock_router = MockLLMProvider::new("router");
22367 mock_router.set_response("Router response.");
22368
22369 let agent = AgentBuilder::new()
22370 .system_prompt("You are helpful.")
22371 .llm_alias("default", Arc::new(mock_default))
22372 .llm_alias("router", Arc::new(mock_router))
22373 .build()
22374 .unwrap();
22375
22376 let response = agent.chat("Hello").await.unwrap();
22377 assert_eq!(response.content, "Default LLM response.");
22378 }
22379
22380 #[tokio::test]
22382 async fn test_integration_agent_reset() {
22383 let mock = mock_with_responses(vec!["Hello!", "Hello again!"]);
22384 let agent = AgentBuilder::new()
22385 .system_prompt("You are helpful.")
22386 .llm(Arc::new(mock))
22387 .build()
22388 .unwrap();
22389
22390 let _ = agent.chat("Hi").await.unwrap();
22391 let messages = agent.memory.get_messages(None).await.unwrap();
22392 assert_eq!(messages.len(), 2); agent.reset().await.unwrap();
22395 let messages = agent.memory.get_messages(None).await.unwrap();
22396 assert_eq!(messages.len(), 0);
22397 }
22398
22399 #[tokio::test]
22401 async fn test_integration_process_validate_reject() {
22402 use ai_agents_process::{ProcessConfig, ProcessProcessor};
22403
22404 let validate_config = ai_agents_process::ValidateStage {
22405 id: Some("length_check".to_string()),
22406 condition: None,
22407 config: ai_agents_process::ValidateConfig {
22408 rules: vec![ai_agents_process::ValidationRule::MinLength {
22409 min_length: 10,
22410 on_fail: ai_agents_process::ValidationAction {
22411 action: ai_agents_process::ValidationActionType::Reject,
22412 message: None,
22413 },
22414 }],
22415 ..Default::default()
22416 },
22417 };
22418 let process_config = ProcessConfig {
22419 input: vec![ai_agents_process::ProcessStage::Validate(validate_config)],
22420 ..Default::default()
22421 };
22422 let processor = ProcessProcessor::new(process_config);
22423
22424 let mock = mock_with_response("Should not reach here.");
22425 let agent = AgentBuilder::new()
22426 .system_prompt("You are helpful.")
22427 .llm(Arc::new(mock))
22428 .process_processor(processor)
22429 .build()
22430 .unwrap();
22431
22432 let response = agent.chat("Hi").await.unwrap();
22433 assert!(
22435 response.content.contains("rejected")
22436 || response.content.contains("Input rejected")
22437 || response.content.contains("too short")
22438 || response.content.contains("Too short")
22439 || response.content.len() < 50, "Expected rejection response, got: {}",
22441 response.content
22442 );
22443 }
22444
22445 #[tokio::test]
22447 async fn test_llm_fallback_on_failure() {
22448 use ai_agents_recovery::{ErrorRecoveryConfig, LLMFailureAction, LLMRecoveryConfig};
22449
22450 let mut primary = MockLLMProvider::new("primary");
22451 primary.set_error("Primary LLM is unavailable");
22452
22453 let mut fallback = MockLLMProvider::new("fallback");
22454 fallback.set_response("Fallback response works!");
22455
22456 let agent = AgentBuilder::new()
22457 .system_prompt("You are helpful.")
22458 .llm_alias("default", Arc::new(primary))
22459 .llm_alias("backup", Arc::new(fallback))
22460 .recovery_manager(RecoveryManager::new(ErrorRecoveryConfig {
22461 llm: LLMRecoveryConfig {
22462 on_failure: LLMFailureAction::FallbackLlm {
22463 fallback_llm: "backup".to_string(),
22464 },
22465 ..Default::default()
22466 },
22467 ..Default::default()
22468 }))
22469 .build()
22470 .unwrap();
22471
22472 let response = agent.chat("Hello").await.unwrap();
22473 assert!(
22474 response.content.contains("Fallback response"),
22475 "Expected fallback response, got: {}",
22476 response.content
22477 );
22478 }
22479
22480 #[tokio::test]
22482 async fn test_llm_fallback_response_static_message() {
22483 use ai_agents_recovery::{ErrorRecoveryConfig, LLMFailureAction, LLMRecoveryConfig};
22484
22485 let mut primary = MockLLMProvider::new("primary");
22486 primary.set_error("Primary LLM is unavailable");
22487
22488 let agent = AgentBuilder::new()
22489 .system_prompt("You are helpful.")
22490 .llm(Arc::new(primary))
22491 .recovery_manager(RecoveryManager::new(ErrorRecoveryConfig {
22492 llm: LLMRecoveryConfig {
22493 on_failure: LLMFailureAction::FallbackResponse {
22494 message: "I am temporarily unavailable. Please try again later."
22495 .to_string(),
22496 },
22497 ..Default::default()
22498 },
22499 ..Default::default()
22500 }))
22501 .build()
22502 .unwrap();
22503
22504 let response = agent.chat("Hello").await.unwrap();
22505 assert!(
22506 response.content.contains("temporarily unavailable"),
22507 "Expected static fallback message, got: {}",
22508 response.content
22509 );
22510 }
22511
22512 #[tokio::test]
22515 async fn test_tool_failure_skip() {
22516 use ai_agents_recovery::{
22517 ErrorRecoveryConfig, ToolFailureAction, ToolRecoveryConfig, ToolRetryConfig,
22518 };
22519
22520 let mock = mock_with_responses(vec![
22521 r#"{"tool": "calculator", "arguments": {"expression": "not a number +"}}"#,
22522 "The calculation was skipped, but I can still help you.",
22523 ]);
22524 let observed = mock.clone();
22525 let mut tools = ai_agents_tools::ToolRegistry::new();
22526 tools
22527 .register(Arc::new(ai_agents_tools::CalculatorTool))
22528 .unwrap();
22529
22530 let agent = AgentBuilder::new()
22531 .system_prompt("You are helpful.")
22532 .llm(Arc::new(mock))
22533 .tools(tools)
22534 .recovery_manager(RecoveryManager::new(ErrorRecoveryConfig {
22535 tools: ToolRecoveryConfig {
22536 default: ToolRetryConfig {
22537 max_retries: 0,
22538 timeout_ms: None,
22539 on_failure: ToolFailureAction::Skip,
22540 },
22541 ..Default::default()
22542 },
22543 ..Default::default()
22544 }))
22545 .build()
22546 .unwrap();
22547
22548 let response = agent.chat("Compute this").await.unwrap();
22549
22550 assert_eq!(
22551 response.content,
22552 "The calculation was skipped, but I can still help you."
22553 );
22554 assert_eq!(observed.call_count(), 2);
22555 let history = agent.tool_call_history();
22557 assert_eq!(history.len(), 1);
22558 assert_eq!(history[0].tool_id, "calculator");
22559 assert_eq!(
22560 history[0].result.get("skipped"),
22561 Some(&serde_json::json!(true)),
22562 "{:?}",
22563 history[0].result
22564 );
22565 }
22566
22567 #[tokio::test]
22569 async fn test_unregistered_tool_call_records_unavailable_and_continues() {
22570 let mock = mock_with_responses(vec![
22571 r#"{"tool": "nonexistent_tool", "arguments": {}}"#,
22572 "The tool was unavailable, but I can still help you.",
22573 ]);
22574 let observed = mock.clone();
22575
22576 let agent = AgentBuilder::new()
22577 .system_prompt("You are helpful.")
22578 .llm(Arc::new(mock))
22579 .build()
22580 .unwrap();
22581
22582 let response = agent.chat("Use the nonexistent tool").await.unwrap();
22583
22584 assert_eq!(
22585 response.content,
22586 "The tool was unavailable, but I can still help you."
22587 );
22588 assert_eq!(observed.call_count(), 2);
22589 let history = agent.tool_call_history();
22590 assert_eq!(history.len(), 1);
22591 assert_eq!(history[0].tool_id, "nonexistent_tool");
22592 assert_eq!(
22593 history[0].result.pointer("/error/kind"),
22594 Some(&serde_json::json!("tool_unavailable")),
22595 "{:?}",
22596 history[0].result
22597 );
22598 }
22599
22600 fn fallback_llm_recovery(fallback_llm: &str) -> RecoveryManager {
22605 use ai_agents_recovery::{ErrorRecoveryConfig, LLMFailureAction, LLMRecoveryConfig};
22606 RecoveryManager::new(ErrorRecoveryConfig {
22607 llm: LLMRecoveryConfig {
22608 on_failure: LLMFailureAction::FallbackLlm {
22609 fallback_llm: fallback_llm.to_string(),
22610 },
22611 ..Default::default()
22612 },
22613 ..Default::default()
22614 })
22615 }
22616
22617 #[tokio::test]
22618 async fn test_stream_llm_fallback_on_open_failure() {
22619 let mut primary = MockLLMProvider::new("primary");
22620 primary.set_error("Primary LLM is unavailable");
22621 let mut fallback = MockLLMProvider::new("fallback");
22622 fallback.set_response("Fallback response works!");
22623
22624 let agent = AgentBuilder::new()
22625 .system_prompt("You are helpful.")
22626 .llm_alias("default", Arc::new(primary))
22627 .llm_alias("backup", Arc::new(fallback))
22628 .recovery_manager(fallback_llm_recovery("backup"))
22629 .build()
22630 .unwrap();
22631
22632 let (content, chunks, final_response) = collect_stream_events(&agent, "Hello").await;
22633 assert!(
22634 !chunks.iter().any(StreamChunk::is_error),
22635 "fallback must not surface as a stream error: {chunks:?}"
22636 );
22637 let final_response = final_response.expect("Final must be emitted after fallback");
22638 assert!(content.contains("Fallback response"));
22639 assert!(final_response.content.contains("Fallback response"));
22640 }
22641
22642 #[tokio::test]
22643 async fn test_stream_llm_fallback_response_static_message() {
22644 use ai_agents_recovery::{ErrorRecoveryConfig, LLMFailureAction, LLMRecoveryConfig};
22645
22646 let mut primary = MockLLMProvider::new("primary");
22647 primary.set_error("Primary LLM is unavailable");
22648
22649 let agent = AgentBuilder::new()
22650 .system_prompt("You are helpful.")
22651 .llm(Arc::new(primary))
22652 .recovery_manager(RecoveryManager::new(ErrorRecoveryConfig {
22653 llm: LLMRecoveryConfig {
22654 on_failure: LLMFailureAction::FallbackResponse {
22655 message: "Service is temporarily unavailable.".to_string(),
22656 },
22657 ..Default::default()
22658 },
22659 ..Default::default()
22660 }))
22661 .build()
22662 .unwrap();
22663
22664 let (content, chunks, final_response) = collect_stream_events(&agent, "Hello").await;
22665 assert!(!chunks.iter().any(StreamChunk::is_error));
22666 let content_chunks = chunks.iter().filter(|c| c.is_content()).count();
22667 assert_eq!(content_chunks, 1, "static fallback is one content chunk");
22668 assert_eq!(content, "Service is temporarily unavailable.");
22669 assert_eq!(
22670 final_response.expect("Final").content,
22671 "Service is temporarily unavailable."
22672 );
22673 }
22674
22675 struct FailOnceStreamProvider {
22677 remaining_failures: Arc<std::sync::atomic::AtomicUsize>,
22678 open_attempts: Arc<std::sync::atomic::AtomicUsize>,
22679 }
22680
22681 #[async_trait]
22682 impl LLMProvider for FailOnceStreamProvider {
22683 async fn complete(
22684 &self,
22685 _messages: &[ChatMessage],
22686 _config: Option<&LLMConfig>,
22687 ) -> std::result::Result<LLMResponse, LLMError> {
22688 Ok(LLMResponse::new("blocking path", FinishReason::Stop))
22689 }
22690
22691 async fn complete_stream(
22692 &self,
22693 _messages: &[ChatMessage],
22694 _config: Option<&LLMConfig>,
22695 ) -> std::result::Result<
22696 Box<dyn Stream<Item = std::result::Result<LLMChunk, LLMError>> + Unpin + Send>,
22697 LLMError,
22698 > {
22699 self.open_attempts.fetch_add(1, Ordering::SeqCst);
22700 if self
22701 .remaining_failures
22702 .fetch_update(Ordering::SeqCst, Ordering::SeqCst, |n| n.checked_sub(1))
22703 .is_ok()
22704 {
22705 return Err(LLMError::Network("connection reset".to_string()));
22706 }
22707 Ok(Box::new(futures::stream::iter(vec![Ok(LLMChunk::new(
22708 "Recovered after retry",
22709 true,
22710 ))])))
22711 }
22712
22713 fn provider_name(&self) -> &str {
22714 "fail-once-stream"
22715 }
22716
22717 fn supports(&self, feature: LLMFeature) -> bool {
22718 matches!(feature, LLMFeature::Streaming)
22719 }
22720 }
22721
22722 #[tokio::test]
22723 async fn test_stream_llm_retry_then_success() {
22724 use ai_agents_recovery::{BackoffConfig, ErrorRecoveryConfig, RetryConfig};
22725
22726 let open_attempts = Arc::new(std::sync::atomic::AtomicUsize::new(0));
22727 let provider = FailOnceStreamProvider {
22728 remaining_failures: Arc::new(std::sync::atomic::AtomicUsize::new(1)),
22729 open_attempts: Arc::clone(&open_attempts),
22730 };
22731
22732 let agent = AgentBuilder::new()
22733 .system_prompt("You are helpful.")
22734 .llm(Arc::new(provider))
22735 .recovery_manager(RecoveryManager::new(ErrorRecoveryConfig {
22736 default: RetryConfig {
22737 max_retries: 1,
22738 backoff: BackoffConfig {
22739 initial_ms: 1,
22740 max_ms: 1,
22741 ..Default::default()
22742 },
22743 ..Default::default()
22744 },
22745 ..Default::default()
22746 }))
22747 .build()
22748 .unwrap();
22749
22750 let (content, chunks, final_response) = collect_stream_events(&agent, "Hello").await;
22751 assert!(!chunks.iter().any(StreamChunk::is_error), "{chunks:?}");
22752 assert_eq!(open_attempts.load(Ordering::SeqCst), 2);
22753 assert_eq!(content, "Recovered after retry");
22754 assert_eq!(
22755 final_response.expect("Final").content,
22756 "Recovered after retry"
22757 );
22758 }
22759
22760 #[tokio::test]
22761 async fn test_stream_llm_error_action_error_emits_terminal_error() {
22762 let mut primary = MockLLMProvider::new("primary");
22763 primary.set_error("Primary LLM is unavailable");
22764
22765 let agent = AgentBuilder::new()
22766 .system_prompt("You are helpful.")
22767 .llm(Arc::new(primary))
22768 .build()
22769 .unwrap();
22770
22771 let (_, chunks, final_response) = collect_stream_events(&agent, "Hello").await;
22772 assert!(
22773 final_response.is_none(),
22774 "default Error action must not produce Final"
22775 );
22776 assert!(
22777 chunks.iter().any(StreamChunk::is_error),
22778 "default Error action must surface a stream error"
22779 );
22780 }
22781
22782 struct MidStreamFailureProvider;
22784
22785 #[async_trait]
22786 impl LLMProvider for MidStreamFailureProvider {
22787 async fn complete(
22788 &self,
22789 _messages: &[ChatMessage],
22790 _config: Option<&LLMConfig>,
22791 ) -> std::result::Result<LLMResponse, LLMError> {
22792 Ok(LLMResponse::new("blocking path", FinishReason::Stop))
22793 }
22794
22795 async fn complete_stream(
22796 &self,
22797 _messages: &[ChatMessage],
22798 _config: Option<&LLMConfig>,
22799 ) -> std::result::Result<
22800 Box<dyn Stream<Item = std::result::Result<LLMChunk, LLMError>> + Unpin + Send>,
22801 LLMError,
22802 > {
22803 Ok(Box::new(futures::stream::iter(vec![
22804 Ok(LLMChunk::new("Partial ", false)),
22805 Err(LLMError::Network("connection dropped".to_string())),
22806 ])))
22807 }
22808
22809 fn provider_name(&self) -> &str {
22810 "mid-stream-failure"
22811 }
22812
22813 fn supports(&self, feature: LLMFeature) -> bool {
22814 matches!(feature, LLMFeature::Streaming)
22815 }
22816 }
22817
22818 #[tokio::test]
22819 async fn test_stream_mid_stream_failure_is_terminal() {
22820 let mut fallback = MockLLMProvider::new("fallback");
22821 fallback.set_response("Fallback must not run");
22822 let fallback_calls = fallback.clone();
22823
22824 let agent = AgentBuilder::new()
22825 .system_prompt("You are helpful.")
22826 .llm_alias("default", Arc::new(MidStreamFailureProvider))
22827 .llm_alias("backup", Arc::new(fallback))
22828 .recovery_manager(fallback_llm_recovery("backup"))
22829 .build()
22830 .unwrap();
22831
22832 let (content, chunks, final_response) = collect_stream_events(&agent, "Hello").await;
22833 assert_eq!(content, "Partial ");
22834 assert!(chunks.iter().any(StreamChunk::is_error));
22835 assert!(final_response.is_none());
22836 assert_eq!(
22837 fallback_calls.call_count(),
22838 0,
22839 "fallback must not run after a visible delta"
22840 );
22841 }
22842
22843 #[tokio::test]
22844 async fn test_buffered_streaming_draft_uses_fallback_llm() {
22845 use futures::StreamExt;
22846
22847 let mut primary = MockLLMProvider::new("primary");
22848 primary.set_error("Primary LLM is unavailable");
22849 let fallback = mock_with_response("fallback one two");
22850 let yaml = r#"
22851name: BufferedFallbackAgent
22852system_prompt: "You stream safely."
22853llm:
22854 default: default
22855streaming:
22856 enabled: true
22857 buffer_size: 8
22858runtime:
22859 optimization:
22860 enabled: true
22861 max_speculative_llm_calls_per_turn: 2
22862 speculative_state_transitions: true
22863 streaming_policy: buffer_until_routing_done
22864 max_parallel_runtime_tasks: 2
22865states:
22866 initial: triage
22867 states:
22868 triage:
22869 prompt: "Answer from triage."
22870 transitions:
22871 - to: billing
22872 guard:
22873 context:
22874 route:
22875 eq: billing
22876 timing: parallel
22877 billing:
22878 prompt: "Billing state."
22879"#;
22880 let agent = AgentBuilder::from_yaml(yaml)
22881 .unwrap()
22882 .llm_alias("default", Arc::new(primary))
22883 .llm_alias("backup", Arc::new(fallback))
22884 .recovery_manager(fallback_llm_recovery("backup"))
22885 .build()
22886 .unwrap();
22887
22888 let mut stream = agent.chat_stream("hello").await.unwrap();
22889 let mut content = String::new();
22890 let mut error = None;
22891 while let Some(chunk) = stream.next().await {
22892 match chunk {
22893 StreamChunk::Content { text } => content.push_str(&text),
22894 StreamChunk::Error { message } => error = Some(message),
22895 StreamChunk::Done {} => break,
22896 _ => {}
22897 }
22898 }
22899
22900 assert_eq!(error, None);
22901 assert_eq!(content, "fallback one two");
22902 }
22903
22904 #[tokio::test]
22905 async fn parity_llm_fallback_llm() {
22906 let build = || {
22907 let mut primary = MockLLMProvider::new("primary");
22908 primary.set_error("Primary LLM is unavailable");
22909 let mut fallback = MockLLMProvider::new("fallback");
22910 fallback.set_response("Fallback response works!");
22911 AgentBuilder::new()
22912 .system_prompt("You are helpful.")
22913 .llm_alias("default", Arc::new(primary))
22914 .llm_alias("backup", Arc::new(fallback))
22915 .recovery_manager(fallback_llm_recovery("backup"))
22916 .build()
22917 .unwrap()
22918 };
22919 assert_blocking_streaming_parity(build, "Hello").await;
22920 }
22921
22922 #[tokio::test]
22923 async fn parity_llm_fallback_response() {
22924 use ai_agents_recovery::{ErrorRecoveryConfig, LLMFailureAction, LLMRecoveryConfig};
22925 let build = || {
22926 let mut primary = MockLLMProvider::new("primary");
22927 primary.set_error("Primary LLM is unavailable");
22928 AgentBuilder::new()
22929 .system_prompt("You are helpful.")
22930 .llm(Arc::new(primary))
22931 .recovery_manager(RecoveryManager::new(ErrorRecoveryConfig {
22932 llm: LLMRecoveryConfig {
22933 on_failure: LLMFailureAction::FallbackResponse {
22934 message: "Service is temporarily unavailable.".to_string(),
22935 },
22936 ..Default::default()
22937 },
22938 ..Default::default()
22939 }))
22940 .build()
22941 .unwrap()
22942 };
22943 assert_blocking_streaming_parity(build, "Hello").await;
22944 }
22945
22946 #[tokio::test]
22947 async fn parity_basic_chat() {
22948 let build = || {
22949 AgentBuilder::new()
22950 .system_prompt("You are helpful.")
22951 .llm(Arc::new(mock_with_response("Plain answer")))
22952 .build()
22953 .unwrap()
22954 };
22955 assert_blocking_streaming_parity(build, "Hello").await;
22956 }
22957
22958 fn skills_with_parallel_transition_yaml(extra_optimization: &str, streaming: &str) -> String {
22964 format!(
22965 r#"
22966name: SkillsBesideTransitionAgent
22967system_prompt: "Use skills when they match."
22968llm:
22969 default: default
22970 router: router
22971observability:
22972 enabled: true
22973 export:
22974 write_raw_events: true
22975{streaming}
22976runtime:
22977 optimization:
22978 enabled: true
22979 speculative_state_transitions: true
22980{extra_optimization}
22981states:
22982 initial: triage
22983 states:
22984 triage:
22985 prompt: "Triage state."
22986 transitions:
22987 - to: billing
22988 guard:
22989 context:
22990 route:
22991 eq: billing
22992 timing: parallel
22993 billing:
22994 prompt: "Billing state."
22995skills:
22996 - id: helper
22997 description: "Answer helper requests"
22998 trigger: "User asks for helper"
22999 steps:
23000 - prompt: "Answer the helper request: {{{{ user_input }}}}"
23001 llm: skill
23002"#
23003 )
23004 }
23005
23006 struct RoleMocks {
23010 main: MockLLMProvider,
23011 router: MockLLMProvider,
23012 skill: MockLLMProvider,
23013 }
23014
23015 fn role_mocks(main: MockLLMProvider, router: MockLLMProvider) -> RoleMocks {
23016 RoleMocks {
23017 main,
23018 router,
23019 skill: mock_with_response("Skill step response"),
23020 }
23021 }
23022
23023 fn build_skills_beside_transition_agent(yaml: &str, mocks: RoleMocks) -> RuntimeAgent {
23024 AgentBuilder::from_yaml(yaml)
23025 .unwrap()
23026 .llm_alias("default", Arc::new(mocks.main))
23027 .llm_alias("router", Arc::new(mocks.router))
23028 .llm_alias("skill", Arc::new(mocks.skill))
23029 .build()
23030 .unwrap()
23031 }
23032
23033 fn branch_events_with_commit_behavior(agent: &RuntimeAgent, behavior: &str) -> usize {
23034 agent
23035 .observability()
23036 .unwrap()
23037 .raw_events()
23038 .iter()
23039 .filter(|event| event.dimensions.get("commit_behavior") == Some(&behavior.to_string()))
23040 .count()
23041 }
23042
23043 #[tokio::test]
23044 async fn test_speculative_transition_with_skills_and_no_skill_branch_routes_skill_serially() {
23045 let default_mock = mock_with_response("Draft response");
23046 let router_mock = mock_with_response("helper");
23047 let router_counter = router_mock.clone();
23048 let yaml = skills_with_parallel_transition_yaml(
23049 " max_speculative_llm_calls_per_turn: 2\n max_parallel_runtime_tasks: 2",
23050 "",
23051 );
23052 let agent =
23053 build_skills_beside_transition_agent(&yaml, role_mocks(default_mock, router_mock));
23054
23055 let response = agent.chat("please use helper").await.unwrap();
23056
23057 assert_eq!(
23058 response.metadata.as_ref().and_then(|m| m.get("skill_id")),
23059 Some(&serde_json::json!("helper")),
23060 "skill must route even without a skill branch: {response:?}"
23061 );
23062 assert_eq!(router_counter.call_count(), 1);
23063 assert!(branch_events_with_commit_behavior(&agent, "transition_decision") > 0);
23065 assert_eq!(
23066 branch_events_with_commit_behavior(&agent, "skill_selection"),
23067 0
23068 );
23069 }
23070
23071 #[tokio::test]
23072 async fn test_speculative_transition_with_skills_no_match_commits_draft() {
23073 let default_mock = mock_with_response("Draft response");
23074 let router_mock = mock_with_response("none");
23075 let router_counter = router_mock.clone();
23076 let yaml = skills_with_parallel_transition_yaml(
23077 " max_speculative_llm_calls_per_turn: 2\n max_parallel_runtime_tasks: 2",
23078 "",
23079 );
23080 let agent =
23081 build_skills_beside_transition_agent(&yaml, role_mocks(default_mock, router_mock));
23082
23083 let response = agent.chat("just chat").await.unwrap();
23084
23085 assert_eq!(response.content, "Draft response");
23086 assert!(
23087 response
23088 .metadata
23089 .as_ref()
23090 .is_none_or(|m| !m.contains_key("skill_id"))
23091 );
23092 assert_eq!(router_counter.call_count(), 1);
23093 assert!(branch_events_with_commit_behavior(&agent, "final_response") > 0);
23094 }
23095
23096 #[tokio::test]
23097 async fn test_speculative_transition_win_skips_serial_skill_selection() {
23098 let default_mock = mock_with_response("Billing answer");
23099 let router_mock = mock_with_response("none");
23100 let router_counter = router_mock.clone();
23101 let yaml = skills_with_parallel_transition_yaml(
23102 " max_speculative_llm_calls_per_turn: 2\n max_parallel_runtime_tasks: 2",
23103 "",
23104 );
23105 let agent =
23106 build_skills_beside_transition_agent(&yaml, role_mocks(default_mock, router_mock));
23107 agent
23108 .set_context("route", serde_json::json!("billing"))
23109 .unwrap();
23110
23111 let response = agent.chat("billing please").await.unwrap();
23112
23113 assert_eq!(agent.current_state().as_deref(), Some("billing"));
23114 assert_eq!(response.content, "Billing answer");
23115 assert_eq!(router_counter.call_count(), 1);
23117 }
23118
23119 #[tokio::test]
23120 async fn test_speculative_skill_capacity_exhausted_still_routes_skill_serially() {
23121 let default_mock = mock_with_response("Draft response");
23122 let router_mock = mock_with_response("helper");
23123 let router_counter = router_mock.clone();
23124 let yaml = skills_with_parallel_transition_yaml(
23126 " speculative_skill_routing: true\n max_speculative_llm_calls_per_turn: 2\n max_parallel_runtime_tasks: 2",
23127 "",
23128 );
23129 let agent =
23130 build_skills_beside_transition_agent(&yaml, role_mocks(default_mock, router_mock));
23131
23132 let response = agent.chat("please use helper").await.unwrap();
23133
23134 assert_eq!(
23135 response.metadata.as_ref().and_then(|m| m.get("skill_id")),
23136 Some(&serde_json::json!("helper"))
23137 );
23138 assert_eq!(router_counter.call_count(), 1);
23139 assert_eq!(
23140 branch_events_with_commit_behavior(&agent, "skill_selection"),
23141 0
23142 );
23143 }
23144
23145 #[tokio::test]
23146 async fn test_speculative_transition_and_skill_both_enabled_unchanged() {
23147 let default_mock = mock_with_response("Draft response");
23148 let router_mock = mock_with_response("helper");
23149 let router_counter = router_mock.clone();
23150 let yaml = skills_with_parallel_transition_yaml(
23151 " speculative_skill_routing: true\n max_speculative_llm_calls_per_turn: 3\n max_parallel_runtime_tasks: 3",
23152 "",
23153 );
23154 let agent =
23155 build_skills_beside_transition_agent(&yaml, role_mocks(default_mock, router_mock));
23156
23157 let response = agent.chat("please use helper").await.unwrap();
23158
23159 assert_eq!(
23160 response.metadata.as_ref().and_then(|m| m.get("skill_id")),
23161 Some(&serde_json::json!("helper"))
23162 );
23163 assert_eq!(router_counter.call_count(), 1);
23164 assert!(branch_events_with_commit_behavior(&agent, "skill_selection") > 0);
23166 }
23167
23168 const BUFFERED_STREAMING_YAML_FRAGMENT: &str = "streaming:\n enabled: true\n buffer_size: 16";
23169 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";
23170
23171 #[tokio::test]
23172 async fn test_buffered_streaming_skill_wins_after_transition_miss() {
23173 let mut default_mock = mock_with_response("draft one two");
23174 default_mock.set_latency(10);
23175 let router_mock = mock_with_response("helper");
23176 let router_counter = router_mock.clone();
23177 let yaml = skills_with_parallel_transition_yaml(
23178 BUFFERED_OPTIMIZATION_FRAGMENT,
23179 BUFFERED_STREAMING_YAML_FRAGMENT,
23180 );
23181 let agent =
23182 build_skills_beside_transition_agent(&yaml, role_mocks(default_mock, router_mock));
23183
23184 let (content, chunks, final_response) =
23185 collect_stream_events(&agent, "please use helper").await;
23186
23187 assert!(!chunks.iter().any(StreamChunk::is_error), "{chunks:?}");
23188 assert!(
23189 !content.contains("draft"),
23190 "buffered draft must be discarded when a skill wins: {content:?}"
23191 );
23192 let final_response = final_response.expect("Final");
23193 assert_eq!(
23194 final_response
23195 .metadata
23196 .as_ref()
23197 .and_then(|m| m.get("skill_id")),
23198 Some(&serde_json::json!("helper"))
23199 );
23200 assert_eq!(content, final_response.content);
23201 assert_eq!(router_counter.call_count(), 1);
23202 }
23203
23204 #[tokio::test]
23205 async fn test_buffered_streaming_skill_miss_releases_buffer_and_commits_draft() {
23206 let mut default_mock = mock_with_response("draft one two");
23207 default_mock.set_latency(10);
23208 let router_mock = mock_with_response("none");
23209 let router_counter = router_mock.clone();
23210 let yaml = skills_with_parallel_transition_yaml(
23211 BUFFERED_OPTIMIZATION_FRAGMENT,
23212 BUFFERED_STREAMING_YAML_FRAGMENT,
23213 );
23214 let agent =
23215 build_skills_beside_transition_agent(&yaml, role_mocks(default_mock, router_mock));
23216
23217 let (content, chunks, final_response) = collect_stream_events(&agent, "just chat").await;
23218
23219 assert!(!chunks.iter().any(StreamChunk::is_error), "{chunks:?}");
23220 assert_eq!(content, "draft one two");
23221 assert_eq!(final_response.expect("Final").content, "draft one two");
23222 assert_eq!(router_counter.call_count(), 1);
23223 }
23224
23225 #[tokio::test]
23226 async fn test_buffered_streaming_transition_win_skips_skill_selection() {
23227 let default_mock = mock_with_response("Billing answer");
23228 let router_mock = mock_with_response("none");
23229 let router_counter = router_mock.clone();
23230 let yaml = skills_with_parallel_transition_yaml(
23231 BUFFERED_OPTIMIZATION_FRAGMENT,
23232 BUFFERED_STREAMING_YAML_FRAGMENT,
23233 );
23234 let agent =
23235 build_skills_beside_transition_agent(&yaml, role_mocks(default_mock, router_mock));
23236 agent
23237 .set_context("route", serde_json::json!("billing"))
23238 .unwrap();
23239
23240 let (content, chunks, final_response) =
23241 collect_stream_events(&agent, "billing please").await;
23242
23243 assert!(!chunks.iter().any(StreamChunk::is_error), "{chunks:?}");
23244 assert_eq!(agent.current_state().as_deref(), Some("billing"));
23245 assert_eq!(content, "Billing answer");
23246 assert_eq!(final_response.expect("Final").content, "Billing answer");
23247 assert_eq!(router_counter.call_count(), 1);
23249 }
23250
23251 #[tokio::test]
23252 async fn parity_buffered_policy_with_skills() {
23253 let yaml = skills_with_parallel_transition_yaml(
23254 BUFFERED_OPTIMIZATION_FRAGMENT,
23255 BUFFERED_STREAMING_YAML_FRAGMENT,
23256 );
23257 let build = || {
23258 build_skills_beside_transition_agent(
23259 &yaml,
23260 role_mocks(
23261 mock_with_response("draft one two"),
23262 mock_with_response("helper"),
23263 ),
23264 )
23265 };
23266 let (blocking, _, _) = assert_blocking_streaming_parity(build, "please use helper").await;
23267 assert_eq!(
23268 blocking.metadata.as_ref().and_then(|m| m.get("skill_id")),
23269 Some(&serde_json::json!("helper"))
23270 );
23271 }
23272
23273 #[tokio::test]
23274 async fn parity_buffered_policy_with_cot() {
23275 let yaml = format!(
23276 r#"
23277name: BufferedCotAgent
23278system_prompt: "Think first."
23279llm:
23280 default: default
23281streaming:
23282 enabled: true
23283 buffer_size: 16
23284reasoning:
23285 mode: cot
23286runtime:
23287 optimization:
23288 enabled: true
23289 speculative_state_transitions: true
23290{BUFFERED_OPTIMIZATION_FRAGMENT}
23291states:
23292 initial: triage
23293 states:
23294 triage:
23295 prompt: "Triage state."
23296 transitions:
23297 - to: billing
23298 guard:
23299 context:
23300 route:
23301 eq: billing
23302 timing: parallel
23303 billing:
23304 prompt: "Billing state."
23305"#
23306 );
23307 let build = || {
23308 AgentBuilder::from_yaml(&yaml)
23309 .unwrap()
23310 .llm_alias(
23311 "default",
23312 Arc::new(mock_with_response(
23313 "<thinking>step by step</thinking>Reasoned answer",
23314 )),
23315 )
23316 .build()
23317 .unwrap()
23318 };
23319 let (blocking, streamed, _) = assert_blocking_streaming_parity(build, "hello").await;
23320 assert_eq!(blocking.content, "Reasoned answer");
23321 let mode = streamed
23322 .metadata
23323 .as_ref()
23324 .and_then(|m| m.get("reasoning"))
23325 .and_then(|r| r.get("mode_used"))
23326 .cloned();
23327 assert_eq!(
23329 mode,
23330 Some(serde_json::to_value(ReasoningMode::CoT).unwrap())
23331 );
23332 }
23333
23334 fn post_response_transition_yaml(states_extra: &str, billing_extra: &str) -> String {
23341 format!(
23342 r#"
23343name: PostResponseTransitionAgent
23344system_prompt: "You are helpful."
23345streaming:
23346 enabled: true
23347states:
23348 initial: intake
23349{states_extra}
23350 states:
23351 intake:
23352 prompt: "Intake"
23353 transitions:
23354 - to: billing
23355 guard:
23356 context:
23357 route:
23358 eq: billing
23359 billing:
23360 prompt: "Billing"
23361{billing_extra}
23362"#
23363 )
23364 }
23365
23366 fn build_post_response_transition_agent(yaml: &str, mock: MockLLMProvider) -> RuntimeAgent {
23367 let agent = AgentBuilder::from_yaml(yaml)
23368 .unwrap()
23369 .llm(Arc::new(mock))
23370 .build()
23371 .unwrap();
23372 agent
23373 .set_context("route", serde_json::json!("billing"))
23374 .unwrap();
23375 agent
23376 }
23377
23378 fn count_occurrences(haystack: &str, needle: &str) -> usize {
23379 haystack.matches(needle).count()
23380 }
23381
23382 #[tokio::test]
23383 async fn test_stream_transition_without_regeneration_emits_content_once() {
23384 let yaml = post_response_transition_yaml(" regenerate_on_transition: false", "");
23385 let agent =
23386 build_post_response_transition_agent(&yaml, mock_with_response("Intake answer"));
23387
23388 let (content, chunks, final_response) = collect_stream_events(&agent, "hello").await;
23389
23390 assert!(!chunks.iter().any(StreamChunk::is_error), "{chunks:?}");
23391 assert_eq!(agent.current_state().as_deref(), Some("billing"));
23392 assert_eq!(
23393 count_occurrences(&content, "Intake answer"),
23394 1,
23395 "committed content must not be emitted twice: {content:?}"
23396 );
23397 assert_eq!(final_response.expect("Final").content, content);
23398 assert!(
23399 chunks
23400 .iter()
23401 .any(|c| matches!(c, StreamChunk::StateTransition { .. }))
23402 );
23403 }
23404
23405 #[tokio::test]
23406 async fn test_stream_transition_without_regeneration_buffered_emits_content_once() {
23407 let yaml = post_response_transition_yaml(" regenerate_on_transition: false", "");
23408 let mut mock = mock_with_response("Intake answer");
23409 mock.set_tool_choice(Some(ToolChoice::Auto));
23411 let agent = build_post_response_transition_agent(&yaml, mock);
23412
23413 let (content, chunks, final_response) = collect_stream_events(&agent, "hello").await;
23414
23415 assert!(!chunks.iter().any(StreamChunk::is_error), "{chunks:?}");
23416 assert_eq!(agent.current_state().as_deref(), Some("billing"));
23417 assert_eq!(
23418 count_occurrences(&content, "Intake answer"),
23419 1,
23420 "{content:?}"
23421 );
23422 assert_eq!(final_response.expect("Final").content, content);
23423 }
23424
23425 #[tokio::test]
23426 async fn test_stream_state_regenerate_on_enter_false_emits_content_once() {
23427 let yaml = post_response_transition_yaml("", " regenerate_on_enter: false");
23428 let agent =
23429 build_post_response_transition_agent(&yaml, mock_with_response("Intake answer"));
23430
23431 let (content, chunks, final_response) = collect_stream_events(&agent, "hello").await;
23432
23433 assert!(!chunks.iter().any(StreamChunk::is_error), "{chunks:?}");
23434 assert_eq!(agent.current_state().as_deref(), Some("billing"));
23435 assert_eq!(
23436 count_occurrences(&content, "Intake answer"),
23437 1,
23438 "{content:?}"
23439 );
23440 assert_eq!(final_response.expect("Final").content, content);
23441 }
23442
23443 #[tokio::test]
23444 async fn test_stream_transition_with_regeneration_emits_replacement() {
23445 let yaml = post_response_transition_yaml("", "");
23446 let agent = build_post_response_transition_agent(
23447 &yaml,
23448 mock_with_responses(vec!["Intake answer", "Billing answer"]),
23449 );
23450
23451 let (content, chunks, final_response) = collect_stream_events(&agent, "hello").await;
23452
23453 assert!(!chunks.iter().any(StreamChunk::is_error), "{chunks:?}");
23454 assert_eq!(agent.current_state().as_deref(), Some("billing"));
23455 assert_eq!(count_occurrences(&content, "Intake answer"), 1);
23457 assert_eq!(count_occurrences(&content, "Billing answer"), 1);
23458 assert_eq!(final_response.expect("Final").content, "Billing answer");
23459 }
23460
23461 #[tokio::test]
23462 async fn test_blocking_transition_without_regeneration_unchanged() {
23463 let yaml = post_response_transition_yaml(" regenerate_on_transition: false", "");
23464 let agent =
23465 build_post_response_transition_agent(&yaml, mock_with_response("Intake answer"));
23466
23467 let response = agent.chat("hello").await.unwrap();
23468
23469 assert_eq!(response.content, "Intake answer");
23470 assert_eq!(agent.current_state().as_deref(), Some("billing"));
23471 }
23472
23473 #[tokio::test]
23474 async fn parity_transition_regenerate_off() {
23475 let yaml = post_response_transition_yaml(" regenerate_on_transition: false", "");
23476 let build =
23477 || build_post_response_transition_agent(&yaml, mock_with_response("Intake answer"));
23478 assert_blocking_streaming_parity(build, "hello").await;
23479 }
23480
23481 #[tokio::test]
23482 async fn parity_transition_regenerate_on() {
23483 let yaml = post_response_transition_yaml("", "");
23484 let build = || {
23485 build_post_response_transition_agent(
23486 &yaml,
23487 mock_with_responses(vec!["Intake answer", "Billing answer"]),
23488 )
23489 };
23490 let (blocking, _, _) = assert_blocking_streaming_parity(build, "hello").await;
23491 assert_eq!(blocking.content, "Billing answer");
23492 }
23493
23494 fn rejecting_process_processor() -> ProcessProcessor {
23499 use ai_agents_process::ProcessConfig;
23500 let validate_config = ai_agents_process::ValidateStage {
23501 id: Some("length_check".to_string()),
23502 condition: None,
23503 config: ai_agents_process::ValidateConfig {
23504 rules: vec![ai_agents_process::ValidationRule::MinLength {
23505 min_length: 10,
23506 on_fail: ai_agents_process::ValidationAction {
23507 action: ai_agents_process::ValidationActionType::Reject,
23508 message: None,
23509 },
23510 }],
23511 ..Default::default()
23512 },
23513 };
23514 ProcessProcessor::new(ProcessConfig {
23515 input: vec![ai_agents_process::ProcessStage::Validate(validate_config)],
23516 ..Default::default()
23517 })
23518 }
23519
23520 fn looks_like_rejection(content: &str) -> bool {
23522 content.contains("rejected")
23523 || content.contains("Input rejected")
23524 || content.contains("too short")
23525 || content.contains("Too short")
23526 || content.len() < 50
23527 }
23528
23529 #[tokio::test]
23530 async fn test_stream_input_rejection_is_final_response() {
23531 let responses = Arc::new(std::sync::atomic::AtomicUsize::new(0));
23532 let hooks = Arc::new(ResponseCountingHooks {
23533 responses: Arc::clone(&responses),
23534 });
23535 let mock = mock_with_response("Should not reach here.");
23536 let llm_calls = mock.clone();
23537 let agent = AgentBuilder::new()
23538 .system_prompt("You are helpful.")
23539 .llm(Arc::new(mock))
23540 .process_processor(rejecting_process_processor())
23541 .hooks(hooks.clone())
23542 .build()
23543 .unwrap();
23544
23545 let (content, chunks, final_response) = collect_stream_events(&agent, "Hi").await;
23546
23547 assert!(
23548 !chunks.iter().any(StreamChunk::is_error),
23549 "rejection is a response, not a stream error: {chunks:?}"
23550 );
23551 let final_response = final_response.expect("rejection must finalize as Final");
23552 assert!(
23553 looks_like_rejection(&final_response.content),
23554 "Expected rejection response, got: {}",
23555 final_response.content
23556 );
23557 assert_eq!(content, final_response.content);
23558 assert_eq!(
23559 llm_calls.call_count(),
23560 0,
23561 "rejected input must not reach the LLM"
23562 );
23563 assert_eq!(responses.load(Ordering::SeqCst), 1, "on_response must fire");
23564 }
23565
23566 #[tokio::test]
23567 async fn parity_input_rejection() {
23568 let build = || {
23569 AgentBuilder::new()
23570 .system_prompt("You are helpful.")
23571 .llm(Arc::new(mock_with_response("Should not reach here.")))
23572 .process_processor(rejecting_process_processor())
23573 .build()
23574 .unwrap()
23575 };
23576 let (blocking, _, _) = assert_blocking_streaming_parity(build, "Hi").await;
23577 assert!(
23578 looks_like_rejection(&blocking.content),
23579 "{}",
23580 blocking.content
23581 );
23582 }
23583
23584 fn pre_response_transition_yaml(streaming_policy: &str) -> String {
23585 format!(
23586 r#"
23587name: StreamingPreflightAgent
23588system_prompt: "You route before streaming."
23589runtime:
23590 optimization:
23591 enabled: true
23592 pre_response_deterministic_transitions: true
23593 streaming_policy: {streaming_policy}
23594streaming:
23595 enabled: true
23596 buffer_size: 16
23597states:
23598 initial: greeting
23599 states:
23600 greeting:
23601 prompt: "OLD_STATE_SENTINEL"
23602 transitions:
23603 - to: billing
23604 guard:
23605 context:
23606 topic:
23607 eq: billing
23608 timing: pre_response
23609 billing:
23610 prompt: "Billing state."
23611"#
23612 )
23613 }
23614
23615 fn build_pre_response_transition_agent(yaml: &str) -> RuntimeAgent {
23616 let agent = AgentBuilder::from_yaml(yaml)
23617 .unwrap()
23618 .llm(Arc::new(mock_with_response("Billing streamed response")))
23619 .build()
23620 .unwrap();
23621 agent
23622 .set_context("topic", serde_json::json!("billing"))
23623 .unwrap();
23624 agent
23625 }
23626
23627 #[tokio::test]
23628 async fn test_stream_buffered_policy_runs_pre_response_deterministic_transition() {
23629 let yaml = pre_response_transition_yaml("buffer_until_routing_done");
23630 let agent = build_pre_response_transition_agent(&yaml);
23631
23632 let (content, chunks, final_response) =
23633 collect_stream_events(&agent, "billing please").await;
23634
23635 assert!(!chunks.iter().any(StreamChunk::is_error), "{chunks:?}");
23636 assert_eq!(agent.current_state().as_deref(), Some("billing"));
23637 assert!(content.contains("Billing streamed response"));
23638 assert!(!content.contains("OLD_STATE_SENTINEL"));
23639 assert_eq!(final_response.expect("Final").content, content);
23640 }
23641
23642 #[tokio::test]
23643 async fn test_stream_disabled_policy_skips_preflight() {
23644 let yaml = pre_response_transition_yaml("disabled");
23645 let agent = build_pre_response_transition_agent(&yaml);
23646
23647 let (_, chunks, final_response) = collect_stream_events(&agent, "billing please").await;
23648
23649 assert!(!chunks.iter().any(StreamChunk::is_error), "{chunks:?}");
23650 assert!(final_response.is_some());
23651 assert_eq!(agent.current_state().as_deref(), Some("greeting"));
23655 }
23656
23657 #[tokio::test]
23658 async fn parity_pre_response_transition_buffered_policy() {
23659 let yaml = pre_response_transition_yaml("buffer_until_routing_done");
23660 let build = || build_pre_response_transition_agent(&yaml);
23661 assert_blocking_streaming_parity(build, "billing please").await;
23662 }
23663
23664 fn calculator_agent_with(mock: MockLLMProvider) -> RuntimeAgent {
23669 let mut tools = ai_agents_tools::ToolRegistry::new();
23670 tools
23671 .register(Arc::new(ai_agents_tools::CalculatorTool))
23672 .unwrap();
23673 AgentBuilder::new()
23674 .system_prompt("You are a calculator assistant.")
23675 .llm(Arc::new(mock))
23676 .tools(tools)
23677 .build()
23678 .unwrap()
23679 }
23680
23681 #[tokio::test]
23682 async fn test_stream_tool_start_events_precede_results_for_batch() {
23683 let mock = mock_with_responses(vec![
23684 r#"[{"tool": "calculator", "arguments": {"expression": "1+1"}}, {"tool": "calculator", "arguments": {"expression": "2+2"}}]"#,
23685 "Both answers are ready.",
23686 ]);
23687 let agent = calculator_agent_with(mock);
23688
23689 let (_, chunks, final_response) = collect_stream_events(&agent, "compute both").await;
23690
23691 assert!(!chunks.iter().any(StreamChunk::is_error), "{chunks:?}");
23692 let final_response = final_response.expect("Final");
23693 assert_eq!(final_response.tool_calls.as_ref().map(Vec::len), Some(2));
23694
23695 let tool_events: Vec<&StreamChunk> = chunks
23696 .iter()
23697 .filter(|c| {
23698 matches!(
23699 c,
23700 StreamChunk::ToolCallStart { .. }
23701 | StreamChunk::ToolResult { .. }
23702 | StreamChunk::ToolCallEnd { .. }
23703 )
23704 })
23705 .collect();
23706 assert_eq!(tool_events.len(), 6, "{tool_events:?}");
23707 assert!(matches!(tool_events[0], StreamChunk::ToolCallStart { .. }));
23709 assert!(matches!(tool_events[1], StreamChunk::ToolCallStart { .. }));
23710 assert!(matches!(
23711 tool_events[2],
23712 StreamChunk::ToolResult { success: true, .. }
23713 ));
23714 assert!(matches!(tool_events[3], StreamChunk::ToolCallEnd { .. }));
23715 assert!(matches!(
23716 tool_events[4],
23717 StreamChunk::ToolResult { success: true, .. }
23718 ));
23719 assert!(matches!(tool_events[5], StreamChunk::ToolCallEnd { .. }));
23720 }
23721
23722 #[tokio::test]
23723 async fn test_stream_clarification_final_carries_options_and_detection() {
23724 let responses = || {
23725 vec![
23726 r#"{"is_ambiguous":true,"confidence":0.2,"ambiguity_type":"missing_target","reasoning":"target missing","what_is_unclear":["target"],"detected_language":"en"}"#,
23727 r#"{"question":"What should I send?","options":["report","invoice"]}"#,
23728 ]
23729 };
23730 let (blocking_agent, _) = state_disambiguation_agent(responses(), true, None, true);
23731 let (streaming_agent, _) = state_disambiguation_agent(responses(), true, None, true);
23732
23733 let blocking = blocking_agent.chat("Send it").await.unwrap();
23734 let (_, chunks, streamed) = collect_stream_events(&streaming_agent, "Send it").await;
23735 assert!(!chunks.iter().any(StreamChunk::is_error), "{chunks:?}");
23736 let streamed = streamed.expect("clarification must finalize as Final");
23737
23738 assert_eq!(streamed.content, "What should I send?");
23739 let streamed_meta = streamed
23740 .metadata
23741 .as_ref()
23742 .and_then(|m| m.get("disambiguation"))
23743 .cloned()
23744 .expect("disambiguation metadata");
23745 for key in ["status", "options", "clarifying", "detection"] {
23746 assert!(
23747 streamed_meta.get(key).is_some(),
23748 "missing {key}: {streamed_meta}"
23749 );
23750 }
23751 assert_eq!(
23752 streamed_meta.get("detection").and_then(|d| d.get("type")),
23753 Some(&serde_json::json!("missing_target"))
23754 );
23755 assert_eq!(
23756 blocking
23757 .metadata
23758 .as_ref()
23759 .and_then(|m| m.get("disambiguation")),
23760 Some(&streamed_meta),
23761 "blocking and streaming clarification metadata must be identical"
23762 );
23763 }
23764
23765 struct FailingMemory {
23767 messages: parking_lot::RwLock<Vec<ChatMessage>>,
23768 fail_on_add: usize,
23769 adds: std::sync::atomic::AtomicUsize,
23770 }
23771
23772 #[async_trait]
23773 impl ai_agents_core::Memory for FailingMemory {
23774 async fn add_message(&self, message: ChatMessage) -> Result<()> {
23775 let n = self.adds.fetch_add(1, Ordering::SeqCst) + 1;
23776 if n == self.fail_on_add {
23777 return Err(AgentError::Other(format!(
23778 "simulated memory failure on add #{n}"
23779 )));
23780 }
23781 self.messages.write().push(message);
23782 Ok(())
23783 }
23784
23785 async fn get_messages(&self, limit: Option<usize>) -> Result<Vec<ChatMessage>> {
23786 let messages = self.messages.read();
23787 Ok(match limit {
23788 Some(n) if n < messages.len() => messages[messages.len() - n..].to_vec(),
23789 _ => messages.clone(),
23790 })
23791 }
23792
23793 async fn clear(&self) -> Result<()> {
23794 self.messages.write().clear();
23795 Ok(())
23796 }
23797
23798 fn len(&self) -> usize {
23799 self.messages.read().len()
23800 }
23801
23802 async fn restore(&self, snapshot: ai_agents_core::MemorySnapshot) -> Result<()> {
23803 *self.messages.write() = snapshot.messages;
23804 Ok(())
23805 }
23806 }
23807
23808 impl ai_agents_memory::Memory for FailingMemory {}
23809
23810 #[tokio::test]
23811 async fn test_stream_memory_write_failure_surfaces_as_error() {
23812 let yaml = r#"
23815name: TransitionOnToolCallAgent
23816system_prompt: "You are helpful."
23817streaming:
23818 enabled: true
23819states:
23820 initial: intake
23821 states:
23822 intake:
23823 prompt: "Intake"
23824 transitions:
23825 - to: billing
23826 guard:
23827 context:
23828 route:
23829 eq: billing
23830 billing:
23831 prompt: "Billing"
23832"#;
23833 let build = |fail_on_add: usize| {
23834 let mut tools = ai_agents_tools::ToolRegistry::new();
23835 tools
23836 .register(Arc::new(ai_agents_tools::CalculatorTool))
23837 .unwrap();
23838 let agent = AgentBuilder::from_yaml(yaml)
23839 .unwrap()
23840 .llm(Arc::new(mock_with_responses(vec![
23841 r#"{"tool": "calculator", "arguments": {"expression": "1+1"}}"#,
23842 "Billing answer",
23843 ])))
23844 .tools(tools)
23845 .memory(Arc::new(FailingMemory {
23846 messages: parking_lot::RwLock::new(Vec::new()),
23847 fail_on_add,
23848 adds: std::sync::atomic::AtomicUsize::new(0),
23849 }))
23850 .build()
23851 .unwrap();
23852 agent
23853 .set_context("route", serde_json::json!("billing"))
23854 .unwrap();
23855 agent
23856 };
23857
23858 let blocking = build(2).chat("compute").await;
23859 assert!(
23860 blocking.is_err(),
23861 "blocking must surface the memory failure"
23862 );
23863
23864 let (_, chunks, final_response) = collect_stream_events(&build(2), "compute").await;
23865 assert!(
23866 final_response.is_none(),
23867 "streaming must not finalize after a memory failure"
23868 );
23869 assert!(
23870 chunks.iter().any(|c| matches!(c, StreamChunk::Error { message } if message.contains("simulated memory failure"))),
23871 "streaming must surface the memory failure: {chunks:?}"
23872 );
23873
23874 assert!(build(usize::MAX).chat("compute").await.is_ok());
23876 }
23877
23878 #[tokio::test]
23879 async fn parity_tool_execution() {
23880 let build = || {
23881 calculator_agent_with(mock_with_responses(vec![
23882 r#"{"tool": "calculator", "arguments": {"expression": "2+2"}}"#,
23883 "The answer is 4.",
23884 ]))
23885 };
23886 let (blocking, _, chunks) = assert_blocking_streaming_parity(build, "What is 2+2?").await;
23887 assert_eq!(blocking.content, "The answer is 4.");
23888 assert!(
23889 chunks
23890 .iter()
23891 .any(|c| matches!(c, StreamChunk::ToolResult { .. }))
23892 );
23893 }
23894
23895 #[tokio::test]
23896 async fn parity_disambiguation_clarification() {
23897 let build = || {
23898 state_disambiguation_agent(
23899 vec![
23900 r#"{"is_ambiguous":true,"confidence":0.2,"ambiguity_type":"missing_target","reasoning":"target missing","what_is_unclear":["target"],"detected_language":"en"}"#,
23901 r#"{"question":"What should I send?","options":null}"#,
23902 ],
23903 true,
23904 None,
23905 true,
23906 )
23907 .0
23908 };
23909 let (blocking, _, _) = assert_blocking_streaming_parity(build, "Send it").await;
23910 assert_eq!(blocking.content, "What should I send?");
23911 }
23912
23913 #[tokio::test]
23914 async fn parity_reflection_enabled() {
23915 let yaml = r#"
23916name: ReflectionAgent
23917system_prompt: "You are careful."
23918reflection:
23919 enabled: true
23920 criteria:
23921 - "Is the answer helpful?"
23922"#;
23923 let build = || {
23924 AgentBuilder::from_yaml(yaml)
23925 .unwrap()
23926 .llm(Arc::new(mock_with_responses(vec![
23927 "Main answer",
23928 "OVERALL: PASS\nCONFIDENCE: 0.9",
23929 ])))
23930 .build()
23931 .unwrap()
23932 };
23933 let (blocking, streamed, _) = assert_blocking_streaming_parity(build, "hello").await;
23934 assert_eq!(blocking.content, "Main answer");
23935 assert!(metadata_keys(&streamed).contains("reflection"));
23936 }
23937
23938 #[tokio::test]
23939 async fn parity_cot_hidden_thinking() {
23940 let yaml = r#"
23941name: CotHiddenAgent
23942system_prompt: "Think first."
23943reasoning:
23944 mode: cot
23945 output: hidden
23946"#;
23947 let build = || {
23948 AgentBuilder::from_yaml(yaml)
23949 .unwrap()
23950 .llm(Arc::new(mock_with_response(
23951 "<thinking>step by step</thinking>Visible answer",
23952 )))
23953 .build()
23954 .unwrap()
23955 };
23956 let (blocking, streamed, chunks) = assert_blocking_streaming_parity(build, "hello").await;
23957 assert_eq!(blocking.content, "Visible answer");
23958 assert_eq!(content_chunks(&chunks).concat(), streamed.content);
23960 }
23961
23962 fn content_chunks(chunks: &[StreamChunk]) -> Vec<String> {
23967 chunks
23968 .iter()
23969 .filter_map(|c| match c {
23970 StreamChunk::Content { text } => Some(text.clone()),
23971 _ => None,
23972 })
23973 .collect()
23974 }
23975
23976 fn reflection_auto_agent(main: MockLLMProvider, judge: MockLLMProvider) -> RuntimeAgent {
23977 let yaml = r#"
23978name: ReflectionAutoAgent
23979system_prompt: "You are careful."
23980llm:
23981 default: default
23982 router: router
23983reflection:
23984 enabled: auto
23985 evaluator_llm: router
23986 criteria:
23987 - "Is the answer helpful?"
23988"#;
23989 AgentBuilder::from_yaml(yaml)
23990 .unwrap()
23991 .llm_alias("default", Arc::new(main))
23992 .llm_alias("router", Arc::new(judge))
23993 .build()
23994 .unwrap()
23995 }
23996
23997 #[tokio::test]
23998 async fn test_stream_reflection_auto_buffers_and_calls_judge_once_per_iteration() {
23999 let judge = mock_with_responses(vec!["YES", "OVERALL: PASS\nCONFIDENCE: 0.9"]);
24000 let judge_calls = judge.clone();
24001 let agent = reflection_auto_agent(mock_with_response("Main answer one two"), judge);
24002
24003 let (content, chunks, final_response) = collect_stream_events(&agent, "hello").await;
24004
24005 assert!(!chunks.iter().any(StreamChunk::is_error), "{chunks:?}");
24006 assert_eq!(
24007 content_chunks(&chunks).len(),
24008 1,
24009 "auto reflection must buffer the main response: {chunks:?}"
24010 );
24011 assert_eq!(content, "Main answer one two");
24012 assert_eq!(final_response.expect("Final").content, content);
24013 assert_eq!(judge_calls.call_count(), 2);
24015 }
24016
24017 #[tokio::test]
24018 async fn test_stream_reflection_auto_rewrite_is_streamed() {
24019 let judge = mock_with_responses(vec![
24020 "YES",
24021 "OVERALL: FAIL\nCONFIDENCE: 0.1",
24022 "OVERALL: PASS\nCONFIDENCE: 0.9",
24023 ]);
24024 let agent = reflection_auto_agent(
24025 mock_with_responses(vec!["First attempt", "Improved answer"]),
24026 judge,
24027 );
24028
24029 let (content, chunks, final_response) = collect_stream_events(&agent, "hello").await;
24030
24031 assert!(!chunks.iter().any(StreamChunk::is_error), "{chunks:?}");
24032 assert_eq!(content_chunks(&chunks).len(), 1, "{chunks:?}");
24033 assert_eq!(
24034 content, "Improved answer",
24035 "the rewritten answer is what streams"
24036 );
24037 assert_eq!(final_response.expect("Final").content, "Improved answer");
24038 }
24039
24040 fn reasoning_agent(mode: &str, output: &str) -> RuntimeAgent {
24041 let yaml = format!(
24042 r#"
24043name: ReasoningStreamAgent
24044system_prompt: "Think first."
24045reasoning:
24046 mode: {mode}
24047 output: {output}
24048"#
24049 );
24050 AgentBuilder::from_yaml(&yaml)
24051 .unwrap()
24052 .llm(Arc::new(mock_with_response(
24053 "<thinking>step by step</thinking>Visible answer",
24054 )))
24055 .build()
24056 .unwrap()
24057 }
24058
24059 #[tokio::test]
24060 async fn test_stream_cot_hidden_emits_no_thinking_tags() {
24061 let agent = reasoning_agent("cot", "hidden");
24062 let (content, chunks, final_response) = collect_stream_events(&agent, "hello").await;
24063 assert!(!chunks.iter().any(StreamChunk::is_error), "{chunks:?}");
24064 assert_eq!(content_chunks(&chunks).len(), 1, "{chunks:?}");
24065 assert!(!content.contains("<thinking>"), "{content:?}");
24066 assert_eq!(content, "Visible answer");
24067 assert_eq!(final_response.expect("Final").content, content);
24068 }
24069
24070 #[tokio::test]
24071 async fn test_stream_cot_visible_matches_final_format() {
24072 let agent = reasoning_agent("cot", "visible");
24073 let (content, chunks, final_response) = collect_stream_events(&agent, "hello").await;
24074 assert!(!chunks.iter().any(StreamChunk::is_error), "{chunks:?}");
24075 assert!(content.starts_with("Thinking:"), "{content:?}");
24076 assert!(content.contains("Answer:\nVisible answer"), "{content:?}");
24077 assert_eq!(final_response.expect("Final").content, content);
24078 }
24079
24080 #[tokio::test]
24081 async fn test_stream_react_buffers() {
24082 let agent = reasoning_agent("react", "hidden");
24083 let (content, chunks, final_response) = collect_stream_events(&agent, "hello").await;
24084 assert!(!chunks.iter().any(StreamChunk::is_error), "{chunks:?}");
24085 assert_eq!(content_chunks(&chunks).len(), 1, "{chunks:?}");
24086 assert_eq!(content, "Visible answer");
24087 assert_eq!(final_response.expect("Final").content, content);
24088 }
24089
24090 #[tokio::test]
24091 async fn test_stream_plain_mode_still_streams_deltas() {
24092 let agent = AgentBuilder::new()
24093 .system_prompt("You are helpful.")
24094 .llm(Arc::new(mock_with_response("one two three")))
24095 .build()
24096 .unwrap();
24097 let (content, chunks, final_response) = collect_stream_events(&agent, "hello").await;
24098 assert!(
24099 content_chunks(&chunks).len() >= 2,
24100 "plain turns must keep token-level streaming: {chunks:?}"
24101 );
24102 assert_eq!(content, "one two three");
24103 assert_eq!(final_response.expect("Final").content, content);
24104 }
24105
24106 struct ActorProbeHooks {
24108 seen: parking_lot::Mutex<Option<crate::TurnActorContext>>,
24109 }
24110
24111 #[async_trait]
24112 impl AgentHooks for ActorProbeHooks {
24113 async fn on_message_received(&self, _input: &str) {
24114 *self.seen.lock() = current_turn_actor_context();
24115 }
24116 }
24117
24118 fn actor_probe_agent(hooks: Arc<ActorProbeHooks>) -> RuntimeAgent {
24119 let yaml = r#"
24120name: ActorStreamAgent
24121system_prompt: "You are helpful."
24122observability:
24123 enabled: true
24124 export:
24125 write_raw_events: true
24126"#;
24127 AgentBuilder::from_yaml(yaml)
24128 .unwrap()
24129 .llm(Arc::new(mock_with_response("Hello actor")))
24130 .hooks(hooks)
24131 .build()
24132 .unwrap()
24133 }
24134
24135 async fn collect_actor_stream_final(
24136 agent: &RuntimeAgent,
24137 input: &str,
24138 actor_context: crate::TurnActorContext,
24139 ) -> AgentResponse {
24140 use futures::StreamExt;
24141 let mut events = agent
24142 .chat_stream_events_with_actor_context(input, actor_context)
24143 .await
24144 .expect("stream opens");
24145 let mut final_response = None;
24146 while let Some(event) = events.next().await {
24147 match event {
24148 AgentStreamEvent::Final(response) => final_response = Some(response),
24149 AgentStreamEvent::Chunk(StreamChunk::Error { message }) => {
24150 panic!("unexpected stream error: {message}")
24151 }
24152 AgentStreamEvent::Chunk(_) => {}
24153 }
24154 }
24155 final_response.expect("Final")
24156 }
24157
24158 #[tokio::test]
24159 async fn test_stream_events_with_actor_context_scopes_actor_for_turn() {
24160 let hooks = Arc::new(ActorProbeHooks {
24161 seen: parking_lot::Mutex::new(None),
24162 });
24163 let agent = actor_probe_agent(Arc::clone(&hooks));
24164 let actor_context = crate::TurnActorContext::new().with_origin_actor("customer_42");
24165
24166 let final_response = collect_actor_stream_final(&agent, "hi", actor_context).await;
24167
24168 assert_eq!(final_response.content, "Hello actor");
24169 assert_eq!(
24170 hooks
24171 .seen
24172 .lock()
24173 .as_ref()
24174 .and_then(|context| context.effective_actor_id().map(str::to_string)),
24175 Some("customer_42".to_string()),
24176 "the actor context must be visible inside the streaming turn"
24177 );
24178 assert!(
24179 agent.actor_id().is_none(),
24180 "a turn-scoped actor must not mutate the global actor ID"
24181 );
24182 let events = agent.observability().unwrap().raw_events();
24183 assert!(
24184 events
24185 .iter()
24186 .any(|event| event.dimensions.get("actor") == Some(&"customer_42".to_string())),
24187 "observation events must carry the actor dimension"
24188 );
24189 }
24190
24191 #[tokio::test]
24192 async fn test_stream_events_with_actor_context_matches_blocking_actor_context() {
24193 let actor_context = crate::TurnActorContext::new()
24194 .with_origin_actor("customer_42")
24195 .with_sender_agent("coordinator");
24196
24197 let blocking_hooks = Arc::new(ActorProbeHooks {
24198 seen: parking_lot::Mutex::new(None),
24199 });
24200 let blocking_agent = actor_probe_agent(Arc::clone(&blocking_hooks));
24201 let blocking = blocking_agent
24202 .chat_with_actor_context("hi", actor_context.clone())
24203 .await
24204 .unwrap();
24205
24206 let streaming_hooks = Arc::new(ActorProbeHooks {
24207 seen: parking_lot::Mutex::new(None),
24208 });
24209 let streaming_agent = actor_probe_agent(Arc::clone(&streaming_hooks));
24210 let streamed =
24211 collect_actor_stream_final(&streaming_agent, "hi", actor_context.clone()).await;
24212
24213 assert_eq!(blocking.content, streamed.content);
24214 assert_eq!(metadata_keys(&blocking), metadata_keys(&streamed));
24215 assert_eq!(
24216 *blocking_hooks.seen.lock(),
24217 *streaming_hooks.seen.lock(),
24218 "both entry points must expose the same turn actor context"
24219 );
24220 assert_eq!(*streaming_hooks.seen.lock(), Some(actor_context));
24221 }
24222
24223 #[tokio::test]
24224 async fn test_stream_events_with_actor_context_releases_root_turn_on_drop() {
24225 use futures::StreamExt;
24226 let agent = AgentBuilder::new()
24227 .system_prompt("You are helpful.")
24228 .llm(Arc::new(mock_with_response("one two three")))
24229 .build()
24230 .unwrap();
24231 {
24232 let mut events = agent
24233 .chat_stream_events_with_actor_context(
24234 "hi",
24235 crate::TurnActorContext::new().with_origin_actor("customer_42"),
24236 )
24237 .await
24238 .unwrap();
24239 let _first = events.next().await;
24241 }
24242 let next = tokio::time::timeout(Duration::from_secs(5), agent.chat("next")).await;
24243 assert!(
24244 matches!(next, Ok(Ok(_))),
24245 "the root turn must be released when the actor stream is dropped: {next:?}"
24246 );
24247 }
24248
24249 #[tokio::test]
24250 async fn parity_skill_route() {
24251 let yaml = skills_with_parallel_transition_yaml(
24252 " max_speculative_llm_calls_per_turn: 2\n max_parallel_runtime_tasks: 2",
24253 "",
24254 );
24255 let build = || {
24256 build_skills_beside_transition_agent(
24257 &yaml,
24258 role_mocks(
24259 mock_with_response("Draft response"),
24260 mock_with_response("helper"),
24261 ),
24262 )
24263 };
24264 assert_blocking_streaming_parity(build, "please use helper").await;
24265 }
24266}