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> {
3848 let context_config = self
3849 .disambiguation_manager
3850 .as_ref()
3851 .map(|manager| manager.config().context.clone())
3852 .unwrap_or_default();
3853 let recent_messages = if context_config.recent_messages == 0 {
3854 Vec::new()
3855 } else {
3856 Self::readable_native_messages(
3857 self.memory
3858 .get_messages(Some(context_config.recent_messages))
3859 .await?,
3860 )?
3861 .iter()
3862 .map(|message| format!("{:?}: {}", message.role, message.content))
3863 .collect()
3864 };
3865
3866 let current_state = self.current_state().map(|s| s.to_string());
3867
3868 let state_prompt: Option<String> = self
3871 .state_machine
3872 .as_ref()
3873 .and_then(|sm| sm.current_definition())
3874 .and_then(|def| def.prompt.clone());
3875
3876 let available_tools = if context_config.include_available_tools {
3877 self.get_available_tool_ids().await?
3878 } else {
3879 Vec::new()
3880 };
3881
3882 let available_skills: Vec<String> = self.skills.iter().map(|s| s.id.clone()).collect();
3883
3884 let mut user_context = self.build_context_with_overlays();
3885 user_context.remove(DISAMBIGUATION_STATE_GENERATION_KEY);
3886 if let Some(state_generation) = self
3887 .state_machine
3888 .as_ref()
3889 .map(|state_machine| state_machine.generation())
3890 {
3891 user_context.insert(
3892 DISAMBIGUATION_STATE_GENERATION_KEY.to_string(),
3893 serde_json::json!(state_generation),
3894 );
3895 }
3896
3897 let available_intents: Vec<String> = if let Some(ref sm) = self.state_machine {
3899 sm.current_definition()
3900 .map(|def| {
3901 def.transitions
3902 .iter()
3903 .filter_map(|t| t.intent.clone())
3904 .collect()
3905 })
3906 .unwrap_or_default()
3907 } else {
3908 Vec::new()
3909 };
3910
3911 Ok(DisambiguationContext::from_agent_state(
3912 recent_messages,
3913 current_state,
3914 state_prompt,
3915 available_tools,
3916 available_skills,
3917 available_intents,
3918 user_context,
3919 ))
3920 }
3921
3922 fn get_available_skills(&self) -> Vec<&SkillDefinition> {
3923 if let Some(ref sm) = self.state_machine
3924 && let Some(state_def) = sm.current_definition()
3925 {
3926 let parent_def = sm.get_parent_definition();
3927 let effective_skills = state_def.get_effective_skills(parent_def.as_ref());
3928 if !effective_skills.is_empty() {
3929 return self
3930 .skills
3931 .iter()
3932 .filter(|s| effective_skills.contains(&&s.id))
3933 .collect();
3934 }
3935 }
3936 self.skills.iter().collect()
3937 }
3938
3939 async fn build_messages(&self) -> Result<Vec<ChatMessage>> {
3940 self.build_messages_internal(true, None, true).await
3941 }
3942
3943 async fn build_messages_for_draft(&self, user_message: &str) -> Result<Vec<ChatMessage>> {
3944 self.build_messages_internal(false, Some(user_message), true)
3945 .await
3946 }
3947
3948 async fn build_messages_internal(
3949 &self,
3950 fire_persona_hooks: bool,
3951 ephemeral_user_message: Option<&str>,
3952 include_tool_prompt: bool,
3953 ) -> Result<Vec<ChatMessage>> {
3954 let system_prompt = self
3955 .get_effective_system_prompt_with_persona_hooks(fire_persona_hooks, include_tool_prompt)
3956 .await?;
3957 let mut messages = vec![ChatMessage::system(&system_prompt)];
3958
3959 let context = self.memory.get_context().await?;
3960 let history = if let Some(ref budget) = self.memory_token_budget {
3961 context.to_llm_messages_with_allocation(&budget.allocation)
3962 } else {
3963 context.to_llm_messages()
3964 };
3965 messages.extend(history);
3966 if let Some(user_message) = ephemeral_user_message {
3967 messages.push(ChatMessage::user(user_message));
3968 }
3969
3970 let total_tokens = self.estimate_total_tokens(&messages);
3971
3972 if total_tokens > self.max_context_tokens {
3973 debug!(
3974 total = total_tokens,
3975 limit = self.max_context_tokens,
3976 "Context overflow"
3977 );
3978
3979 match &self.recovery_manager.config().llm.on_context_overflow {
3980 ContextOverflowAction::Error => {
3981 return Err(AgentError::LLM(format!(
3982 "Context overflow: {} tokens > {} limit",
3983 total_tokens, self.max_context_tokens
3984 )));
3985 }
3986 ContextOverflowAction::Truncate { keep_recent } => {
3987 self.truncate_context(&mut messages, *keep_recent)?;
3988 }
3989 ContextOverflowAction::Summarize {
3990 summarizer_llm,
3991 max_summary_tokens,
3992 custom_prompt,
3993 keep_recent,
3994 filter,
3995 } => {
3996 self.summarize_context(
3997 &mut messages,
3998 summarizer_llm.as_deref(),
3999 *max_summary_tokens,
4000 custom_prompt.as_deref(),
4001 *keep_recent,
4002 filter.as_ref(),
4003 )
4004 .await?;
4005 }
4006 }
4007 }
4008
4009 self.validate_active_native_history(&messages, true)?;
4010 Ok(messages)
4011 }
4012
4013 async fn main_tool_protocol(
4014 &self,
4015 llm: &dyn LLMProvider,
4016 ephemeral_new_turn: bool,
4017 ) -> Result<MainToolProtocol> {
4018 let mut choice = llm.configured_tool_choice();
4019 if matches!(choice.as_ref(), Some(ToolChoice::None)) {
4020 return Ok(MainToolProtocol {
4021 choice,
4022 tool_ids: Vec::new(),
4023 definitions: Vec::new(),
4024 });
4025 }
4026
4027 let mut tool_ids = self.get_available_tool_ids().await?;
4028 tool_ids.sort();
4029 tool_ids.dedup();
4030 if let Some(ToolChoice::Specific(expected)) = choice.as_ref() {
4031 let canonical = self.tools.canonical_id(expected).ok_or_else(|| {
4032 AgentError::Config(format!(
4033 "specific tool choice '{expected}' is not registered"
4034 ))
4035 })?;
4036 if canonical != *expected {
4037 return Err(AgentError::Config(format!(
4038 "specific tool choice must use canonical ID '{canonical}', not '{expected}'"
4039 )));
4040 }
4041 if !tool_ids.iter().any(|tool_id| tool_id == expected) {
4042 return Err(AgentError::Config(format!(
4043 "specific tool choice '{expected}' is outside the effective tool grant"
4044 )));
4045 }
4046 }
4047 if matches!(
4048 choice.as_ref(),
4049 Some(ToolChoice::Required | ToolChoice::Specific(_))
4050 ) && tool_ids.is_empty()
4051 {
4052 return Err(AgentError::Config(
4053 "required tool choice has no tool inside the effective grant".to_string(),
4054 ));
4055 }
4056 if !ephemeral_new_turn
4057 && let Some(configured_choice) = choice.as_ref()
4058 && matches!(
4059 configured_choice,
4060 ToolChoice::Required | ToolChoice::Specific(_)
4061 )
4062 && self
4063 .tool_choice_satisfied_in_current_turn(configured_choice, &tool_ids)
4064 .await?
4065 {
4066 choice = Some(ToolChoice::Auto);
4067 }
4068 if let Some(ToolChoice::Specific(expected)) = choice.as_ref() {
4069 tool_ids.retain(|tool_id| tool_id == expected);
4070 }
4071
4072 let definitions = tool_ids
4073 .iter()
4074 .map(|tool_id| {
4075 let tool = self.tools.get(tool_id).ok_or_else(|| {
4076 AgentError::Config(format!(
4077 "effective tool '{tool_id}' disappeared before provider exposure"
4078 ))
4079 })?;
4080 Ok(LLMToolDefinition {
4081 name: tool_id.clone(),
4082 description: tool.description().to_string(),
4083 input_schema: tool.input_schema(),
4084 })
4085 })
4086 .collect::<Result<Vec<_>>>()?;
4087
4088 Ok(MainToolProtocol {
4092 choice,
4093 tool_ids,
4094 definitions,
4095 })
4096 }
4097
4098 async fn tool_choice_satisfied_in_current_turn(
4099 &self,
4100 choice: &ToolChoice,
4101 effective_tool_ids: &[String],
4102 ) -> Result<bool> {
4103 let messages = self.memory.get_messages(None).await?;
4104 let mut saw_tool_result = false;
4105 for message in messages.iter().rev() {
4106 match message.role {
4107 ai_agents_core::Role::Tool | ai_agents_core::Role::Function => {
4108 saw_tool_result = true;
4109 }
4110 ai_agents_core::Role::Assistant if saw_tool_result => {
4111 let Some(calls) = self.parse_tool_calls(&message.content)? else {
4112 continue;
4113 };
4114 let calls_are_effective = !calls.is_empty()
4115 && calls.iter().all(|call| {
4116 self.tools
4117 .canonical_id(&call.name)
4118 .is_some_and(|canonical| effective_tool_ids.contains(&canonical))
4119 });
4120 return Ok(calls_are_effective
4121 && match choice {
4122 ToolChoice::Required => true,
4123 ToolChoice::Specific(expected) => calls.iter().all(|call| {
4124 self.tools.canonical_id(&call.name).as_deref()
4125 == Some(expected.as_str())
4126 }),
4127 _ => false,
4128 });
4129 }
4130 ai_agents_core::Role::User => return Ok(false),
4131 _ => {}
4132 }
4133 }
4134 Ok(false)
4135 }
4136
4137 fn provider_can_use_native_tools(
4138 &self,
4139 llm: &dyn LLMProvider,
4140 protocol: &MainToolProtocol,
4141 ) -> bool {
4142 let Some(choice) = protocol.choice.as_ref() else {
4143 return false;
4144 };
4145 if matches!(choice, ToolChoice::None) || protocol.definitions.is_empty() {
4146 return false;
4147 }
4148 llm.supports_tool_choice(choice)
4149 && protocol.definitions.iter().all(|definition| {
4150 !definition.name.is_empty()
4151 && definition.name.len() <= 64
4152 && definition
4153 .name
4154 .bytes()
4155 .all(|byte| byte.is_ascii_alphanumeric() || matches!(byte, b'_' | b'-'))
4156 })
4157 }
4158
4159 fn prompt_messages_for_tool_protocol(
4160 &self,
4161 messages: &[ChatMessage],
4162 protocol: &MainToolProtocol,
4163 corrective: bool,
4164 ) -> Vec<ChatMessage> {
4165 let mut messages = messages.to_vec();
4166 let Some(choice) = protocol.choice.as_ref() else {
4167 return messages;
4168 };
4169 if matches!(choice, ToolChoice::None) || protocol.tool_ids.is_empty() {
4170 return messages;
4171 }
4172
4173 let mut tool_prompt = self.tools.generate_scoped_prompt_with_mode(
4174 &protocol.tool_ids,
4175 None,
4176 self.parallel_tools.enabled,
4177 self.runtime_config.tool_schema_prompt_mode,
4178 );
4179 match choice {
4180 ToolChoice::Required => tool_prompt.push_str(
4181 "\n\nYou must call at least one listed tool before giving a final answer.",
4182 ),
4183 ToolChoice::Specific(tool_id) => tool_prompt.push_str(&format!(
4184 "\n\nYou must call the '{tool_id}' tool before giving a final answer."
4185 )),
4186 ToolChoice::Auto => {}
4187 ToolChoice::None => return messages,
4188 _ => return messages,
4189 }
4190 if let Some(system) = messages
4191 .iter_mut()
4192 .find(|message| message.role == ai_agents_core::Role::System)
4193 {
4194 system.content.push_str("\n\n");
4195 system.content.push_str(&tool_prompt);
4196 } else {
4197 messages.insert(0, ChatMessage::system(tool_prompt));
4198 }
4199 if corrective {
4200 let instruction = match choice {
4201 ToolChoice::Required => {
4202 "Your previous response did not call a required tool. Call at least one listed tool now and return only the JSON tool call."
4203 }
4204 ToolChoice::Specific(tool_id) => {
4205 messages.push(ChatMessage::user(format!(
4206 "Your previous response did not call the required '{tool_id}' tool. Call it now and return only the JSON tool call."
4207 )));
4208 return messages;
4209 }
4210 _ => return messages,
4211 };
4212 messages.push(ChatMessage::user(instruction));
4213 }
4214 messages
4215 }
4216
4217 async fn invoke_main_provider(
4218 &self,
4219 llm: Arc<dyn LLMProvider>,
4220 messages: &[ChatMessage],
4221 protocol: &MainToolProtocol,
4222 corrective: bool,
4223 ) -> std::result::Result<MainProviderResponse, LLMError> {
4224 let use_native = self.provider_can_use_native_tools(llm.as_ref(), protocol);
4225 let response = if use_native {
4226 let request = LLMToolRequest {
4227 tools: protocol.definitions.clone(),
4228 choice: protocol
4229 .choice
4230 .clone()
4231 .expect("native tool requests require an explicit choice"),
4232 };
4233 self.observe_purpose(
4234 ObservationPurpose::MainResponse,
4235 llm.complete_with_tools(messages, None, &request),
4236 )
4237 .await?
4238 } else {
4239 let prompt_messages =
4240 self.prompt_messages_for_tool_protocol(messages, protocol, corrective);
4241 self.observe_purpose(
4242 ObservationPurpose::MainResponse,
4243 llm.complete(&prompt_messages, None),
4244 )
4245 .await?
4246 };
4247 Ok(MainProviderResponse {
4248 response,
4249 used_native_tools: use_native,
4250 })
4251 }
4252
4253 async fn complete_main_attempt_with_recovery(
4254 &self,
4255 llm: Arc<dyn LLMProvider>,
4256 messages: &[ChatMessage],
4257 protocol: &MainToolProtocol,
4258 corrective: bool,
4259 ) -> Result<MainProviderResponse> {
4260 let primary_result = self
4262 .recovery_manager
4263 .with_llm_retry(
4264 "llm_call",
4265 None,
4266 || {
4267 let llm = Arc::clone(&llm);
4268 async move {
4269 self.invoke_main_provider(llm, messages, protocol, corrective)
4270 .await
4271 }
4272 },
4273 |error| llm.is_terminal_error(error),
4274 )
4275 .await;
4276
4277 match primary_result {
4278 Ok(response) => Ok(response),
4279 Err(ai_agents_recovery::RetryFailure::Terminal { error, .. }) => {
4280 Err(AgentError::LLM(error.to_string()))
4281 }
4282 Err(failure) => {
4283 let primary_error = AgentError::LLM(failure.into_error().to_string());
4284 match &self.recovery_manager.config().llm.on_failure {
4285 LLMFailureAction::FallbackLlm { fallback_llm } => {
4286 let fallback = self.llm_registry.get(fallback_llm).map_err(|error| {
4287 AgentError::Config(format!(
4288 "Fallback LLM '{fallback_llm}' not found: {error}"
4289 ))
4290 })?;
4291 self.invoke_main_provider(fallback, messages, protocol, corrective)
4292 .await
4293 .map_err(|error| AgentError::LLM(error.to_string()))
4294 }
4295 LLMFailureAction::FallbackResponse { message } => {
4296 if matches!(
4297 protocol.choice.as_ref(),
4298 Some(ToolChoice::Required | ToolChoice::Specific(_))
4299 ) {
4300 Err(AgentError::LLM(format!(
4301 "Required tool selection failed and cannot be satisfied by a static fallback response: {primary_error}"
4302 )))
4303 } else {
4304 Ok(MainProviderResponse {
4305 response: LLMResponse::new(message.clone(), FinishReason::Stop),
4306 used_native_tools: false,
4307 })
4308 }
4309 }
4310 LLMFailureAction::Error => Err(primary_error),
4311 }
4312 }
4313 }
4314 }
4315
4316 fn normalize_main_provider_response(
4317 &self,
4318 mut response: LLMResponse,
4319 protocol: &MainToolProtocol,
4320 ) -> Result<(LLMResponse, bool)> {
4321 let provider_state = response
4322 .take_provider_state()
4323 .map_err(|error| AgentError::LLM(error.to_string()))?;
4324 let native_calls = response
4325 .tool_calls()
4326 .map_err(|error| AgentError::LLM(error.to_string()))?;
4327 let calls = match native_calls {
4328 Some(calls) => {
4329 response.content = encode_native_tool_call_markers(&calls, provider_state.as_ref())
4330 .map_err(|error| AgentError::LLM(error.to_string()))?;
4331 Some(calls)
4332 }
4333 None if provider_state.is_some() => {
4334 return Err(AgentError::LLM(
4335 "Provider returned replay state without native tool calls".to_string(),
4336 ));
4337 }
4338 None if !matches!(protocol.choice.as_ref(), Some(ToolChoice::None)) => {
4339 self.parse_tool_calls(response.content.trim())?
4340 }
4341 None => None,
4342 };
4343
4344 if protocol.choice.is_some()
4345 && let Some(calls) = calls.as_ref()
4346 && calls.iter().any(|call| {
4347 self.tools
4348 .canonical_id(&call.name)
4349 .is_none_or(|canonical| !protocol.tool_ids.contains(&canonical))
4350 })
4351 {
4352 return Err(AgentError::LLM(
4353 "Provider returned a tool call outside the effective grant".to_string(),
4354 ));
4355 }
4356
4357 let compliant = match protocol.choice.as_ref() {
4358 Some(ToolChoice::Required) => calls.as_ref().is_some_and(|calls| !calls.is_empty()),
4359 Some(ToolChoice::Specific(expected)) => calls.as_ref().is_some_and(|calls| {
4360 !calls.is_empty()
4361 && calls.iter().all(|call| {
4362 self.tools.canonical_id(&call.name).as_deref() == Some(expected.as_str())
4363 })
4364 }),
4365 _ => true,
4366 };
4367 Ok((response, compliant))
4368 }
4369
4370 async fn complete_main_llm_with_recovery(
4371 &self,
4372 llm: Arc<dyn LLMProvider>,
4373 messages: &[ChatMessage],
4374 protocol: &MainToolProtocol,
4375 ) -> Result<LLMResponse> {
4376 let first = self
4377 .complete_main_attempt_with_recovery(Arc::clone(&llm), messages, protocol, false)
4378 .await?;
4379 let (response, compliant) =
4380 self.normalize_main_provider_response(first.response, protocol)?;
4381 if compliant {
4382 return Ok(response);
4383 }
4384 if first.used_native_tools {
4385 return Err(AgentError::LLM(
4386 "Provider returned no compliant native call for required tool choice".to_string(),
4387 ));
4388 }
4389
4390 let corrected = self
4391 .complete_main_attempt_with_recovery(llm, messages, protocol, true)
4392 .await?;
4393 let (response, compliant) =
4394 self.normalize_main_provider_response(corrected.response, protocol)?;
4395 if compliant {
4396 return Ok(response);
4397 }
4398 Err(AgentError::LLM(
4399 "Provider returned no compliant tool call after one corrective retry".to_string(),
4400 ))
4401 }
4402
4403 async fn open_main_stream_with_recovery(
4410 &self,
4411 llm: Arc<dyn LLMProvider>,
4412 messages: &[ChatMessage],
4413 protocol: &MainToolProtocol,
4414 ) -> Result<MainStreamSource> {
4415 debug_assert!(
4416 protocol.choice.is_none(),
4417 "streaming raw path must not run with explicit tool choice"
4418 );
4419 let primary = self
4420 .recovery_manager
4421 .with_llm_retry(
4422 "llm_stream_open",
4423 None,
4424 || {
4425 let llm = Arc::clone(&llm);
4426 async move {
4427 self.observe_purpose(
4428 ObservationPurpose::MainResponse,
4429 llm.complete_stream(messages, None),
4430 )
4431 .await
4432 }
4433 },
4434 |error| llm.is_terminal_error(error),
4435 )
4436 .await;
4437
4438 match primary {
4439 Ok(stream) => Ok(MainStreamSource::Stream(stream)),
4440 Err(ai_agents_recovery::RetryFailure::Terminal { error, .. }) => {
4441 Err(AgentError::LLM(error.to_string()))
4442 }
4443 Err(failure) => {
4444 let primary_error = AgentError::LLM(failure.into_error().to_string());
4445 match &self.recovery_manager.config().llm.on_failure {
4446 LLMFailureAction::FallbackLlm { fallback_llm } => {
4447 let fallback = self.llm_registry.get(fallback_llm).map_err(|error| {
4448 AgentError::Config(format!(
4449 "Fallback LLM '{fallback_llm}' not found: {error}"
4450 ))
4451 })?;
4452 if fallback.supports(LLMFeature::Streaming) {
4453 let stream = self
4454 .observe_purpose(
4455 ObservationPurpose::MainResponse,
4456 fallback.complete_stream(messages, None),
4457 )
4458 .await
4459 .map_err(|error| AgentError::LLM(error.to_string()))?;
4460 Ok(MainStreamSource::Stream(stream))
4461 } else {
4462 let response = self
4463 .observe_purpose(
4464 ObservationPurpose::MainResponse,
4465 fallback.complete(messages, None),
4466 )
4467 .await
4468 .map_err(|error| AgentError::LLM(error.to_string()))?;
4469 Ok(MainStreamSource::StaticResponse(response.content))
4470 }
4471 }
4472 LLMFailureAction::FallbackResponse { message } => {
4473 Ok(MainStreamSource::StaticResponse(message.clone()))
4474 }
4475 LLMFailureAction::Error => Err(primary_error),
4476 }
4477 }
4478 }
4479 }
4480
4481 fn main_stream_must_buffer(
4490 &self,
4491 reasoning_mode: &ReasoningMode,
4492 protocol: &MainToolProtocol,
4493 ) -> bool {
4494 protocol.choice.is_some()
4495 || self.get_effective_reflection_config().requires_evaluation()
4496 || matches!(reasoning_mode, ReasoningMode::CoT | ReasoningMode::React)
4497 }
4498
4499 fn is_native_tool_call_content(content: &str) -> Result<bool> {
4501 decode_native_tool_call_markers(content)
4502 .map(|batch| batch.is_some())
4503 .map_err(|error| AgentError::LLM(error.to_string()))
4504 }
4505
4506 fn tool_result_message(
4508 tool_call: &ToolCall,
4509 output: &str,
4510 native_tool_call: bool,
4511 ) -> Result<ChatMessage> {
4512 if !native_tool_call {
4513 return Ok(ChatMessage::function(&tool_call.name, output));
4514 }
4515 let output = serde_json::from_str::<serde_json::Value>(output)
4516 .unwrap_or_else(|_| serde_json::Value::String(output.to_string()));
4517 let content = encode_native_tool_result_marker(tool_call, output)
4518 .map_err(|error| AgentError::LLM(error.to_string()))?;
4519 Ok(ChatMessage::function(&tool_call.name, content))
4520 }
4521
4522 fn remember_active_native_exchange(&self, content: &str) -> Result<()> {
4524 let Some(batch) = decode_native_tool_call_markers(content)
4525 .map_err(|error| AgentError::LLM(error.to_string()))?
4526 else {
4527 return Ok(());
4528 };
4529 let Some(state) = batch.provider_state() else {
4530 return Ok(());
4531 };
4532 let expected = ActiveNativeExchange {
4533 exchange_id: state.exchange_id().to_string(),
4534 call_ids: batch.calls().iter().map(|call| call.id.clone()).collect(),
4535 };
4536 let mut active = self.active_native_exchanges.write();
4537 if let Some(existing) = active
4538 .iter()
4539 .find(|existing| existing.exchange_id == expected.exchange_id)
4540 {
4541 if existing.call_ids != expected.call_ids {
4542 return Err(AgentError::LLM(format!(
4543 "Active native exchange '{}' changed its call identities",
4544 expected.exchange_id
4545 )));
4546 }
4547 } else {
4548 active.push(expected);
4549 }
4550 Ok(())
4551 }
4552
4553 fn validate_active_native_history(
4555 &self,
4556 messages: &[ChatMessage],
4557 require_complete: bool,
4558 ) -> Result<()> {
4559 let expected = self.active_native_exchanges.read().clone();
4560 if expected.is_empty() {
4561 return Ok(());
4562 }
4563 let inspection =
4564 inspect_native_history(messages).map_err(|error| AgentError::LLM(error.to_string()))?;
4565 let expected_count = expected.len();
4566 for (index, expected) in expected.iter().enumerate() {
4567 let Some(exchange) = inspection
4568 .exchanges()
4569 .iter()
4570 .find(|exchange| exchange.state().exchange_id() == expected.exchange_id)
4571 else {
4572 return Err(AgentError::LLM(format!(
4573 "Active native exchange '{}' was removed before provider continuation",
4574 expected.exchange_id
4575 )));
4576 };
4577 let must_be_complete = require_complete || index + 1 < expected_count;
4578 if exchange.call_ids() != expected.call_ids
4579 || (must_be_complete && !exchange.is_complete())
4580 {
4581 return Err(AgentError::LLM(format!(
4582 "Active native exchange '{}' is incomplete before provider continuation",
4583 expected.exchange_id
4584 )));
4585 }
4586 }
4587 Ok(())
4588 }
4589
4590 async fn remember_committed_native_exchange(&self, content: &str) -> Result<()> {
4592 self.remember_active_native_exchange(content)?;
4593 if !self.active_native_exchanges.read().is_empty() {
4594 let messages = self.memory.get_messages(None).await?;
4595 self.validate_active_native_history(&messages, false)?;
4596 }
4597 Ok(())
4598 }
4599
4600 fn readable_native_messages(mut messages: Vec<ChatMessage>) -> Result<Vec<ChatMessage>> {
4602 for message in &mut messages {
4603 if matches!(
4604 message.role,
4605 ai_agents_core::Role::Assistant
4606 | ai_agents_core::Role::Tool
4607 | ai_agents_core::Role::Function
4608 ) {
4609 message.content = native_readable_projection(&message.content)
4610 .map_err(|error| AgentError::LLM(error.to_string()))?;
4611 }
4612 }
4613 Ok(messages)
4614 }
4615
4616 fn parse_main_tool_calls(
4618 &self,
4619 content: &str,
4620 protocol: &MainToolProtocol,
4621 ) -> Result<Option<Vec<ToolCall>>> {
4622 if matches!(protocol.choice.as_ref(), Some(ToolChoice::None)) {
4623 Ok(None)
4624 } else {
4625 self.parse_tool_calls(content)
4626 }
4627 }
4628
4629 fn parse_tool_calls(&self, content: &str) -> Result<Option<Vec<ToolCall>>> {
4631 if let Some(batch) = decode_native_tool_call_markers(content)
4632 .map_err(|error| AgentError::LLM(error.to_string()))?
4633 {
4634 return Ok(Some(batch.into_parts().0));
4635 }
4636 if let Ok(parsed) = serde_json::from_str::<serde_json::Value>(content) {
4638 if let Some(arr) = parsed.as_array() {
4640 let calls: Vec<ToolCall> = arr
4641 .iter()
4642 .filter_map(|v| self.extract_tool_call_from_value(v))
4643 .collect();
4644 if !calls.is_empty() {
4645 return Ok(Some(calls));
4646 }
4647 }
4648 if let Some(tool_call) = self.extract_tool_call_from_value(&parsed) {
4650 return Ok(Some(vec![tool_call]));
4651 }
4652 }
4653
4654 if let Some(json_str) = self.extract_json_from_content(content)
4656 && let Ok(parsed) = serde_json::from_str::<serde_json::Value>(&json_str)
4657 {
4658 if let Some(arr) = parsed.as_array() {
4660 let calls: Vec<ToolCall> = arr
4661 .iter()
4662 .filter_map(|v| self.extract_tool_call_from_value(v))
4663 .collect();
4664 if !calls.is_empty() {
4665 return Ok(Some(calls));
4666 }
4667 }
4668 if let Some(tool_call) = self.extract_tool_call_from_value(&parsed) {
4670 return Ok(Some(vec![tool_call]));
4671 }
4672 }
4673
4674 Ok(None)
4675 }
4676
4677 fn extract_tool_call_from_value(&self, parsed: &serde_json::Value) -> Option<ToolCall> {
4678 if let Some(tool_name) = parsed.get("tool").and_then(|v| v.as_str()) {
4679 let arguments = parsed
4680 .get("arguments")
4681 .cloned()
4682 .unwrap_or(serde_json::json!({}));
4683 return Some(ToolCall {
4684 id: parsed
4685 .get("id")
4686 .and_then(|value| value.as_str())
4687 .filter(|id| !id.is_empty())
4688 .map(str::to_string)
4689 .unwrap_or_else(|| uuid::Uuid::new_v4().to_string()),
4690 name: tool_name.to_string(),
4691 arguments,
4692 });
4693 }
4694 None
4695 }
4696
4697 fn extract_json_from_content(&self, content: &str) -> Option<String> {
4699 if let Some(result) = self.extract_json_array_from_content(content) {
4701 return Some(result);
4702 }
4703 self.extract_json_object_from_content(content)
4704 }
4705
4706 fn extract_json_array_from_content(&self, content: &str) -> Option<String> {
4708 let start = content.find('[')?;
4709 let content_from_start = &content[start..];
4710
4711 let mut depth = 0;
4712 let mut end = 0;
4713 for (i, ch) in content_from_start.char_indices() {
4714 match ch {
4715 '[' => depth += 1,
4716 ']' => {
4717 depth -= 1;
4718 if depth == 0 {
4719 end = i + 1;
4720 break;
4721 }
4722 }
4723 _ => {}
4724 }
4725 }
4726
4727 if end > 0 {
4728 let json_str = &content_from_start[..end];
4729 if json_str.contains("\"tool\"") {
4731 return Some(json_str.to_string());
4732 }
4733 }
4734
4735 None
4736 }
4737
4738 fn extract_json_object_from_content(&self, content: &str) -> Option<String> {
4740 let start = content.find('{')?;
4741 let content_from_start = &content[start..];
4742
4743 let mut depth = 0;
4745 let mut end = 0;
4746 for (i, ch) in content_from_start.char_indices() {
4747 match ch {
4748 '{' => depth += 1,
4749 '}' => {
4750 depth -= 1;
4751 if depth == 0 {
4752 end = i + 1;
4753 break;
4754 }
4755 }
4756 _ => {}
4757 }
4758 }
4759
4760 if end > 0 {
4761 let json_str = &content_from_start[..end];
4762 if json_str.contains("\"tool\"") {
4764 return Some(json_str.to_string());
4765 }
4766 }
4767
4768 None
4769 }
4770
4771 #[allow(clippy::too_many_arguments)]
4775 fn record_from_parts(
4776 &self,
4777 request: &ToolExecutionRequest,
4778 canonical_id: String,
4779 executed_arguments: Value,
4780 started_at: chrono::DateTime<chrono::Utc>,
4781 start: Instant,
4782 executed: bool,
4783 success: bool,
4784 output: String,
4785 metadata: HashMap<String, Value>,
4786 policy: ToolPolicyDecisionRecord,
4787 approval: Option<ToolApprovalRecord>,
4788 timed_out: bool,
4789 output_truncated: bool,
4790 ) -> ToolExecutionRecord {
4791 let versions = ToolDecisionVersions {
4792 policy: self.active_tool_security().policy_version(),
4793 registry: self.tools.version(),
4794 runtime_control: self.runtime_control.version.load(Ordering::SeqCst),
4795 state: self
4796 .state_machine
4797 .as_ref()
4798 .map(|state_machine| state_machine.generation()),
4799 };
4800 self.record_from_parts_at(
4801 request,
4802 canonical_id,
4803 executed_arguments,
4804 started_at,
4805 start,
4806 executed,
4807 success,
4808 output,
4809 metadata,
4810 policy,
4811 approval,
4812 timed_out,
4813 output_truncated,
4814 versions,
4815 )
4816 }
4817
4818 #[allow(clippy::too_many_arguments)]
4820 fn record_from_parts_at(
4821 &self,
4822 request: &ToolExecutionRequest,
4823 canonical_id: String,
4824 executed_arguments: Value,
4825 started_at: chrono::DateTime<chrono::Utc>,
4826 start: Instant,
4827 executed: bool,
4828 success: bool,
4829 output: String,
4830 metadata: HashMap<String, Value>,
4831 policy: ToolPolicyDecisionRecord,
4832 approval: Option<ToolApprovalRecord>,
4833 timed_out: bool,
4834 output_truncated: bool,
4835 versions: ToolDecisionVersions,
4836 ) -> ToolExecutionRecord {
4837 ToolExecutionRecord {
4838 call_id: request.call_id.clone(),
4839 requested_name: request.requested_name.clone(),
4840 canonical_id,
4841 source: request.source.clone(),
4842 arguments: request.arguments.clone(),
4843 executed_arguments,
4844 policy_version: versions.policy,
4845 registry_version: versions.registry,
4846 runtime_config_version: versions.runtime_control,
4847 executed,
4848 success,
4849 output,
4850 metadata,
4851 policy,
4852 approval,
4853 started_at,
4854 duration_ms: start.elapsed().as_millis() as u64,
4855 timed_out,
4856 cancelled: false,
4857 cancellation_reason: None,
4858 output_truncated,
4859 }
4860 }
4861
4862 async fn finish_tool_record(&self, record: &ToolExecutionRecord) {
4864 let result = ToolResult {
4865 success: record.success,
4866 output: record.model_output_string(),
4867 metadata: if record.metadata.is_empty() {
4868 None
4869 } else {
4870 Some(record.metadata.clone())
4871 },
4872 };
4873 self.hooks
4874 .on_tool_complete(&record.canonical_id, &result, record.duration_ms)
4875 .await;
4876 self.hooks.on_tool_execution_record(record).await;
4877 self.record_tool_call(&record.canonical_id, record.model_output_value());
4878 if !record.success {
4879 self.hooks
4880 .on_error(&AgentError::Tool(record.output.clone()))
4881 .await;
4882 }
4883 }
4884
4885 async fn finish_tool_record_after_resource_guards(
4887 &self,
4888 resource_guards: ToolResourceGuards,
4889 record: &ToolExecutionRecord,
4890 ) {
4891 drop(resource_guards);
4892 self.finish_tool_record(record).await;
4893 }
4894
4895 fn validated_tool_timeout(timeout_ms: u64) -> Result<ValidatedToolTimeout> {
4899 if timeout_ms > MAX_TOOL_TIMEOUT_MS {
4900 return Err(AgentError::Config(format!(
4901 "effective tool timeout_ms must be no greater than {MAX_TOOL_TIMEOUT_MS} milliseconds"
4902 )));
4903 }
4904 let timer = Duration::from_millis(timeout_ms);
4905 let deadline_delta = chrono::Duration::from_std(timer).map_err(|_| {
4906 AgentError::Config(format!(
4907 "effective tool timeout_ms cannot be represented as a UTC deadline: {timeout_ms}"
4908 ))
4909 })?;
4910 Ok(ValidatedToolTimeout {
4911 timer,
4912 deadline_delta,
4913 })
4914 }
4915
4916 fn effective_tool_limits(
4920 security_engine: &ToolSecurityEngine,
4921 canonical_id: &str,
4922 safety: &ToolSafetyMetadata,
4923 classification: &ToolCallClassification,
4924 recovery_timeout_ms: Option<u64>,
4925 ) -> Result<(ToolExecutionLimits, ValidatedToolTimeout)> {
4926 if let Some(timeout_ms) = classification.timeout_ms {
4927 Self::validated_tool_timeout(timeout_ms)?;
4928 }
4929 if let Some(timeout_ms) = recovery_timeout_ms {
4930 Self::validated_tool_timeout(timeout_ms)?;
4931 }
4932
4933 let mut limits = security_engine.effective_limits(canonical_id, safety, classification);
4934 if let Some(recovery_timeout_ms) = recovery_timeout_ms {
4935 limits.timeout_ms = Some(limits.timeout_ms.map_or(recovery_timeout_ms, |timeout_ms| {
4936 timeout_ms.min(recovery_timeout_ms)
4937 }));
4938 }
4939 let timeout_ms = limits
4940 .timeout_ms
4941 .unwrap_or_else(|| security_engine.get_tool_timeout(canonical_id));
4942 let timeout = Self::validated_tool_timeout(timeout_ms)?;
4943 Ok((limits, timeout))
4944 }
4945
4946 async fn execute_resolved_tool_once(
4948 &self,
4949 tool: Arc<dyn ai_agents_core::Tool>,
4950 args: Value,
4951 mut ctx: ToolExecutionContext,
4952 timeout: ValidatedToolTimeout,
4953 ) -> Result<(ToolResult, bool, bool, bool)> {
4954 if self.runtime_control.emergency_deny.load(Ordering::SeqCst) {
4955 return Ok((
4956 ToolResult::error("Tool execution cancelled by runtime control"),
4957 false,
4958 true,
4959 false,
4960 ));
4961 }
4962 ctx.deadline = Some(
4967 chrono::Utc::now()
4968 .checked_add_signed(timeout.deadline_delta)
4969 .ok_or_else(|| {
4970 AgentError::Config(
4971 "effective tool timeout_ms exceeds the current UTC deadline range"
4972 .to_string(),
4973 )
4974 })?,
4975 );
4976 let invoked = Arc::new(AtomicBool::new(false));
4980 let invoked_by_future = Arc::clone(&invoked);
4981 let actor_context = current_turn_actor_context();
4982 let future = async move {
4983 invoked_by_future.store(true, Ordering::SeqCst);
4984 if let Some(actor_context) = actor_context {
4985 scope_actor_context(actor_context, tool.execute(args, ctx)).await
4986 } else {
4987 tool.execute(args, ctx).await
4988 }
4989 };
4990 tokio::pin!(future);
4991 let timer = tokio::time::sleep(timeout.timer);
4992 tokio::pin!(timer);
4993 let mut cancel_tick = tokio::time::interval(std::time::Duration::from_millis(50));
4994
4995 loop {
4996 tokio::select! {
4997 result = &mut future => return Ok((result, false, false, true)),
4998 _ = &mut timer => {
4999 return Ok((
5000 ToolResult::error("Tool execution timed out"),
5001 true,
5002 false,
5003 invoked.load(Ordering::SeqCst),
5004 ));
5005 }
5006 _ = cancel_tick.tick() => {
5007 if self.runtime_control.emergency_deny.load(Ordering::SeqCst) {
5008 return Ok((
5009 ToolResult::error("Tool execution cancelled by runtime control"),
5010 false,
5011 true,
5012 invoked.load(Ordering::SeqCst),
5013 ));
5014 }
5015 }
5016 }
5017 }
5018 }
5019
5020 fn truncate_tool_output(output: String, max_chars: Option<usize>) -> (String, bool) {
5022 let Some(max_chars) = max_chars else {
5023 return (output, false);
5024 };
5025 let mut chars = output.chars();
5026 let truncated: String = chars.by_ref().take(max_chars).collect();
5027 if chars.next().is_some() {
5028 (truncated, true)
5029 } else {
5030 (output, false)
5031 }
5032 }
5033
5034 async fn acquire_tool_resource_locks(&self, keys: &[String]) -> Option<ToolResourceGuards> {
5036 let locks = {
5037 let mut table = self.resource_locks.write();
5038 table.retain(|_, lock| lock.strong_count() > 0);
5039 keys.iter()
5040 .map(|key| {
5041 if let Some(lock) = table.get(key).and_then(Weak::upgrade) {
5042 lock
5043 } else {
5044 let lock = Arc::new(tokio::sync::Mutex::new(()));
5045 table.insert(key.clone(), Arc::downgrade(&lock));
5046 lock
5047 }
5048 })
5049 .collect::<Vec<_>>()
5050 };
5051 let mut resource_guards = ToolResourceGuards {
5052 guards: Vec::with_capacity(locks.len()),
5053 locks: Arc::clone(&self.resource_locks),
5054 };
5055 let mut locks = locks.into_iter();
5056 while let Some(lock) = locks.next() {
5057 let mut lock = Box::pin(lock.lock_owned());
5058 loop {
5059 tokio::select! {
5060 guard = &mut lock => {
5061 resource_guards.guards.push(guard);
5062 break;
5063 }
5064 _ = tokio::time::sleep(std::time::Duration::from_millis(10)) => {
5065 if self.runtime_control.emergency_deny.load(Ordering::SeqCst) {
5066 drop(lock);
5067 drop(locks);
5068 drop(resource_guards);
5069 return None;
5070 }
5071 }
5072 }
5073 }
5074 }
5075 Some(resource_guards)
5076 }
5077
5078 async fn run_tool_with_retries(
5082 &self,
5083 canonical_id: &str,
5084 tool: Arc<dyn ai_agents_core::Tool>,
5085 args: Value,
5086 ctx: ToolExecutionContext,
5087 timeout: ValidatedToolTimeout,
5088 max_retries: u32,
5089 ) -> Result<(ToolResult, bool, bool, bool)> {
5090 let max_retries = if ctx.classification.safely_retryable {
5091 max_retries
5092 } else {
5093 0
5094 };
5095 let mut attempts = 0;
5096 let mut invoked = false;
5097 loop {
5098 let (result, timed_out, cancelled, attempt_invoked) = self
5099 .execute_resolved_tool_once(tool.clone(), args.clone(), ctx.clone(), timeout)
5100 .await?;
5101 invoked |= attempt_invoked;
5102 if result.success || timed_out || cancelled || attempts >= max_retries {
5103 return Ok((result, timed_out, cancelled, invoked));
5104 }
5105 attempts += 1;
5106 warn!(tool = %canonical_id, attempt = attempts, error = %result.output, "Retrying failed tool call");
5107 }
5108 }
5109
5110 fn host_tool_unavailability(&self, canonical_id: &str) -> Option<(&'static str, &'static str)> {
5112 match canonical_id {
5113 "command" if !self.tools.command_runner_available() => Some((
5114 "Command runner is unavailable",
5115 "command runner is unavailable",
5116 )),
5117 "diagnostics" if !self.tools.diagnostics_available() => Some((
5118 "Diagnostics provider is unavailable",
5119 "diagnostics provider is unavailable",
5120 )),
5121 "web_search" if !self.tools.web_search_available() => Some((
5122 "Web search provider is unavailable",
5123 "web search provider is unavailable",
5124 )),
5125 _ => None,
5126 }
5127 }
5128
5129 fn execute_tool_record(
5131 &self,
5132 request: ToolExecutionRequest,
5133 ) -> Pin<Box<dyn Future<Output = Result<ToolExecutionRecord>> + Send + '_>> {
5134 Box::pin(self.execute_tool_record_inner(request, ToolFallbackState::default()))
5135 }
5136
5137 async fn execute_tool_record_inner(
5141 &self,
5142 request: ToolExecutionRequest,
5143 fallback_state: ToolFallbackState,
5144 ) -> Result<ToolExecutionRecord> {
5145 let started_at = chrono::Utc::now();
5146 let start = Instant::now();
5147 info!(tool = %request.requested_name, args = %request.arguments, "Executing tool");
5148
5149 if self.runtime_control.emergency_deny.load(Ordering::SeqCst) {
5150 let record = self.record_from_parts(
5151 &request,
5152 request.requested_name.clone(),
5153 request.arguments.clone(),
5154 started_at,
5155 start,
5156 false,
5157 false,
5158 "Tool execution is disabled by runtime control".to_string(),
5159 HashMap::new(),
5160 ToolPolicyDecisionRecord::deny("runtime emergency deny is enabled"),
5161 None,
5162 false,
5163 false,
5164 );
5165 self.finish_tool_record(&record).await;
5166 return Ok(record);
5167 }
5168
5169 let Some(resolved) = self.tools.resolve(&request.requested_name) else {
5170 let record = self.record_from_parts(
5171 &request,
5172 request.requested_name.clone(),
5173 request.arguments.clone(),
5174 started_at,
5175 start,
5176 false,
5177 false,
5178 format!("Tool '{}' is unavailable", request.requested_name),
5179 HashMap::new(),
5180 ToolPolicyDecisionRecord::unavailable(format!(
5181 "Tool '{}' is not registered",
5182 request.requested_name
5183 )),
5184 None,
5185 false,
5186 false,
5187 );
5188 self.finish_tool_record(&record).await;
5189 return Ok(record);
5190 };
5191
5192 let canonical_id = resolved.identity.canonical_id.clone();
5193
5194 let initial_scope_snapshot = self.get_available_tool_ids_snapshot().await?;
5195 if !initial_scope_snapshot
5196 .tool_ids
5197 .iter()
5198 .any(|id| id == &canonical_id)
5199 {
5200 let record = self.record_from_parts(
5201 &request,
5202 canonical_id.clone(),
5203 request.arguments.clone(),
5204 started_at,
5205 start,
5206 false,
5207 false,
5208 format!(
5209 "Tool '{}' is not available in the current scope",
5210 canonical_id
5211 ),
5212 HashMap::new(),
5213 ToolPolicyDecisionRecord::deny(format!(
5214 "Tool '{}' is not granted by the current top-level and state tool scope",
5215 canonical_id
5216 )),
5217 None,
5218 false,
5219 false,
5220 );
5221 self.finish_tool_record(&record).await;
5222 return Ok(record);
5223 }
5224
5225 let approval_control_snapshot = self.runtime_safety_snapshot();
5226 let security_engine = approval_control_snapshot.tool_security.clone();
5227 if let Some(reason) = fallback_state.rejection_reason(&canonical_id) {
5228 let mut metadata = HashMap::new();
5229 metadata.insert(
5230 "fallback_chain".to_string(),
5231 serde_json::to_value(&fallback_state.visited_canonical_ids).unwrap_or(Value::Null),
5232 );
5233 let record = self.record_from_parts(
5234 &request,
5235 canonical_id,
5236 request.arguments.clone(),
5237 started_at,
5238 start,
5239 false,
5240 false,
5241 format!("Denied: {reason}"),
5242 metadata,
5243 ToolPolicyDecisionRecord::deny(reason),
5244 None,
5245 false,
5246 false,
5247 );
5248 self.finish_tool_record(&record).await;
5249 return Ok(record);
5250 }
5251 let admitted_canonical_id = canonical_id.clone();
5252 let fallback_state = fallback_state.with_current(canonical_id.clone());
5253 let bindings = resolved.tool.policy_bindings();
5254 let mut executed_arguments = security_engine.prepare_tool_arguments_with_bindings(
5255 &canonical_id,
5256 &request.arguments,
5257 &bindings,
5258 );
5259 let mut metadata = HashMap::new();
5260 let safety = resolved.tool.safety_metadata();
5261 let classification = resolved.tool.classify_call(&executed_arguments);
5262 let initial_recovery_timeout_ms = self.recovery_manager.get_tool_timeout(&canonical_id);
5263 let (limits, _) = Self::effective_tool_limits(
5264 &security_engine,
5265 &canonical_id,
5266 &safety,
5267 &classification,
5268 initial_recovery_timeout_ms,
5269 )?;
5270 self.hooks
5271 .on_tool_start(&canonical_id, &executed_arguments)
5272 .await;
5273 metadata.insert(
5274 "classification".to_string(),
5275 serde_json::to_value(&classification).unwrap_or(Value::Null),
5276 );
5277 metadata.insert(
5278 "effective_limits".to_string(),
5279 serde_json::to_value(&limits).unwrap_or(Value::Null),
5280 );
5281 let policy_snapshot = security_engine.policy_snapshot(&canonical_id);
5282 if !policy_snapshot.is_null() {
5283 metadata.insert("policy_snapshot".to_string(), policy_snapshot.clone());
5284 }
5285
5286 let mut approval_record = Some(ToolApprovalRecord {
5287 status: ToolApprovalStatus::NotRequired,
5288 reason: None,
5289 modified_arguments: None,
5290 });
5291
5292 let mut security_result = security_engine
5293 .validate_tool_execution_with_bindings(&canonical_id, &executed_arguments, &bindings)
5294 .await?;
5295 if (security_result.is_allowed()
5300 || matches!(
5301 &security_result,
5302 SecurityCheckResult::RequireConfirmation { .. }
5303 ))
5304 && let Some((output, reason)) = self.host_tool_unavailability(&canonical_id)
5305 {
5306 let record = self.record_from_parts(
5307 &request,
5308 canonical_id,
5309 executed_arguments,
5310 started_at,
5311 start,
5312 false,
5313 false,
5314 output.to_string(),
5315 metadata,
5316 ToolPolicyDecisionRecord::unavailable(reason),
5317 Some(ToolApprovalRecord {
5318 status: ToolApprovalStatus::Unavailable,
5319 reason: Some(reason.to_string()),
5320 modified_arguments: None,
5321 }),
5322 false,
5323 false,
5324 );
5325 self.finish_tool_record(&record).await;
5326 return Ok(record);
5327 }
5328 match &security_result {
5329 SecurityCheckResult::Allow => {}
5330 SecurityCheckResult::Warn { message } => {
5331 warn!(tool = %canonical_id, message = %message, "Tool security warning");
5332 }
5333 SecurityCheckResult::Block { reason } => {
5334 let record = self.record_from_parts(
5335 &request,
5336 canonical_id,
5337 executed_arguments,
5338 started_at,
5339 start,
5340 false,
5341 false,
5342 format!("Denied: {}", reason),
5343 metadata,
5344 ToolPolicyDecisionRecord::deny(reason.clone()),
5345 approval_record,
5346 false,
5347 false,
5348 );
5349 self.finish_tool_record(&record).await;
5350 return Ok(record);
5351 }
5352 SecurityCheckResult::Unavailable { reason } => {
5353 let record = self.record_from_parts(
5354 &request,
5355 canonical_id,
5356 executed_arguments,
5357 started_at,
5358 start,
5359 false,
5360 false,
5361 format!("Unavailable: {}", reason),
5362 metadata,
5363 ToolPolicyDecisionRecord::unavailable(reason.clone()),
5364 approval_record,
5365 false,
5366 false,
5367 );
5368 self.finish_tool_record(&record).await;
5369 return Ok(record);
5370 }
5371 SecurityCheckResult::RequireConfirmation { message } => {
5372 if self.hitl_engine.is_none() {
5373 approval_record = Some(ToolApprovalRecord {
5374 status: ToolApprovalStatus::Unavailable,
5375 reason: Some("No HITL engine configured".to_string()),
5376 modified_arguments: None,
5377 });
5378 let record = self.record_from_parts(
5379 &request,
5380 canonical_id,
5381 executed_arguments,
5382 started_at,
5383 start,
5384 false,
5385 false,
5386 format!("Approval unavailable: {}", message),
5387 metadata,
5388 ToolPolicyDecisionRecord::approval(message.clone()),
5389 approval_record,
5390 false,
5391 false,
5392 );
5393 self.finish_tool_record(&record).await;
5394 return Ok(record);
5395 }
5396
5397 let check_result = HITLCheckResult::required(
5398 ApprovalTrigger::tool(&canonical_id, executed_arguments.clone()),
5399 HashMap::new(),
5400 message.clone(),
5401 None,
5402 );
5403 match self.request_hitl_approval(check_result).await? {
5404 ApprovalResult::Approved => {
5405 merge_approved_record(&mut approval_record);
5406 }
5407 ApprovalResult::Modified { changes } => {
5408 if let Some(obj) = executed_arguments.as_object_mut() {
5409 for (key, value) in changes {
5410 obj.insert(key, value);
5411 }
5412 }
5413 security_result = security_engine
5414 .validate_tool_execution_with_bindings(
5415 &canonical_id,
5416 &executed_arguments,
5417 &bindings,
5418 )
5419 .await?;
5420 if !matches!(
5421 security_result,
5422 SecurityCheckResult::Allow
5423 | SecurityCheckResult::Warn { .. }
5424 | SecurityCheckResult::RequireConfirmation { .. }
5425 ) {
5426 let reason = security_result
5427 .reason()
5428 .unwrap_or("modified arguments failed policy")
5429 .to_string();
5430 let record = self.record_from_parts(
5431 &request,
5432 canonical_id,
5433 executed_arguments.clone(),
5434 started_at,
5435 start,
5436 false,
5437 false,
5438 reason.clone(),
5439 metadata,
5440 ToolPolicyDecisionRecord::deny(reason),
5441 Some(ToolApprovalRecord {
5442 status: ToolApprovalStatus::Modified,
5443 reason: None,
5444 modified_arguments: Some(executed_arguments),
5445 }),
5446 false,
5447 false,
5448 );
5449 self.finish_tool_record(&record).await;
5450 return Ok(record);
5451 }
5452 approval_record = Some(ToolApprovalRecord {
5453 status: ToolApprovalStatus::Modified,
5454 reason: None,
5455 modified_arguments: Some(executed_arguments.clone()),
5456 });
5457 }
5458 ApprovalResult::Rejected { reason } => {
5459 let reason = reason.unwrap_or_else(|| "rejected".to_string());
5460 approval_record = Some(ToolApprovalRecord {
5461 status: ToolApprovalStatus::Rejected,
5462 reason: Some(reason.clone()),
5463 modified_arguments: None,
5464 });
5465 let record = self.record_from_parts(
5466 &request,
5467 canonical_id,
5468 executed_arguments,
5469 started_at,
5470 start,
5471 false,
5472 false,
5473 format!("Approval rejected: {}", reason),
5474 metadata,
5475 ToolPolicyDecisionRecord::approval(reason),
5476 approval_record,
5477 false,
5478 false,
5479 );
5480 self.finish_tool_record(&record).await;
5481 return Ok(record);
5482 }
5483 ApprovalResult::Timeout => {
5484 approval_record = Some(ToolApprovalRecord {
5485 status: ToolApprovalStatus::Timeout,
5486 reason: Some("approval timeout".to_string()),
5487 modified_arguments: None,
5488 });
5489 let record = self.record_from_parts(
5490 &request,
5491 canonical_id,
5492 executed_arguments,
5493 started_at,
5494 start,
5495 false,
5496 false,
5497 "Approval timed out".to_string(),
5498 metadata,
5499 ToolPolicyDecisionRecord::approval("approval timeout"),
5500 approval_record,
5501 false,
5502 false,
5503 );
5504 self.finish_tool_record(&record).await;
5505 return Ok(record);
5506 }
5507 }
5508 }
5509 }
5510
5511 if approval_record
5512 .as_ref()
5513 .is_some_and(|record| matches!(record.status, ToolApprovalStatus::NotRequired))
5514 && let Some(message) =
5515 security_engine.classification_approval_message(&canonical_id, &classification)
5516 {
5517 if self.hitl_engine.is_none() {
5518 approval_record = Some(ToolApprovalRecord {
5519 status: ToolApprovalStatus::Unavailable,
5520 reason: Some("No HITL engine configured".to_string()),
5521 modified_arguments: None,
5522 });
5523 let record = self.record_from_parts(
5524 &request,
5525 canonical_id,
5526 executed_arguments,
5527 started_at,
5528 start,
5529 false,
5530 false,
5531 format!("Approval unavailable: {}", message),
5532 metadata,
5533 ToolPolicyDecisionRecord::approval(message),
5534 approval_record,
5535 false,
5536 false,
5537 );
5538 self.finish_tool_record(&record).await;
5539 return Ok(record);
5540 }
5541 let check_result = HITLCheckResult::required(
5542 ApprovalTrigger::tool(&canonical_id, executed_arguments.clone()),
5543 HashMap::new(),
5544 message.clone(),
5545 None,
5546 );
5547 match self.request_hitl_approval(check_result).await? {
5548 ApprovalResult::Approved => {
5549 merge_approved_record(&mut approval_record);
5550 }
5551 ApprovalResult::Modified { changes } => {
5552 if let Some(obj) = executed_arguments.as_object_mut() {
5553 for (key, value) in changes {
5554 obj.insert(key, value);
5555 }
5556 }
5557 let modified_security = security_engine
5558 .validate_tool_execution_with_bindings(
5559 &canonical_id,
5560 &executed_arguments,
5561 &bindings,
5562 )
5563 .await?;
5564 if !matches!(
5565 modified_security,
5566 SecurityCheckResult::Allow | SecurityCheckResult::Warn { .. }
5567 ) {
5568 let reason = modified_security
5569 .reason()
5570 .unwrap_or("modified arguments failed policy")
5571 .to_string();
5572 let record = self.record_from_parts(
5573 &request,
5574 canonical_id,
5575 executed_arguments.clone(),
5576 started_at,
5577 start,
5578 false,
5579 false,
5580 reason.clone(),
5581 metadata,
5582 ToolPolicyDecisionRecord::deny(reason),
5583 Some(ToolApprovalRecord {
5584 status: ToolApprovalStatus::Modified,
5585 reason: None,
5586 modified_arguments: Some(executed_arguments),
5587 }),
5588 false,
5589 false,
5590 );
5591 self.finish_tool_record(&record).await;
5592 return Ok(record);
5593 }
5594 approval_record = Some(ToolApprovalRecord {
5595 status: ToolApprovalStatus::Modified,
5596 reason: None,
5597 modified_arguments: Some(executed_arguments.clone()),
5598 });
5599 }
5600 ApprovalResult::Rejected { reason } => {
5601 let reason = reason.unwrap_or_else(|| "rejected".to_string());
5602 let record = self.record_from_parts(
5603 &request,
5604 canonical_id,
5605 executed_arguments,
5606 started_at,
5607 start,
5608 false,
5609 false,
5610 format!("Approval rejected: {}", reason),
5611 metadata,
5612 ToolPolicyDecisionRecord::approval(reason.clone()),
5613 Some(ToolApprovalRecord {
5614 status: ToolApprovalStatus::Rejected,
5615 reason: Some(reason),
5616 modified_arguments: None,
5617 }),
5618 false,
5619 false,
5620 );
5621 self.finish_tool_record(&record).await;
5622 return Ok(record);
5623 }
5624 ApprovalResult::Timeout => {
5625 let record = self.record_from_parts(
5626 &request,
5627 canonical_id,
5628 executed_arguments,
5629 started_at,
5630 start,
5631 false,
5632 false,
5633 "Approval timed out".to_string(),
5634 metadata,
5635 ToolPolicyDecisionRecord::approval("approval timeout"),
5636 Some(ToolApprovalRecord {
5637 status: ToolApprovalStatus::Timeout,
5638 reason: Some("approval timeout".to_string()),
5639 modified_arguments: None,
5640 }),
5641 false,
5642 false,
5643 );
5644 self.finish_tool_record(&record).await;
5645 return Ok(record);
5646 }
5647 }
5648 }
5649
5650 let hitl_lang_ctx = self.build_hitl_language_context();
5651 if let Some(ref hitl_engine) = self.hitl_engine {
5652 let check_result = self
5653 .observe_purpose(
5654 ObservationPurpose::HitlLocalization,
5655 hitl_engine.check_tool_with_localization(
5656 &canonical_id,
5657 &executed_arguments,
5658 &hitl_lang_ctx,
5659 self.approval_handler.as_ref(),
5660 Some(&self.llm_registry),
5661 ),
5662 )
5663 .await?;
5664 if check_result.is_required() {
5665 match self.request_hitl_approval(check_result).await? {
5666 ApprovalResult::Approved => {
5667 merge_approved_record(&mut approval_record);
5668 }
5669 ApprovalResult::Modified { changes } => {
5670 if let Some(obj) = executed_arguments.as_object_mut() {
5671 for (key, value) in changes {
5672 obj.insert(key, value);
5673 }
5674 }
5675 let modified_security = security_engine
5676 .validate_tool_execution_with_bindings(
5677 &canonical_id,
5678 &executed_arguments,
5679 &bindings,
5680 )
5681 .await?;
5682 if !matches!(
5683 modified_security,
5684 SecurityCheckResult::Allow | SecurityCheckResult::Warn { .. }
5685 ) {
5686 let reason = modified_security
5687 .reason()
5688 .unwrap_or("modified arguments failed policy")
5689 .to_string();
5690 let record = self.record_from_parts(
5691 &request,
5692 canonical_id,
5693 executed_arguments.clone(),
5694 started_at,
5695 start,
5696 false,
5697 false,
5698 reason.clone(),
5699 metadata,
5700 ToolPolicyDecisionRecord::deny(reason),
5701 Some(ToolApprovalRecord {
5702 status: ToolApprovalStatus::Modified,
5703 reason: None,
5704 modified_arguments: Some(executed_arguments),
5705 }),
5706 false,
5707 false,
5708 );
5709 self.finish_tool_record(&record).await;
5710 return Ok(record);
5711 }
5712 approval_record = Some(ToolApprovalRecord {
5713 status: ToolApprovalStatus::Modified,
5714 reason: None,
5715 modified_arguments: Some(executed_arguments.clone()),
5716 });
5717 }
5718 ApprovalResult::Rejected { reason } => {
5719 let reason = reason.unwrap_or_else(|| "rejected".to_string());
5720 let record = self.record_from_parts(
5721 &request,
5722 canonical_id,
5723 executed_arguments,
5724 started_at,
5725 start,
5726 false,
5727 false,
5728 format!("Approval rejected: {}", reason),
5729 metadata,
5730 ToolPolicyDecisionRecord::approval(reason.clone()),
5731 Some(ToolApprovalRecord {
5732 status: ToolApprovalStatus::Rejected,
5733 reason: Some(reason),
5734 modified_arguments: None,
5735 }),
5736 false,
5737 false,
5738 );
5739 self.finish_tool_record(&record).await;
5740 return Ok(record);
5741 }
5742 ApprovalResult::Timeout => {
5743 let record = self.record_from_parts(
5744 &request,
5745 canonical_id,
5746 executed_arguments,
5747 started_at,
5748 start,
5749 false,
5750 false,
5751 "Approval timed out".to_string(),
5752 metadata,
5753 ToolPolicyDecisionRecord::approval("approval timeout"),
5754 Some(ToolApprovalRecord {
5755 status: ToolApprovalStatus::Timeout,
5756 reason: Some("approval timeout".to_string()),
5757 modified_arguments: None,
5758 }),
5759 false,
5760 false,
5761 );
5762 self.finish_tool_record(&record).await;
5763 return Ok(record);
5764 }
5765 }
5766 }
5767
5768 let condition_check = self
5769 .observe_purpose(
5770 ObservationPurpose::HitlLocalization,
5771 hitl_engine.check_conditions_with_localization(
5772 &executed_arguments,
5773 &hitl_lang_ctx,
5774 self.approval_handler.as_ref(),
5775 Some(&self.llm_registry),
5776 ),
5777 )
5778 .await?;
5779 if condition_check.is_required() {
5780 match self.request_hitl_approval(condition_check).await? {
5781 ApprovalResult::Approved => {
5782 merge_approved_record(&mut approval_record);
5783 }
5784 ApprovalResult::Modified { changes } => {
5785 if let Some(obj) = executed_arguments.as_object_mut() {
5786 for (key, value) in changes {
5787 obj.insert(key, value);
5788 }
5789 }
5790 let modified_security = security_engine
5791 .validate_tool_execution_with_bindings(
5792 &canonical_id,
5793 &executed_arguments,
5794 &bindings,
5795 )
5796 .await?;
5797 if !matches!(
5798 modified_security,
5799 SecurityCheckResult::Allow | SecurityCheckResult::Warn { .. }
5800 ) {
5801 let reason = modified_security
5802 .reason()
5803 .unwrap_or("modified arguments failed policy")
5804 .to_string();
5805 let record = self.record_from_parts(
5806 &request,
5807 canonical_id,
5808 executed_arguments,
5809 started_at,
5810 start,
5811 false,
5812 false,
5813 reason.clone(),
5814 metadata,
5815 ToolPolicyDecisionRecord::deny(reason),
5816 approval_record,
5817 false,
5818 false,
5819 );
5820 self.finish_tool_record(&record).await;
5821 return Ok(record);
5822 }
5823 approval_record = Some(ToolApprovalRecord {
5824 status: ToolApprovalStatus::Modified,
5825 reason: None,
5826 modified_arguments: Some(executed_arguments.clone()),
5827 });
5828 }
5829 ApprovalResult::Rejected { reason } => {
5830 let reason = reason.unwrap_or_else(|| "rejected".to_string());
5831 let record = self.record_from_parts(
5832 &request,
5833 canonical_id,
5834 executed_arguments,
5835 started_at,
5836 start,
5837 false,
5838 false,
5839 format!("Approval rejected: {}", reason),
5840 metadata,
5841 ToolPolicyDecisionRecord::approval(reason.clone()),
5842 Some(ToolApprovalRecord {
5843 status: ToolApprovalStatus::Rejected,
5844 reason: Some(reason),
5845 modified_arguments: None,
5846 }),
5847 false,
5848 false,
5849 );
5850 self.finish_tool_record(&record).await;
5851 return Ok(record);
5852 }
5853 ApprovalResult::Timeout => {
5854 let record = self.record_from_parts(
5855 &request,
5856 canonical_id,
5857 executed_arguments,
5858 started_at,
5859 start,
5860 false,
5861 false,
5862 "Approval timed out".to_string(),
5863 metadata,
5864 ToolPolicyDecisionRecord::approval("approval timeout"),
5865 Some(ToolApprovalRecord {
5866 status: ToolApprovalStatus::Timeout,
5867 reason: Some("approval timeout".to_string()),
5868 modified_arguments: None,
5869 }),
5870 false,
5871 false,
5872 );
5873 self.finish_tool_record(&record).await;
5874 return Ok(record);
5875 }
5876 }
5877 }
5878 }
5879
5880 executed_arguments = security_engine.prepare_tool_arguments_with_bindings(
5885 &canonical_id,
5886 &executed_arguments,
5887 &bindings,
5888 );
5889 if let Some(record) = approval_record.as_mut()
5890 && matches!(record.status, ToolApprovalStatus::Modified)
5891 {
5892 record.modified_arguments = Some(executed_arguments.clone());
5893 }
5894 let binding_security_result = security_engine
5895 .validate_tool_execution_with_bindings(&canonical_id, &executed_arguments, &bindings)
5896 .await?;
5897 let approval_confirmation_required = matches!(
5898 binding_security_result,
5899 SecurityCheckResult::RequireConfirmation { .. }
5900 ) || security_engine
5901 .classification_approval_message(
5902 &canonical_id,
5903 &resolved.tool.classify_call(&executed_arguments),
5904 )
5905 .is_some();
5906 let approval_binding = approval_record.as_ref().and_then(|record| {
5907 matches!(
5908 record.status,
5909 ToolApprovalStatus::Approved | ToolApprovalStatus::Modified
5910 )
5911 .then(|| ToolApprovalBinding {
5912 canonical_id: canonical_id.clone(),
5913 arguments: executed_arguments.clone(),
5914 confirmation_required: approval_confirmation_required,
5915 policy_version: security_engine.policy_version(),
5916 runtime_control_version: approval_control_snapshot.version,
5917 state_generation: initial_scope_snapshot.state_generation,
5918 reviewed_tool: Arc::clone(&resolved.tool),
5919 })
5920 });
5921
5922 let control_snapshot = self.runtime_safety_snapshot();
5927 let resolved = self.tools.resolve(&request.requested_name);
5928 let registry_version = self.tools.version();
5929 let mut versions = ToolDecisionVersions {
5930 policy: control_snapshot.tool_security.policy_version(),
5931 registry: registry_version,
5932 runtime_control: control_snapshot.version,
5933 state: None,
5934 };
5935 metadata.insert(
5936 "runtime_scope_snapshot".to_string(),
5937 serde_json::to_value(&control_snapshot.tool_scope_override).unwrap_or(Value::Null),
5938 );
5939 let resolved = match resolved {
5940 Some(resolved) => resolved,
5941 None => {
5942 let reason = format!(
5943 "Tool '{}' became unavailable after approval",
5944 request.requested_name
5945 );
5946 let record = self.record_from_parts_at(
5947 &request,
5948 request.requested_name.clone(),
5949 executed_arguments,
5950 started_at,
5951 start,
5952 false,
5953 false,
5954 reason.clone(),
5955 metadata,
5956 ToolPolicyDecisionRecord::unavailable(reason),
5957 approval_record,
5958 false,
5959 false,
5960 versions,
5961 );
5962 self.finish_tool_record(&record).await;
5963 return Ok(record);
5964 }
5965 };
5966
5967 let canonical_id = resolved.identity.canonical_id.clone();
5968 if let Some(reason) =
5969 fallback_state.final_rejection_reason(&admitted_canonical_id, &canonical_id)
5970 {
5971 metadata.insert(
5975 "fallback_chain".to_string(),
5976 serde_json::to_value(&fallback_state.visited_canonical_ids).unwrap_or(Value::Null),
5977 );
5978 metadata.insert(
5979 "final_resolved_canonical_id".to_string(),
5980 Value::String(canonical_id),
5981 );
5982 let record = self.record_from_parts_at(
5983 &request,
5984 admitted_canonical_id,
5985 executed_arguments,
5986 started_at,
5987 start,
5988 false,
5989 false,
5990 format!("Denied: {reason}"),
5991 metadata,
5992 ToolPolicyDecisionRecord::deny(reason),
5993 approval_record,
5994 false,
5995 false,
5996 versions,
5997 );
5998 self.finish_tool_record(&record).await;
5999 return Ok(record);
6000 }
6001 let bindings = resolved.tool.policy_bindings();
6002 let final_arguments = control_snapshot
6003 .tool_security
6004 .prepare_tool_arguments_with_bindings(&canonical_id, &executed_arguments, &bindings);
6005 if let Some(record) = approval_record.as_mut()
6006 && matches!(record.status, ToolApprovalStatus::Modified)
6007 {
6008 record.modified_arguments = Some(final_arguments.clone());
6009 }
6010 let classification = resolved.tool.classify_call(&final_arguments);
6011 let safety = resolved.tool.safety_metadata();
6012 let security_engine = control_snapshot.tool_security;
6013 let tool_config = self.recovery_manager.get_tool_config(&canonical_id).clone();
6014 let recovery_timeout_ms = self.recovery_manager.get_tool_timeout(&canonical_id);
6015 metadata.insert(
6016 "classification".to_string(),
6017 serde_json::to_value(&classification).unwrap_or(Value::Null),
6018 );
6019 let (limits, timeout) = match Self::effective_tool_limits(
6023 &security_engine,
6024 &canonical_id,
6025 &safety,
6026 &classification,
6027 recovery_timeout_ms,
6028 ) {
6029 Ok(effective) => effective,
6030 Err(error) => {
6031 let reason = error.to_string();
6032 metadata.insert(
6033 "configuration_error".to_string(),
6034 Value::String(reason.clone()),
6035 );
6036 let record = self.record_from_parts_at(
6037 &request,
6038 canonical_id,
6039 final_arguments,
6040 started_at,
6041 start,
6042 false,
6043 false,
6044 format!("Denied: {reason}"),
6045 metadata,
6046 ToolPolicyDecisionRecord::deny(reason),
6047 approval_record,
6048 false,
6049 false,
6050 versions,
6051 );
6052 self.finish_tool_record(&record).await;
6053 return Ok(record);
6054 }
6055 };
6056 let policy_snapshot = security_engine.policy_snapshot(&canonical_id);
6057 let resource_lock_keys =
6058 tool_resource_lock_keys(&canonical_id, &final_arguments, &bindings, &classification);
6059 metadata.insert(
6060 "effective_limits".to_string(),
6061 serde_json::to_value(&limits).unwrap_or(Value::Null),
6062 );
6063 metadata.insert(
6064 "resource_lock_keys".to_string(),
6065 serde_json::to_value(&resource_lock_keys).unwrap_or(Value::Null),
6066 );
6067 if policy_snapshot.is_null() {
6068 metadata.remove("policy_snapshot");
6069 } else {
6070 metadata.insert("policy_snapshot".to_string(), policy_snapshot.clone());
6071 }
6072
6073 let final_denial = |canonical_id: String,
6074 output: String,
6075 policy: ToolPolicyDecisionRecord,
6076 metadata: HashMap<String, Value>,
6077 decision_versions: ToolDecisionVersions| {
6078 self.record_from_parts_at(
6079 &request,
6080 canonical_id,
6081 final_arguments.clone(),
6082 started_at,
6083 start,
6084 false,
6085 false,
6086 output,
6087 metadata,
6088 policy,
6089 approval_record.clone(),
6090 false,
6091 false,
6092 decision_versions,
6093 )
6094 };
6095
6096 if control_snapshot.emergency_deny {
6097 let reason = "Tool execution is disabled by runtime control".to_string();
6098 let record = final_denial(
6099 canonical_id,
6100 reason.clone(),
6101 ToolPolicyDecisionRecord::deny(reason),
6102 metadata,
6103 versions,
6104 );
6105 self.finish_tool_record(&record).await;
6106 return Ok(record);
6107 }
6108
6109 let available_snapshot = self
6114 .get_available_tool_ids_snapshot_for_scope(
6115 control_snapshot.tool_scope_override.as_deref(),
6116 )
6117 .await?;
6118 versions.state = available_snapshot.state_generation;
6119 metadata.insert(
6120 "available_tool_ids_snapshot".to_string(),
6121 serde_json::to_value(&available_snapshot.tool_ids).unwrap_or(Value::Null),
6122 );
6123 metadata.insert(
6124 "state_generation_snapshot".to_string(),
6125 serde_json::to_value(available_snapshot.state_generation).unwrap_or(Value::Null),
6126 );
6127 if !available_snapshot
6128 .tool_ids
6129 .iter()
6130 .any(|tool_id| tool_id == &canonical_id)
6131 {
6132 let reason = format!(
6133 "Tool '{}' is not available in the final runtime scope",
6134 canonical_id
6135 );
6136 let record = final_denial(
6137 canonical_id,
6138 reason.clone(),
6139 ToolPolicyDecisionRecord::deny(reason),
6140 metadata,
6141 versions,
6142 );
6143 self.finish_tool_record(&record).await;
6144 return Ok(record);
6145 }
6146
6147 let final_security_result = security_engine
6152 .validate_tool_execution_with_bindings(&canonical_id, &final_arguments, &bindings)
6153 .await?;
6154 match &final_security_result {
6155 SecurityCheckResult::Block { reason } => {
6156 let record = final_denial(
6157 canonical_id,
6158 format!("Denied: {}", reason),
6159 ToolPolicyDecisionRecord::deny(reason.clone()),
6160 metadata,
6161 versions,
6162 );
6163 self.finish_tool_record(&record).await;
6164 return Ok(record);
6165 }
6166 SecurityCheckResult::Unavailable { reason } => {
6167 let record = final_denial(
6168 canonical_id,
6169 format!("Unavailable: {}", reason),
6170 ToolPolicyDecisionRecord::unavailable(reason.clone()),
6171 metadata,
6172 versions,
6173 );
6174 self.finish_tool_record(&record).await;
6175 return Ok(record);
6176 }
6177 SecurityCheckResult::Warn { message } => {
6178 warn!(tool = %canonical_id, message = %message, "Tool security warning after approval");
6179 }
6180 SecurityCheckResult::Allow | SecurityCheckResult::RequireConfirmation { .. } => {}
6181 }
6182 let final_confirmation_required = matches!(
6183 final_security_result,
6184 SecurityCheckResult::RequireConfirmation { .. }
6185 ) || security_engine
6186 .classification_approval_message(&canonical_id, &classification)
6187 .is_some();
6188 let stale_approval = approval_binding.as_ref().is_some_and(|binding| {
6189 binding.is_stale(
6190 &canonical_id,
6191 &final_arguments,
6192 final_confirmation_required,
6193 versions,
6194 &resolved.tool,
6195 )
6196 });
6197 if stale_approval {
6198 let reason = "Approval became stale before final admission".to_string();
6199 let record = final_denial(
6200 canonical_id,
6201 reason.clone(),
6202 ToolPolicyDecisionRecord::deny(reason),
6203 metadata,
6204 versions,
6205 );
6206 self.finish_tool_record(&record).await;
6207 return Ok(record);
6208 }
6209 if final_confirmation_required && approval_binding.is_none() {
6210 let reason = "Final policy requires fresh approval".to_string();
6211 let record = final_denial(
6212 canonical_id,
6213 reason.clone(),
6214 ToolPolicyDecisionRecord::approval(reason),
6215 metadata,
6216 versions,
6217 );
6218 self.finish_tool_record(&record).await;
6219 return Ok(record);
6220 }
6221
6222 if let Some((_, reason)) = self.host_tool_unavailability(&canonical_id) {
6223 let record = final_denial(
6224 canonical_id,
6225 reason.to_string(),
6226 ToolPolicyDecisionRecord::unavailable(reason),
6227 metadata,
6228 versions,
6229 );
6230 self.finish_tool_record(&record).await;
6231 return Ok(record);
6232 }
6233
6234 let Some(resource_guards) = self.acquire_tool_resource_locks(&resource_lock_keys).await
6239 else {
6240 let reason = "Tool execution cancelled while waiting for resource locks".to_string();
6244 let mut record = final_denial(
6245 canonical_id,
6246 reason.clone(),
6247 ToolPolicyDecisionRecord::deny(reason),
6248 metadata,
6249 versions,
6250 );
6251 record.cancelled = true;
6252 record.cancellation_reason = Some("runtime control cancellation".to_string());
6253 self.finish_tool_record(&record).await;
6254 return Ok(record);
6255 };
6256
6257 let admission = self.admit_tool_execution(
6262 versions.runtime_control,
6263 versions.policy,
6264 versions.state,
6265 &canonical_id,
6266 );
6267 if !matches!(admission, SecurityCheckResult::Allow) {
6268 let latest_control = self.runtime_safety_snapshot();
6269 let reason = admission
6270 .reason()
6271 .unwrap_or("tool admission was denied")
6272 .to_string();
6273 let policy = if admission.is_unavailable() {
6274 ToolPolicyDecisionRecord::unavailable(reason.clone())
6275 } else {
6276 ToolPolicyDecisionRecord::deny(reason.clone())
6277 };
6278 let record = self.record_from_parts_at(
6279 &request,
6280 canonical_id,
6281 final_arguments,
6282 started_at,
6283 start,
6284 false,
6285 false,
6286 reason,
6287 metadata,
6288 policy,
6289 approval_record,
6290 false,
6291 false,
6292 ToolDecisionVersions {
6293 policy: latest_control.tool_security.policy_version(),
6294 registry: versions.registry,
6295 runtime_control: latest_control.version,
6296 state: self
6297 .state_machine
6298 .as_ref()
6299 .map(|state_machine| state_machine.generation()),
6300 },
6301 );
6302 self.finish_tool_record_after_resource_guards(resource_guards, &record)
6303 .await;
6304 return Ok(record);
6305 }
6306 let executed_arguments = final_arguments;
6307
6308 let turn_actor = current_turn_actor_context();
6309 let actor = ToolActorContext {
6310 actor_id: turn_actor
6311 .as_ref()
6312 .and_then(|context| context.effective_actor_id().map(str::to_string))
6313 .or_else(|| self.actor_id()),
6314 origin_actor_id: turn_actor
6315 .as_ref()
6316 .and_then(|context| context.origin_actor_id.clone()),
6317 sender_agent_id: turn_actor
6318 .as_ref()
6319 .and_then(|context| context.sender_agent_id.clone()),
6320 };
6321 let tool_context = ToolExecutionContext {
6322 requested_name: request.requested_name.clone(),
6323 canonical_id: canonical_id.clone(),
6324 display_name: resolved.identity.display_name.clone(),
6325 provider_id: resolved.identity.provider_id.clone(),
6326 registry_version: versions.registry,
6327 policy_version: versions.policy,
6328 runtime_control_version: versions.runtime_control,
6329 call_id: request.call_id.clone(),
6330 source: request.source.clone(),
6331 actor,
6332 cancellation: ToolCancellationToken::new(
6333 Arc::clone(&self.runtime_control.emergency_deny),
6334 Some("runtime control cancellation".to_string()),
6335 ),
6336 started_at,
6337 deadline: None,
6338 permission: ToolPolicyDecisionRecord::allow(),
6339 approval: approval_record.clone(),
6340 classification: classification.clone(),
6341 safety,
6342 limits: limits.clone(),
6343 policy_snapshot,
6344 custom_config: security_engine.custom_config(&canonical_id),
6345 };
6346 let (mut result, timed_out, cancelled, invoked) = self
6347 .run_tool_with_retries(
6348 &canonical_id,
6349 resolved.tool.clone(),
6350 executed_arguments.clone(),
6351 tool_context,
6352 timeout,
6353 tool_config.max_retries,
6354 )
6355 .await?;
6356
6357 let fallback_tool = if !result.success && !cancelled {
6361 match &tool_config.on_failure {
6362 ToolFailureAction::Skip => {
6363 result = ToolResult::ok(format!(
6364 "{{\"skipped\": true, \"reason\": \"Tool '{}' was skipped after failure\"}}",
6365 canonical_id
6366 ));
6367 None
6368 }
6369 ToolFailureAction::Fallback { fallback_tool } => Some(fallback_tool.clone()),
6370 ToolFailureAction::ReportError => None,
6371 }
6372 } else {
6373 None
6374 };
6375
6376 let output_cap = limits.max_output_chars;
6377 let (output, output_truncated) =
6378 Self::truncate_tool_output(result.output.clone(), output_cap);
6379 if let Some(result_metadata) = result.metadata {
6380 metadata.extend(result_metadata);
6381 }
6382 let mut record = self.record_from_parts_at(
6383 &request,
6384 canonical_id,
6385 executed_arguments,
6386 started_at,
6387 start,
6388 invoked,
6389 result.success,
6390 output,
6391 metadata,
6392 ToolPolicyDecisionRecord::allow(),
6393 approval_record,
6394 timed_out,
6395 output_truncated,
6396 versions,
6397 );
6398 record.cancelled = cancelled;
6399 if cancelled {
6400 record.cancellation_reason = Some("runtime control cancellation".to_string());
6401 }
6402 if let Some(fallback_tool) = fallback_tool {
6403 let fallback_arguments = record.executed_arguments.clone();
6404 let original_tool = record.canonical_id.clone();
6405 self.finish_tool_record_after_resource_guards(resource_guards, &record)
6409 .await;
6410 let fallback_request = ToolExecutionRequest::new(
6411 request.call_id.clone(),
6412 fallback_tool,
6413 fallback_arguments,
6414 ToolCallSource::Fallback { original_tool },
6415 );
6416 return Box::pin(self.execute_tool_record_inner(fallback_request, fallback_state))
6417 .await;
6418 }
6419 self.finish_tool_record_after_resource_guards(resource_guards, &record)
6420 .await;
6421 Ok(record)
6422 }
6423
6424 #[instrument(skip(self, tool_call), fields(tool = %tool_call.name))]
6425 async fn execute_tool_smart(&self, tool_call: &ToolCall) -> Result<String> {
6426 let record = self
6427 .execute_tool_record(ToolExecutionRequest::new(
6428 tool_call.id.clone(),
6429 tool_call.name.clone(),
6430 tool_call.arguments.clone(),
6431 ToolCallSource::Model,
6432 ))
6433 .await?;
6434 if record.success {
6435 Ok(record.model_output_string())
6436 } else if matches!(record.policy.outcome, PermissionOutcome::RequiresApproval) {
6437 Err(AgentError::HITLRejected(record.model_output_string()))
6438 } else {
6439 Err(AgentError::Tool(record.model_output_string()))
6440 }
6441 }
6442
6443 async fn select_skill_candidate(&self, input: &str) -> Result<Option<SkillCandidate>> {
6449 let Some(ref router) = self.skill_router else {
6450 return Ok(None);
6451 };
6452 let available_skills = self.get_available_skills();
6453 if available_skills.is_empty() {
6454 return Ok(None);
6455 }
6456 let skill_ids: Vec<&str> = available_skills.iter().map(|s| s.id.as_str()).collect();
6457 let Some(skill_id) = self
6458 .observe_purpose(
6459 ObservationPurpose::SkillRouting,
6460 router.select_skill_filtered(input, &skill_ids),
6461 )
6462 .await?
6463 else {
6464 return Ok(None);
6465 };
6466 let skill = router
6467 .get_skill(&skill_id)
6468 .cloned()
6469 .ok_or_else(|| AgentError::Skill(format!("Skill not found: {}", skill_id)))?;
6470 info!(skill_id = %skill_id, "Skill selected");
6471 Ok(Some(SkillCandidate::new(skill_id, skill)))
6472 }
6473
6474 async fn commit_skill_candidate_route_result(
6479 &self,
6480 candidate: SkillCandidate,
6481 input: &str,
6482 ) -> Result<SkillRouteResult> {
6483 let skill_id = candidate.skill_id;
6484 let skill = candidate.skill;
6485 let expected_state_generation = self
6486 .state_machine
6487 .as_ref()
6488 .map(|state_machine| state_machine.generation());
6489 let expected_disambiguation_epoch = self.disambiguation_epoch.load(Ordering::SeqCst);
6490 if let Some(ref skill_disambig) = skill.disambiguation
6491 && skill_disambig.enabled.unwrap_or(false)
6492 && let Some(ref disambiguator) = self.disambiguation_manager
6493 {
6494 let context = self.build_disambiguation_context().await?;
6495 let state_override = self
6496 .state_machine
6497 .as_ref()
6498 .and_then(|sm| sm.current_definition())
6499 .and_then(|def| def.disambiguation.clone());
6500
6501 let disambiguation_result = self
6502 .observe_purpose(
6503 ObservationPurpose::DisambiguationDetection,
6504 disambiguator.process_input_with_override(
6505 input,
6506 &context,
6507 state_override.as_ref(),
6508 Some(skill_disambig),
6509 ),
6510 )
6511 .await?;
6512 let current_state_generation = self
6513 .state_machine
6514 .as_ref()
6515 .map(|state_machine| state_machine.generation());
6516 if current_state_generation != expected_state_generation
6517 || self.disambiguation_epoch.load(Ordering::SeqCst) != expected_disambiguation_epoch
6518 {
6519 disambiguator.clear_pending().await;
6520 *self.pending_skill_id.write() = None;
6521 return Err(AgentError::Other(
6522 "State or reset ownership changed during skill disambiguation".to_string(),
6523 ));
6524 }
6525 match disambiguation_result {
6526 DisambiguationResult::Clear => {
6527 debug!(skill_id = %skill_id, "Skill disambiguation: clear");
6528 }
6529 DisambiguationResult::NeedsClarification {
6530 question,
6531 detection,
6532 } => {
6533 let admission = self
6534 .admit_disambiguation_redispatch(
6535 expected_disambiguation_epoch,
6536 expected_state_generation,
6537 )
6538 .await?;
6539 let awaiting_confirmation = disambiguator.has_pending_confirmation().await;
6540 info!(
6541 skill_id = %skill_id,
6542 ambiguity_type = ?detection.ambiguity_type,
6543 confidence = detection.confidence,
6544 "Skill requires clarification before execution"
6545 );
6546 *self.pending_skill_id.write() = Some(skill_id.clone());
6547 let response = AgentResponse::new(&question.question).with_metadata(
6548 "disambiguation",
6549 serde_json::json!({
6550 "status": if awaiting_confirmation { "awaiting_confirmation" } else { "awaiting_clarification" },
6551 "skill_id": skill_id,
6552 "options": question.options,
6553 "clarifying": question.clarifying,
6554 "detection": {
6555 "type": detection.ambiguity_type,
6556 "confidence": detection.confidence,
6557 "what_is_unclear": detection.what_is_unclear,
6558 }
6559 }),
6560 );
6561 drop(admission);
6562 return Ok(SkillRouteResult::NeedsClarification {
6563 response,
6564 ownership: Some(DisambiguationOwnership {
6565 epoch: expected_disambiguation_epoch,
6566 state_generation: expected_state_generation,
6567 }),
6568 });
6569 }
6570 DisambiguationResult::Clarified { enriched_input, .. } => {
6571 info!(skill_id = %skill_id, enriched = %enriched_input, "Skill disambiguation clarified");
6572 let admission = self
6573 .admit_disambiguation_redispatch(
6574 expected_disambiguation_epoch,
6575 expected_state_generation,
6576 )
6577 .await?;
6578 drop(admission);
6579 let content = self.execute_skill(&skill, &enriched_input).await?;
6580 return Ok(SkillRouteResult::Response { skill_id, content });
6581 }
6582 DisambiguationResult::ProceedWithBestGuess { enriched_input } => {
6583 info!(skill_id = %skill_id, "Skill disambiguation best guess");
6584 let admission = self
6585 .admit_disambiguation_redispatch(
6586 expected_disambiguation_epoch,
6587 expected_state_generation,
6588 )
6589 .await?;
6590 drop(admission);
6591 let content = self.execute_skill(&skill, &enriched_input).await?;
6592 return Ok(SkillRouteResult::Response { skill_id, content });
6593 }
6594 DisambiguationResult::GiveUp { reason } => {
6595 warn!(skill_id = %skill_id, reason = %reason, "Skill disambiguation gave up");
6596 let apology = self
6597 .generate_localized_apology(
6598 "Generate a brief, polite apology saying you couldn't understand the request. Be concise.",
6599 &reason,
6600 )
6601 .await
6602 .unwrap_or_else(|_| {
6603 format!("I'm sorry, I couldn't understand your request: {}", reason)
6604 });
6605 return Ok(SkillRouteResult::NeedsClarification {
6606 response: AgentResponse::new(&apology),
6607 ownership: None,
6608 });
6609 }
6610 DisambiguationResult::Escalate { reason } => {
6611 info!(skill_id = %skill_id, reason = %reason, "Skill disambiguation escalating");
6612 let apology = self
6613 .generate_localized_apology(
6614 "Explain briefly that you're transferring the user to a human agent for help.",
6615 &reason,
6616 )
6617 .await
6618 .unwrap_or_else(|_| {
6619 format!("I need human assistance to help with your request: {}", reason)
6620 });
6621 return Ok(SkillRouteResult::NeedsClarification {
6622 response: AgentResponse::new(&apology),
6623 ownership: None,
6624 });
6625 }
6626 DisambiguationResult::Abandoned { .. } => {
6627 debug!(skill_id = %skill_id, "Skill disambiguation abandoned");
6628 return Ok(SkillRouteResult::NoMatch);
6629 }
6630 }
6631 }
6632 let admission = self
6633 .admit_disambiguation_redispatch(
6634 expected_disambiguation_epoch,
6635 expected_state_generation,
6636 )
6637 .await?;
6638 drop(admission);
6639 let content = self.execute_skill(&skill, input).await?;
6640 Ok(SkillRouteResult::Response { skill_id, content })
6641 }
6642
6643 async fn try_skill_route(&self, input: &str) -> Result<SkillRouteResult> {
6645 if let Some(candidate) = self.select_skill_candidate(input).await? {
6646 self.commit_skill_candidate_route_result(candidate, input)
6647 .await
6648 } else {
6649 Ok(SkillRouteResult::NoMatch)
6650 }
6651 }
6652
6653 fn skill_clarification_needs_memory_record(response: &AgentResponse) -> bool {
6656 response
6657 .metadata
6658 .as_ref()
6659 .and_then(|m| m.get("disambiguation"))
6660 .and_then(|d| d.get("status"))
6661 .and_then(|s| s.as_str())
6662 == Some("awaiting_clarification")
6663 }
6664
6665 async fn commit_winning_skill_candidate(
6672 &self,
6673 candidate: SkillCandidate,
6674 processed_input: &str,
6675 input_context: &HashMap<String, Value>,
6676 ) -> Result<Option<AgentResponse>> {
6677 self.commit_root_user_message(processed_input).await?;
6678 match self
6679 .commit_skill_candidate_route_result(candidate, processed_input)
6680 .await?
6681 {
6682 SkillRouteResult::Response { skill_id, content } => self
6683 .handle_skill_response(processed_input, &skill_id, content, input_context)
6684 .await
6685 .map(Some),
6686 SkillRouteResult::NeedsClarification {
6687 response,
6688 ownership,
6689 } => {
6690 let admission = self
6691 .admit_optional_disambiguation_ownership(ownership)
6692 .await?;
6693 if Self::skill_clarification_needs_memory_record(&response) {
6694 self.memory
6695 .add_message(ChatMessage::assistant(&response.content))
6696 .await?;
6697 }
6698 drop(admission);
6699 self.finish_turn_if_root(&response).await?;
6700 Ok(Some(response))
6701 }
6702 SkillRouteResult::NoMatch => Ok(None),
6703 }
6704 }
6705
6706 async fn execute_skill(&self, skill: &SkillDefinition, input: &str) -> Result<String> {
6708 if let Some(ref executor) = self.skill_executor {
6709 let skill_reasoning = self.get_skill_reasoning_config(skill);
6710 let skill_reflection = self.get_skill_reflection_config(skill);
6711
6712 debug!(
6713 skill_id = %skill.id,
6714 reasoning_mode = ?skill_reasoning.mode,
6715 reflection_enabled = ?skill_reflection.enabled,
6716 "Skill reasoning/reflection config"
6717 );
6718
6719 let response = self
6720 .observe_purpose(
6721 ObservationPurpose::SkillPrompt,
6722 executor.execute_with_invoker(skill, input, serde_json::json!({}), self),
6723 )
6724 .await?;
6725
6726 if skill_reflection.requires_evaluation() && skill_reflection.is_enabled() {
6727 let should_reflect = self
6728 .should_reflect_with_config(input, &response, &skill_reflection)
6729 .await?;
6730 if should_reflect {
6731 let evaluated = self
6732 .evaluate_and_retry_with_config(input, response, &skill_reflection)
6733 .await?;
6734 return Ok(evaluated);
6735 }
6736 }
6737
6738 return Ok(response);
6739 }
6740 Err(AgentError::Skill(
6741 "No skill executor configured".to_string(),
6742 ))
6743 }
6744
6745 async fn execute_skill_by_id(&self, skill_id: &str, input: &str) -> Result<String> {
6748 let skill = self
6749 .skill_router
6750 .as_ref()
6751 .and_then(|r| r.get_skill(skill_id).cloned())
6752 .ok_or_else(|| AgentError::Skill(format!("Skill not found: {}", skill_id)))?;
6753 self.execute_skill(&skill, input).await
6754 }
6755
6756 async fn should_reflect_with_config(
6758 &self,
6759 input: &str,
6760 response: &str,
6761 config: &ReflectionConfig,
6762 ) -> Result<bool> {
6763 if !config.requires_evaluation() {
6764 return Ok(false);
6765 }
6766
6767 if config.is_enabled() {
6768 return Ok(true);
6769 }
6770
6771 let evaluator_llm = config
6772 .evaluator_llm
6773 .as_ref()
6774 .and_then(|alias| self.llm_registry.get(alias).ok())
6775 .or_else(|| self.llm_registry.router().ok())
6776 .or_else(|| self.llm_registry.default().ok());
6777
6778 let Some(llm) = evaluator_llm else {
6779 return Ok(false);
6780 };
6781
6782 let response_preview: String = response.chars().take(500).collect();
6783 let prompt = format!(
6784 r#"Should this response be evaluated for quality? Consider if it's a complex or important response.
6785
6786User query: "{}"
6787Response: "{}"
6788
6789Answer YES or NO only."#,
6790 input, response_preview
6791 );
6792
6793 let messages = vec![ChatMessage::user(&prompt)];
6794 let result = self
6795 .observe_purpose(
6796 ObservationPurpose::ReflectionDecision,
6797 llm.complete(&messages, None),
6798 )
6799 .await;
6800
6801 match result {
6802 Ok(resp) => Ok(resp.content.trim().to_uppercase().contains("YES")),
6803 Err(_) => Ok(false),
6804 }
6805 }
6806
6807 async fn evaluate_and_retry_with_config(
6808 &self,
6809 input: &str,
6810 mut response: String,
6811 config: &ReflectionConfig,
6812 ) -> Result<String> {
6813 let llm = self.get_state_llm()?;
6814 let mut attempts = 0u32;
6815 let max_retries = config.max_retries;
6816
6817 loop {
6818 let evaluation = self
6819 .evaluate_response_with_config(input, &response, config)
6820 .await?;
6821
6822 if evaluation.passed || attempts >= max_retries {
6823 info!(
6824 passed = evaluation.passed,
6825 confidence = evaluation.confidence,
6826 attempts = attempts + 1,
6827 "Skill reflection evaluation complete"
6828 );
6829 return Ok(response);
6830 }
6831
6832 debug!(
6833 attempt = attempts + 1,
6834 failed_criteria = evaluation.failed_criteria().count(),
6835 "Skill response did not meet criteria, retrying"
6836 );
6837
6838 let feedback: Vec<String> = evaluation
6839 .failed_criteria()
6840 .map(|c| format!("- {}", c.criterion))
6841 .collect();
6842
6843 let retry_prompt = format!(
6844 "Your previous response did not meet these criteria:\n{}\n\nPlease provide an improved response to: {}",
6845 feedback.join("\n"),
6846 input
6847 );
6848
6849 let messages = vec![ChatMessage::user(&retry_prompt)];
6850 let retry_response = self
6851 .observe_purpose(
6852 ObservationPurpose::ReflectionEvaluation,
6853 llm.complete(&messages, None),
6854 )
6855 .await
6856 .map_err(|e| AgentError::LLM(e.to_string()))?;
6857
6858 response = retry_response.content.trim().to_string();
6859 attempts += 1;
6860 }
6861 }
6862
6863 async fn evaluate_response_with_config(
6864 &self,
6865 input: &str,
6866 response: &str,
6867 config: &ReflectionConfig,
6868 ) -> Result<EvaluationResult> {
6869 let evaluator_llm = config
6870 .evaluator_llm
6871 .as_ref()
6872 .and_then(|alias| self.llm_registry.get(alias).ok())
6873 .or_else(|| self.llm_registry.router().ok())
6874 .or_else(|| self.llm_registry.default().ok())
6875 .ok_or_else(|| AgentError::Config("No LLM available for evaluation".into()))?;
6876
6877 let criteria = &config.criteria;
6878 let criteria_list = criteria
6879 .iter()
6880 .enumerate()
6881 .map(|(i, c)| format!("{}. {}", i + 1, c))
6882 .collect::<Vec<_>>()
6883 .join("\n");
6884
6885 let prompt = format!(
6886 r#"Evaluate this response against the criteria.
6887
6888User query: "{}"
6889
6890Response to evaluate: "{}"
6891
6892Criteria:
6893{}
6894
6895For each criterion, respond with:
6896- criterion number
6897- PASS or FAIL
6898- brief reason
6899
6900Then provide overall confidence (0.0 to 1.0) and whether it passes overall.
6901
6902Format:
69031. PASS/FAIL - reason
69042. PASS/FAIL - reason
6905...
6906CONFIDENCE: 0.X
6907OVERALL: PASS/FAIL"#,
6908 input, response, criteria_list
6909 );
6910
6911 let messages = vec![ChatMessage::user(&prompt)];
6912 let eval_response = self
6913 .observe_purpose(
6914 ObservationPurpose::ReflectionEvaluation,
6915 evaluator_llm.complete(&messages, None),
6916 )
6917 .await
6918 .map_err(|e| AgentError::LLM(format!("Evaluation failed: {}", e)))?;
6919
6920 let content = eval_response.content.to_uppercase();
6921 let llm_pass = content.contains("OVERALL: PASS");
6922
6923 let confidence = content
6924 .lines()
6925 .find(|l| l.contains("CONFIDENCE:"))
6926 .and_then(|l| {
6927 l.split(':')
6928 .nth(1)
6929 .and_then(|v| v.trim().parse::<f32>().ok())
6930 })
6931 .unwrap_or(if llm_pass { 0.8 } else { 0.4 });
6932
6933 let overall_pass = llm_pass && confidence >= config.pass_threshold;
6936
6937 let mut criteria_results = Vec::new();
6938 for (i, criterion) in criteria.iter().enumerate() {
6939 let line_marker = format!("{}.", i + 1);
6940 let passed = eval_response
6941 .content
6942 .lines()
6943 .find(|l| l.contains(&line_marker))
6944 .map(|l| l.to_uppercase().contains("PASS"))
6945 .unwrap_or(overall_pass);
6946
6947 if passed {
6948 criteria_results.push(CriterionResult::pass(criterion));
6949 } else {
6950 criteria_results.push(CriterionResult::fail(criterion, "Did not meet criterion"));
6951 }
6952 }
6953
6954 Ok(EvaluationResult::new(overall_pass, confidence).with_criteria(criteria_results))
6955 }
6956
6957 async fn process_input(&self, input: &str) -> Result<ProcessData> {
6959 if let Some(processor) = self.get_state_process_processor() {
6960 let purpose = observation_purpose_for_process(processor.input_purpose_hint());
6961 return self
6962 .observe_purpose(purpose, processor.process_input(input))
6963 .await;
6964 }
6965 if let Some(ref processor) = self.process_processor {
6966 let purpose = observation_purpose_for_process(processor.input_purpose_hint());
6967 self.observe_purpose(purpose, processor.process_input(input))
6968 .await
6969 } else {
6970 Ok(ProcessData::new(input))
6971 }
6972 }
6973
6974 async fn process_output(
6976 &self,
6977 output: &str,
6978 input_context: &std::collections::HashMap<String, serde_json::Value>,
6979 ) -> Result<ProcessData> {
6980 if let Some(processor) = self.get_state_process_processor() {
6981 let purpose = observation_purpose_for_process(processor.output_purpose_hint());
6982 return self
6983 .observe_purpose(purpose, processor.process_output(output, input_context))
6984 .await;
6985 }
6986 if let Some(ref processor) = self.process_processor {
6987 let purpose = observation_purpose_for_process(processor.output_purpose_hint());
6988 self.observe_purpose(purpose, processor.process_output(output, input_context))
6989 .await
6990 } else {
6991 Ok(ProcessData::new(output))
6992 }
6993 }
6994
6995 fn get_state_process_processor(&self) -> Option<ProcessProcessor> {
6997 let sm = self.state_machine.as_ref()?;
6998 let def = sm.current_definition()?;
6999 let config = def.process.as_ref()?;
7000 let mut processor = ProcessProcessor::new(config.clone());
7001 if let Some(ref registry) = Some(self.llm_registry.clone()) {
7002 processor = processor.with_llm_registry(registry.clone());
7003 }
7004 processor = processor.with_stage_observer(Arc::new(ObservabilityProcessStageObserver));
7005 Some(processor)
7006 }
7007
7008 async fn check_turn_timeout(&self) -> Result<()> {
7010 let Some(ref sm) = self.state_machine else {
7011 return Ok(());
7012 };
7013 let Some(timeout_state) = sm.check_timeout() else {
7014 return Ok(());
7015 };
7016 let claim_admission = self.disambiguation_admission.write().await;
7017 if sm.check_timeout().as_deref() != Some(timeout_state.as_str()) {
7018 return Ok(());
7019 }
7020 let Some(reservation) = self.reserve_state_transition() else {
7021 return Ok(());
7022 };
7023 let from_state = sm.current();
7024 let expected_state_generation = sm.generation();
7025 let expected_disambiguation_epoch = self.disambiguation_epoch.load(Ordering::SeqCst);
7026 let history_before = sm.history();
7027 drop(claim_admission);
7028
7029 self.execute_state_exit_actions(&from_state).await;
7030
7031 let admission = self.disambiguation_admission.write().await;
7032 if sm.current() != from_state
7033 || sm.generation() != expected_state_generation
7034 || self.disambiguation_epoch.load(Ordering::SeqCst) != expected_disambiguation_epoch
7035 || sm.check_timeout().as_deref() != Some(timeout_state.as_str())
7036 {
7037 return Ok(());
7038 }
7039 sm.transition_to(&timeout_state, "max_turns exceeded")?;
7040 self.invalidate_pending_confirmation("state_timeout").await;
7041 let entered = sm.current();
7042 let is_reentry = Self::state_was_previously_entered(&entered, &from_state, &history_before);
7043 drop(admission);
7044
7045 self.execute_state_enter_actions(&entered, is_reentry).await;
7046 drop(reservation);
7047 info!(to = %entered, "Timeout transition");
7048 Ok(())
7049 }
7050
7051 fn increment_turn(&self) {
7052 if let Some(ref sm) = self.state_machine {
7053 sm.increment_turn();
7054 }
7055 }
7056
7057 fn transitions_available_for_commit(&self) -> Option<(Vec<Transition>, String)> {
7058 let sm = self.state_machine.as_ref()?;
7059 let current = sm.current();
7060 let transitions: Vec<_> = sm
7061 .auto_transitions()
7062 .into_iter()
7063 .filter(|t| match t.cooldown_turns {
7064 Some(cd) if cd > 0 => {
7065 let resolved = sm.config().resolve_full_path(¤t, &t.to);
7066 !sm.is_on_cooldown(&resolved, cd)
7067 }
7068 _ => true,
7069 })
7070 .collect();
7071 Some((transitions, current))
7072 }
7073
7074 fn transition_reason(transition: &Transition) -> String {
7075 if transition.when.is_empty() {
7076 "guard condition met".to_string()
7077 } else {
7078 transition.when.clone()
7079 }
7080 }
7081
7082 fn build_transition_context(
7084 &self,
7085 user_message: &str,
7086 response: &str,
7087 current_state: &str,
7088 staged: Option<&HashMap<String, Value>>,
7089 ) -> TransitionContext {
7090 let context_map = staged
7091 .map(|writes| self.build_context_with_staged(writes))
7092 .unwrap_or_else(|| self.build_context_with_overlays());
7093 TransitionContext::new(user_message, response, current_state).with_context(context_map)
7094 }
7095
7096 async fn select_transition_candidate(
7098 &self,
7099 user_message: &str,
7100 response: &str,
7101 ) -> Result<Option<TransitionCandidate>> {
7102 let Some((transitions, current_state)) = self.transitions_available_for_commit() else {
7103 return Ok(None);
7104 };
7105 let transitions: Vec<Transition> = transitions
7106 .into_iter()
7107 .filter(|transition| matches!(transition.timing, TransitionTiming::PostResponse))
7108 .collect();
7109 if transitions.is_empty() {
7110 return Ok(None);
7111 }
7112 let Some(evaluator) = self.transition_evaluator.as_ref() else {
7113 return Ok(None);
7114 };
7115 let context = self.build_transition_context(user_message, response, ¤t_state, None);
7116 let selected = self
7117 .observe_purpose(
7118 ObservationPurpose::StateTransitionEvaluation,
7119 evaluator.select_transition(&transitions, &context),
7120 )
7121 .await?;
7122 Ok(selected.map(|index| {
7123 let transition = transitions[index].clone();
7124 TransitionCandidate::new(
7125 current_state,
7126 transition.clone(),
7127 Self::transition_reason(&transition),
7128 )
7129 }))
7130 }
7131
7132 fn select_deterministic_transition_candidate(
7134 &self,
7135 user_message: &str,
7136 current_state: &str,
7137 transitions: &[Transition],
7138 staged: &HashMap<String, Value>,
7139 ) -> Option<TransitionCandidate> {
7140 let context = self.build_transition_context(user_message, "", current_state, Some(staged));
7141
7142 for transition in transitions {
7143 if let Some(guard) = transition.guard.as_ref()
7144 && evaluate_guard(guard, &context)
7145 {
7146 return Some(TransitionCandidate::new(
7147 current_state,
7148 transition.clone(),
7149 Self::transition_reason(transition),
7150 ));
7151 }
7152 }
7153
7154 let resolved_intent = context
7155 .context
7156 .get("resolved_intent")
7157 .and_then(Value::as_str)
7158 .filter(|value| !value.is_empty());
7159 if let Some(resolved_intent) = resolved_intent {
7160 for transition in transitions {
7161 if transition.intent.as_deref() == Some(resolved_intent) {
7162 return Some(TransitionCandidate::new(
7163 current_state,
7164 transition.clone(),
7165 Self::transition_reason(transition),
7166 ));
7167 }
7168 }
7169 }
7170
7171 None
7172 }
7173
7174 async fn commit_transition_candidate(&self, candidate: &TransitionCandidate) -> Result<bool> {
7176 self.commit_transition_target(&candidate.from_state, candidate.target(), &candidate.reason)
7177 .await
7178 }
7179
7180 async fn approve_transition_target(&self, from_state: &str, target: &str) -> Result<bool> {
7182 let approved = self.check_state_hitl(Some(from_state), target).await?;
7183 if !approved {
7184 info!(to = %target, "State transition rejected by HITL");
7185 }
7186 Ok(approved)
7187 }
7188
7189 async fn apply_transition_target(
7191 &self,
7192 from_state: &str,
7193 target: &str,
7194 reason: &str,
7195 staged: Option<&HashMap<String, Value>>,
7196 ) -> Result<bool> {
7197 let Some(ref sm) = self.state_machine else {
7198 return Ok(false);
7199 };
7200 let claim_admission = self.disambiguation_admission.write().await;
7201 if sm.current() != from_state {
7202 return Ok(false);
7203 }
7204 let Some(reservation) = self.reserve_state_transition() else {
7205 return Ok(false);
7206 };
7207 let expected_state_generation = sm.generation();
7208 let expected_disambiguation_epoch = self.disambiguation_epoch.load(Ordering::SeqCst);
7209 let history_before = sm.history();
7210 drop(claim_admission);
7211
7212 self.execute_state_exit_actions(from_state).await;
7213
7214 let admission = self.disambiguation_admission.write().await;
7215 if sm.current() != from_state
7216 || sm.generation() != expected_state_generation
7217 || self.disambiguation_epoch.load(Ordering::SeqCst) != expected_disambiguation_epoch
7218 {
7219 return Ok(false);
7220 }
7221 sm.transition_to(target, reason)?;
7222 self.invalidate_pending_confirmation("state_transition")
7223 .await;
7224 sm.reset_no_transition();
7225 if let Some(staged) = staged {
7226 self.commit_staged_context_writes(staged);
7227 }
7228 let entered = sm.current();
7229 let is_reentry = Self::state_was_previously_entered(&entered, from_state, &history_before);
7230 drop(admission);
7231
7232 self.execute_state_enter_actions(&entered, is_reentry).await;
7233 drop(reservation);
7234 self.hooks
7235 .on_state_transition(Some(from_state), &entered, reason)
7236 .await;
7237 info!(from = %from_state, to = %entered, "State transition");
7238 Ok(true)
7239 }
7240
7241 async fn commit_transition_target(
7243 &self,
7244 from_state: &str,
7245 target: &str,
7246 reason: &str,
7247 ) -> Result<bool> {
7248 if !self.approve_transition_target(from_state, target).await? {
7249 return Ok(false);
7250 }
7251 self.apply_transition_target(from_state, target, reason, None)
7252 .await
7253 }
7254
7255 async fn apply_pre_response_transition_candidate(
7257 &self,
7258 candidate: &TransitionCandidate,
7259 staged: &HashMap<String, Value>,
7260 processed_input: &str,
7261 ) -> Result<bool> {
7262 self.commit_root_user_message(processed_input).await?;
7263 self.apply_transition_target(
7264 &candidate.from_state,
7265 candidate.target(),
7266 &candidate.reason,
7267 Some(staged),
7268 )
7269 .await
7270 }
7271
7272 async fn commit_pre_response_transition_candidate(
7274 &self,
7275 candidate: &TransitionCandidate,
7276 staged: &HashMap<String, Value>,
7277 processed_input: &str,
7278 ) -> Result<bool> {
7279 if !self
7280 .approve_transition_target(&candidate.from_state, candidate.target())
7281 .await?
7282 {
7283 return Ok(false);
7284 }
7285 self.apply_pre_response_transition_candidate(candidate, staged, processed_input)
7286 .await
7287 }
7288
7289 async fn handle_transition_miss(&self, current_state: &str) -> Result<bool> {
7291 let Some(ref sm) = self.state_machine else {
7292 return Ok(false);
7293 };
7294 sm.increment_no_transition();
7295 let Some(fallback) = sm.check_fallback() else {
7296 return Ok(false);
7297 };
7298 self.commit_transition_target(current_state, &fallback, "fallback after no transitions")
7299 .await
7300 }
7301
7302 async fn evaluate_transitions(&self, user_message: &str, response: &str) -> Result<bool> {
7304 let Some((transitions, current_state)) = self.transitions_available_for_commit() else {
7305 return Ok(false);
7306 };
7307 if transitions.is_empty() {
7308 return Ok(false);
7309 }
7310 if let Some(candidate) = self
7311 .select_transition_candidate(user_message, response)
7312 .await?
7313 {
7314 return self.commit_transition_candidate(&candidate).await;
7315 }
7316 self.handle_transition_miss(¤t_state).await
7317 }
7318
7319 async fn try_pre_response_transition(
7321 &self,
7322 processed_input: &str,
7323 ) -> Result<Option<AgentResponse>> {
7324 let optimization = &self.runtime_config.optimization;
7325 if !optimization.enabled || !optimization.pre_response_deterministic_transitions {
7326 return Ok(None);
7327 }
7328 let Some((transitions, current_state)) = self.transitions_available_for_commit() else {
7329 return Ok(None);
7330 };
7331 let eligible: Vec<Transition> = transitions
7332 .into_iter()
7333 .filter(|transition| !transition.requires_response)
7334 .filter(|transition| matches!(transition.timing, TransitionTiming::PreResponse))
7335 .collect();
7336 if eligible.is_empty() {
7337 return Ok(None);
7338 }
7339
7340 let empty_staged = HashMap::new();
7341 let mut extracted_staged: Option<HashMap<String, Value>> = None;
7342 let mut selected: Option<(TransitionCandidate, HashMap<String, Value>)> = None;
7343
7344 for transition in &eligible {
7345 let use_extractors = optimization.pre_response_extractors || transition.run_extractors;
7346 let staged_for_eval = if use_extractors {
7347 if extracted_staged.is_none() {
7348 extracted_staged =
7349 Some(self.run_context_extractors_staged(processed_input).await);
7350 }
7351 extracted_staged.as_ref().unwrap_or(&empty_staged)
7352 } else {
7353 &empty_staged
7354 };
7355
7356 if let Some(candidate) = self.select_deterministic_transition_candidate(
7357 processed_input,
7358 ¤t_state,
7359 std::slice::from_ref(transition),
7360 staged_for_eval,
7361 ) {
7362 let staged_for_commit = if use_extractors {
7363 staged_for_eval.clone()
7364 } else {
7365 HashMap::new()
7366 };
7367 selected = Some((candidate, staged_for_commit));
7368 break;
7369 }
7370 }
7371
7372 let Some((candidate, staged)) = selected else {
7373 return Ok(None);
7374 };
7375
7376 if !self
7377 .commit_pre_response_transition_candidate(&candidate, &staged, processed_input)
7378 .await?
7379 {
7380 return Ok(None);
7381 }
7382 self.redispatch_current_state(processed_input)
7383 .await
7384 .map(Some)
7385 }
7386
7387 async fn try_speculative_branches(
7392 &self,
7393 processed_input: &str,
7394 input_context: &HashMap<String, Value>,
7395 ) -> Result<Option<AgentResponse>> {
7396 let optimization = &self.runtime_config.optimization;
7397 if !optimization.enabled {
7398 return Ok(None);
7399 }
7400
7401 let effective_reasoning_mode = self.get_effective_reasoning_config().mode.clone();
7402 if !matches!(
7403 effective_reasoning_mode,
7404 ReasoningMode::None | ReasoningMode::Auto
7405 ) {
7406 return Ok(None);
7407 }
7408
7409 let mut transition_enabled =
7410 optimization.speculative_state_transitions && self.has_parallel_transition_candidates();
7411 let mut skill_enabled = optimization.speculative_skill_routing
7412 && self.skill_router.is_some()
7413 && self.pending_skill_id.read().is_none();
7414 let mut reasoning_enabled = optimization.speculative_reasoning_auto
7415 && matches!(effective_reasoning_mode, ReasoningMode::Auto);
7416
7417 if matches!(effective_reasoning_mode, ReasoningMode::Auto)
7418 && (!reasoning_enabled || optimization.max_speculative_llm_calls_per_turn < 2)
7419 {
7420 return Ok(None);
7421 }
7422
7423 if !transition_enabled && !skill_enabled && !reasoning_enabled {
7424 return Ok(None);
7425 }
7426
7427 let mut optional_slots = optimization.max_parallel_runtime_tasks.saturating_sub(1);
7428 let mut speculative_call_slots = optimization
7429 .max_speculative_llm_calls_per_turn
7430 .saturating_sub(1);
7431 if reasoning_enabled {
7432 if optional_slots == 0 || speculative_call_slots == 0 {
7433 return Ok(None);
7434 }
7435 optional_slots -= 1;
7436 speculative_call_slots -= 1;
7437 }
7438 if transition_enabled {
7439 if optional_slots == 0 {
7440 transition_enabled = false;
7441 } else {
7442 optional_slots -= 1;
7443 }
7444 }
7445 if skill_enabled && (optional_slots == 0 || speculative_call_slots == 0) {
7446 skill_enabled = false;
7447 }
7448
7449 if !transition_enabled && !skill_enabled && !reasoning_enabled {
7450 return Ok(None);
7451 }
7452
7453 let main_kind = if transition_enabled {
7454 RuntimeOptimizationKind::ParallelStateTransition
7455 } else if skill_enabled {
7456 RuntimeOptimizationKind::SpeculativeSkillRouting
7457 } else {
7458 RuntimeOptimizationKind::SpeculativeReasoningAuto
7459 };
7460 if !self.reserve_active_speculative_llm_call(main_kind) {
7461 return Ok(None);
7462 }
7463
7464 let mut branch_set = ScheduledBranchSet::new(optimization.max_parallel_runtime_tasks)?;
7465 let main_branch = RuntimeBranch::new(
7466 RuntimeTaskPurpose::MainResponse,
7467 main_kind,
7468 RuntimeTaskPriority::Normal,
7469 RuntimeCommitBehavior::FinalResponse,
7470 );
7471 let transition_branch = RuntimeBranch::new(
7472 RuntimeTaskPurpose::StateTransition,
7473 RuntimeOptimizationKind::ParallelStateTransition,
7474 RuntimeTaskPriority::Critical,
7475 RuntimeCommitBehavior::TransitionDecision,
7476 );
7477 let skill_branch = RuntimeBranch::new(
7478 RuntimeTaskPurpose::SkillRouting,
7479 RuntimeOptimizationKind::SpeculativeSkillRouting,
7480 RuntimeTaskPriority::High,
7481 RuntimeCommitBehavior::SkillSelection,
7482 );
7483 let reasoning_branch = RuntimeBranch::new(
7484 RuntimeTaskPurpose::ReasoningJudge,
7485 RuntimeOptimizationKind::SpeculativeReasoningAuto,
7486 RuntimeTaskPriority::Normal,
7487 RuntimeCommitBehavior::ReasoningDecision,
7488 );
7489 let main_id = main_branch.branch_id();
7490 let transition_id = transition_branch.branch_id();
7491 let skill_id = skill_branch.branch_id();
7492 let reasoning_id = reasoning_branch.branch_id();
7493
7494 let main_id_for_future = main_id.clone();
7495 if !branch_set.schedule(
7496 main_branch,
7497 Box::pin(async move {
7498 match crate::optimization::observability::with_branch_observation(
7499 &main_id_for_future,
7500 main_kind,
7501 RuntimeCommitBehavior::FinalResponse,
7502 self.generate_main_response_draft(processed_input, &ReasoningMode::None),
7503 )
7504 .await
7505 {
7506 Ok(draft) => RuntimeBranchResult::MainDraft(draft),
7507 Err(error) => RuntimeBranchResult::Failed(error),
7508 }
7509 }),
7510 ) {
7511 return Ok(None);
7512 }
7513
7514 if transition_enabled {
7515 let transition_id_for_future = transition_id.clone();
7516 if !branch_set.schedule(
7517 transition_branch,
7518 Box::pin(async move {
7519 match crate::optimization::observability::with_branch_observation(
7520 &transition_id_for_future,
7521 RuntimeOptimizationKind::ParallelStateTransition,
7522 RuntimeCommitBehavior::TransitionDecision,
7523 self.select_parallel_transition_candidate(processed_input),
7524 )
7525 .await
7526 {
7527 Ok(ParallelTransitionSelection::Candidate(candidate)) => {
7528 RuntimeBranchResult::Transition(Some(candidate))
7529 }
7530 Ok(ParallelTransitionSelection::NoMatch) => {
7531 RuntimeBranchResult::Transition(None)
7532 }
7533 Ok(ParallelTransitionSelection::ReservationExhausted) => {
7534 RuntimeBranchResult::Cancelled
7535 }
7536 Err(error) => RuntimeBranchResult::Failed(error),
7537 }
7538 }),
7539 ) {
7540 transition_enabled = false;
7541 }
7542 }
7543
7544 if skill_enabled {
7545 let skill_id_for_future = skill_id.clone();
7546 if !branch_set.schedule(
7547 skill_branch,
7548 Box::pin(async move {
7549 if !self.reserve_active_speculative_llm_call(
7550 RuntimeOptimizationKind::SpeculativeSkillRouting,
7551 ) {
7552 return RuntimeBranchResult::Cancelled;
7553 }
7554 match crate::optimization::observability::with_branch_observation(
7555 &skill_id_for_future,
7556 RuntimeOptimizationKind::SpeculativeSkillRouting,
7557 RuntimeCommitBehavior::SkillSelection,
7558 self.select_skill_candidate(processed_input),
7559 )
7560 .await
7561 {
7562 Ok(candidate) => RuntimeBranchResult::Skill(candidate),
7563 Err(error) => RuntimeBranchResult::Failed(error),
7564 }
7565 }),
7566 ) {
7567 skill_enabled = false;
7568 }
7569 }
7570
7571 if reasoning_enabled {
7572 let reasoning_id_for_future = reasoning_id.clone();
7573 if !branch_set.schedule(
7574 reasoning_branch,
7575 Box::pin(async move {
7576 if !self.reserve_active_speculative_llm_call(
7577 RuntimeOptimizationKind::SpeculativeReasoningAuto,
7578 ) {
7579 return RuntimeBranchResult::Cancelled;
7580 }
7581 match crate::optimization::observability::with_branch_observation(
7582 &reasoning_id_for_future,
7583 RuntimeOptimizationKind::SpeculativeReasoningAuto,
7584 RuntimeCommitBehavior::ReasoningDecision,
7585 self.determine_reasoning_mode_strict(processed_input),
7586 )
7587 .await
7588 {
7589 Ok(mode) => RuntimeBranchResult::Reasoning(mode),
7590 Err(error) => RuntimeBranchResult::Failed(error),
7591 }
7592 }),
7593 ) {
7594 reasoning_enabled = false;
7595 }
7596 }
7597
7598 if matches!(effective_reasoning_mode, ReasoningMode::Auto) && !reasoning_enabled {
7599 self.finalize_pending_branches(branch_set.cancel_pending());
7600 return Ok(None);
7601 }
7602
7603 if !transition_enabled && !skill_enabled && !reasoning_enabled {
7604 self.finalize_pending_branches(branch_set.cancel_pending());
7605 return Ok(None);
7606 }
7607
7608 let mut main_pending = true;
7609 let mut skill_pending = skill_enabled;
7610 let mut reasoning_pending = reasoning_enabled;
7611 let mut transition_finalized = !transition_enabled;
7612 let mut skill_finalized = !skill_enabled && self.skill_router.is_none();
7615 let mut reasoning_finalized = !reasoning_enabled;
7616 let mut main_result: Option<Result<MainResponseDraft>> = None;
7617 let mut transition_candidate: Option<TransitionCandidate> = None;
7618 let mut skill_candidate: Option<SkillCandidate> = None;
7619 let mut reasoning_decision: Option<ReasoningMode> = None;
7620 let mut transition_fallback_required = false;
7621 let mut skill_fallback_required = false;
7622 let mut reasoning_fallback_required = false;
7623
7624 loop {
7625 if let Some(candidate) = transition_candidate.take() {
7626 if self
7627 .approve_transition_target(&candidate.from_state, candidate.target())
7628 .await?
7629 {
7630 self.finalize_pending_branches(branch_set.cancel_pending());
7632 if !main_pending {
7633 self.finalize_branch_loss(
7634 &main_id,
7635 main_kind,
7636 RuntimeCommitBehavior::FinalResponse,
7637 false,
7638 main_result.as_ref().map(|result| result.is_err()),
7639 );
7640 }
7641 if skill_enabled && !skill_pending {
7642 self.finalize_branch_loss(
7643 &skill_id,
7644 RuntimeOptimizationKind::SpeculativeSkillRouting,
7645 RuntimeCommitBehavior::SkillSelection,
7646 false,
7647 Some(false),
7648 );
7649 }
7650 if reasoning_enabled && !reasoning_pending {
7651 self.finalize_branch_loss(
7652 &reasoning_id,
7653 RuntimeOptimizationKind::SpeculativeReasoningAuto,
7654 RuntimeCommitBehavior::ReasoningDecision,
7655 false,
7656 Some(false),
7657 );
7658 }
7659 if !self
7660 .apply_pre_response_transition_candidate(
7661 &candidate,
7662 &HashMap::new(),
7663 processed_input,
7664 )
7665 .await?
7666 {
7667 self.finalize_optional_branch(
7668 &transition_id,
7669 RuntimeOptimizationKind::ParallelStateTransition,
7670 RuntimeCommitBehavior::TransitionDecision,
7671 "discarded",
7672 false,
7673 );
7674 return Ok(None);
7675 }
7676 self.finalize_optional_branch(
7677 &transition_id,
7678 RuntimeOptimizationKind::ParallelStateTransition,
7679 RuntimeCommitBehavior::TransitionDecision,
7680 "committed",
7681 true,
7682 );
7683 return self
7684 .redispatch_current_state(processed_input)
7685 .await
7686 .map(Some);
7687 }
7688 self.finalize_optional_branch(
7689 &transition_id,
7690 RuntimeOptimizationKind::ParallelStateTransition,
7691 RuntimeCommitBehavior::TransitionDecision,
7692 "discarded",
7693 false,
7694 );
7695 transition_finalized = true;
7696 }
7697
7698 if transition_finalized
7707 && !skill_finalized
7708 && !skill_enabled
7709 && self.skill_router.is_some()
7710 {
7711 match self.select_skill_candidate(processed_input).await {
7712 Ok(Some(candidate)) => skill_candidate = Some(candidate),
7713 Ok(None) => {}
7714 Err(error) => {
7715 self.finalize_pending_branches(branch_set.cancel_pending());
7717 return Err(error);
7718 }
7719 }
7720 skill_finalized = true;
7721 }
7722
7723 if transition_finalized && skill_candidate.is_some() {
7724 let candidate = skill_candidate.take().unwrap();
7725 if skill_enabled {
7727 self.finalize_optional_branch(
7728 &skill_id,
7729 RuntimeOptimizationKind::SpeculativeSkillRouting,
7730 RuntimeCommitBehavior::SkillSelection,
7731 "committed",
7732 true,
7733 );
7734 }
7735 if !main_pending {
7736 self.finalize_branch_loss(
7737 &main_id,
7738 main_kind,
7739 RuntimeCommitBehavior::FinalResponse,
7740 false,
7741 main_result.as_ref().map(|result| result.is_err()),
7742 );
7743 }
7744 if reasoning_enabled && !reasoning_pending {
7745 self.finalize_branch_loss(
7746 &reasoning_id,
7747 RuntimeOptimizationKind::SpeculativeReasoningAuto,
7748 RuntimeCommitBehavior::ReasoningDecision,
7749 false,
7750 Some(false),
7751 );
7752 }
7753 self.finalize_pending_branches(branch_set.cancel_pending());
7754 return self
7755 .commit_winning_skill_candidate(candidate, processed_input, input_context)
7756 .await;
7757 }
7758
7759 if transition_finalized
7760 && skill_finalized
7761 && let Some(reasoning_mode) = reasoning_decision.take()
7762 {
7763 if !matches!(reasoning_mode, ReasoningMode::None) {
7764 self.finalize_optional_branch(
7765 &reasoning_id,
7766 RuntimeOptimizationKind::SpeculativeReasoningAuto,
7767 RuntimeCommitBehavior::ReasoningDecision,
7768 "committed",
7769 true,
7770 );
7771 if !main_pending {
7772 self.finalize_branch_loss(
7773 &main_id,
7774 main_kind,
7775 RuntimeCommitBehavior::FinalResponse,
7776 false,
7777 main_result.as_ref().map(|result| result.is_err()),
7778 );
7779 }
7780 self.finalize_pending_branches(branch_set.cancel_pending());
7781 self.commit_root_user_message(processed_input).await?;
7782 return if matches!(reasoning_mode, ReasoningMode::PlanAndExecute) {
7783 self.handle_plan_and_execute(processed_input, input_context, true)
7784 .await
7785 .map(Some)
7786 } else {
7787 self.run_committed_response_loop_with_reasoning(
7788 processed_input,
7789 input_context,
7790 reasoning_mode,
7791 true,
7792 )
7793 .await
7794 .map(Some)
7795 };
7796 }
7797 self.finalize_optional_branch(
7798 &reasoning_id,
7799 RuntimeOptimizationKind::SpeculativeReasoningAuto,
7800 RuntimeCommitBehavior::ReasoningDecision,
7801 "committed",
7802 true,
7803 );
7804 reasoning_finalized = true;
7805 }
7806
7807 if transition_finalized && skill_finalized && reasoning_finalized {
7808 if transition_fallback_required
7809 || skill_fallback_required
7810 || reasoning_fallback_required
7811 {
7812 if !main_pending {
7813 self.finalize_branch_loss(
7814 &main_id,
7815 main_kind,
7816 RuntimeCommitBehavior::FinalResponse,
7817 false,
7818 main_result.as_ref().map(|result| result.is_err()),
7819 );
7820 }
7821 self.finalize_pending_branches(branch_set.cancel_pending());
7822 return Ok(None);
7823 }
7824
7825 if let Some(result) = main_result.take() {
7826 let draft = match result {
7827 Ok(draft) => draft,
7828 Err(error) => {
7829 self.finalize_optional_branch(
7830 &main_id,
7831 main_kind,
7832 RuntimeCommitBehavior::FinalResponse,
7833 "failed",
7834 false,
7835 );
7836 self.finalize_pending_branches(branch_set.cancel_pending());
7837 return Err(error);
7838 }
7839 };
7840 self.finalize_optional_branch(
7841 &main_id,
7842 main_kind,
7843 RuntimeCommitBehavior::FinalResponse,
7844 "committed",
7845 true,
7846 );
7847 self.finalize_pending_branches(branch_set.cancel_pending());
7848 return self
7849 .commit_main_response_draft(
7850 processed_input,
7851 input_context,
7852 draft,
7853 ReasoningMode::None,
7854 reasoning_enabled,
7855 )
7856 .await
7857 .map(Some);
7858 }
7859 }
7860
7861 if branch_set.is_empty() {
7862 return Ok(None);
7863 }
7864
7865 let Some(outcome) = branch_set.next_completed().await else {
7866 return Ok(None);
7867 };
7868 let branch_id = outcome.branch.branch_id();
7869 match outcome.result {
7870 RuntimeBranchResult::MainDraft(draft) => {
7871 main_pending = false;
7872 main_result = Some(Ok(draft));
7873 }
7874 RuntimeBranchResult::Transition(candidate) => {
7875 if let Some(candidate) = candidate {
7876 transition_candidate = Some(candidate);
7877 } else {
7878 self.finalize_optional_branch(
7879 &transition_id,
7880 RuntimeOptimizationKind::ParallelStateTransition,
7881 RuntimeCommitBehavior::TransitionDecision,
7882 "discarded",
7883 false,
7884 );
7885 transition_finalized = true;
7886 }
7887 }
7888 RuntimeBranchResult::Skill(candidate) => {
7889 skill_pending = false;
7890 if let Some(candidate) = candidate {
7891 skill_candidate = Some(candidate);
7892 } else {
7893 self.finalize_optional_branch(
7894 &skill_id,
7895 RuntimeOptimizationKind::SpeculativeSkillRouting,
7896 RuntimeCommitBehavior::SkillSelection,
7897 "discarded",
7898 false,
7899 );
7900 skill_finalized = true;
7901 }
7902 }
7903 RuntimeBranchResult::Reasoning(mode) => {
7904 reasoning_pending = false;
7905 reasoning_decision = Some(mode);
7906 }
7907 RuntimeBranchResult::Failed(error) => {
7908 if branch_id == main_id {
7909 main_pending = false;
7910 main_result = Some(Err(error));
7911 } else if branch_id == transition_id {
7912 self.finalize_optional_branch(
7913 &transition_id,
7914 RuntimeOptimizationKind::ParallelStateTransition,
7915 RuntimeCommitBehavior::TransitionDecision,
7916 "failed",
7917 false,
7918 );
7919 transition_finalized = true;
7920 } else if branch_id == skill_id {
7921 skill_pending = false;
7922 self.finalize_optional_branch(
7923 &skill_id,
7924 RuntimeOptimizationKind::SpeculativeSkillRouting,
7925 RuntimeCommitBehavior::SkillSelection,
7926 "failed",
7927 false,
7928 );
7929 skill_finalized = true;
7930 } else if branch_id == reasoning_id {
7931 reasoning_pending = false;
7932 self.finalize_optional_branch(
7933 &reasoning_id,
7934 RuntimeOptimizationKind::SpeculativeReasoningAuto,
7935 RuntimeCommitBehavior::ReasoningDecision,
7936 "failed",
7937 false,
7938 );
7939 reasoning_finalized = true;
7940 }
7941 }
7942 RuntimeBranchResult::Cancelled => {
7943 self.finalize_optional_branch(
7944 &branch_id,
7945 outcome.branch.optimization,
7946 outcome.branch.commit_behavior,
7947 "cancelled",
7948 false,
7949 );
7950 if branch_id == main_id {
7951 main_pending = false;
7952 main_result =
7953 Some(Err(AgentError::Other("main branch cancelled".to_string())));
7954 } else if branch_id == transition_id {
7955 transition_finalized = true;
7956 transition_fallback_required = true;
7957 } else if branch_id == skill_id {
7958 skill_pending = false;
7959 skill_finalized = true;
7960 skill_fallback_required = true;
7961 } else if branch_id == reasoning_id {
7962 reasoning_pending = false;
7963 reasoning_finalized = true;
7964 reasoning_fallback_required = true;
7965 }
7966 }
7967 }
7968 }
7969 }
7970
7971 fn finalize_pending_branches(&self, branches: Vec<RuntimeBranch>) {
7972 for branch in branches {
7973 self.finalize_optional_branch(
7974 &branch.branch_id(),
7975 branch.optimization,
7976 branch.commit_behavior,
7977 "cancelled",
7978 false,
7979 );
7980 }
7981 }
7982
7983 fn finalize_branch_loss(
7988 &self,
7989 branch_id: &str,
7990 optimization: RuntimeOptimizationKind,
7991 commit_behavior: RuntimeCommitBehavior,
7992 pending: bool,
7993 completed_failed: Option<bool>,
7994 ) {
7995 let status = if pending {
7996 "cancelled"
7997 } else if completed_failed.unwrap_or(false) {
7998 "failed"
7999 } else {
8000 "discarded"
8001 };
8002 self.finalize_optional_branch(branch_id, optimization, commit_behavior, status, false);
8003 }
8004
8005 fn finalize_optional_branch(
8010 &self,
8011 branch_id: &str,
8012 optimization: RuntimeOptimizationKind,
8013 commit_behavior: RuntimeCommitBehavior,
8014 status: &str,
8015 winner: bool,
8016 ) {
8017 crate::optimization::observability::finalize_branch(
8018 self.observability_manager.as_ref(),
8019 branch_id,
8020 status,
8021 winner,
8022 optimization,
8023 commit_behavior,
8024 );
8025 }
8026
8027 fn has_parallel_transition_candidates(&self) -> bool {
8032 self.transitions_available_for_commit()
8033 .map(|(transitions, _)| {
8034 transitions
8035 .iter()
8036 .any(|transition| matches!(transition.timing, TransitionTiming::Parallel))
8037 })
8038 .unwrap_or(false)
8039 }
8040
8041 async fn select_parallel_transition_candidate(
8046 &self,
8047 processed_input: &str,
8048 ) -> Result<ParallelTransitionSelection> {
8049 let Some((transitions, current_state)) = self.transitions_available_for_commit() else {
8050 return Ok(ParallelTransitionSelection::NoMatch);
8051 };
8052 let parallel: Vec<Transition> = transitions
8053 .into_iter()
8054 .filter(|transition| matches!(transition.timing, TransitionTiming::Parallel))
8055 .filter(|transition| !transition.requires_response)
8056 .collect();
8057 if parallel.is_empty() {
8058 return Ok(ParallelTransitionSelection::NoMatch);
8059 }
8060 let empty_staged = HashMap::new();
8061 if let Some(candidate) = self.select_deterministic_transition_candidate(
8062 processed_input,
8063 ¤t_state,
8064 ¶llel,
8065 &empty_staged,
8066 ) {
8067 return Ok(ParallelTransitionSelection::Candidate(candidate));
8068 }
8069 let when_transitions: Vec<(usize, &Transition)> = parallel
8070 .iter()
8071 .enumerate()
8072 .filter(|(_, transition)| !transition.when.trim().is_empty())
8073 .collect();
8074 if when_transitions.is_empty() {
8075 return Ok(ParallelTransitionSelection::NoMatch);
8076 }
8077 let llm = self
8078 .llm_registry
8079 .router()
8080 .or_else(|_| self.llm_registry.default())
8081 .map_err(|e| AgentError::Config(e.to_string()))?;
8082 let conditions = when_transitions
8083 .iter()
8084 .enumerate()
8085 .map(|(display_idx, (_, transition))| {
8086 format!("{}. {}", display_idx + 1, transition.when)
8087 })
8088 .collect::<Vec<_>>()
8089 .join("\n");
8090 if !self
8091 .reserve_active_speculative_llm_call(RuntimeOptimizationKind::ParallelStateTransition)
8092 {
8093 return Ok(ParallelTransitionSelection::ReservationExhausted);
8094 }
8095 let context_preview = self.branch_context_preview();
8096 let prompt = format!(
8097 "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-{}).",
8098 current_state,
8099 processed_input,
8100 context_preview,
8101 conditions,
8102 when_transitions.len()
8103 );
8104 let response = self
8105 .observe_purpose(
8106 ObservationPurpose::StateTransitionEvaluation,
8107 llm.complete(&[ChatMessage::user(prompt)], None),
8108 )
8109 .await
8110 .map_err(|e| AgentError::LLM(e.to_string()))?;
8111 let choice = response.content.trim().parse::<usize>().unwrap_or(0);
8112 if choice == 0 || choice > when_transitions.len() {
8113 return Ok(ParallelTransitionSelection::NoMatch);
8114 }
8115 let transition = when_transitions[choice - 1].1.clone();
8116 Ok(ParallelTransitionSelection::Candidate(
8117 TransitionCandidate::new(
8118 current_state,
8119 transition.clone(),
8120 Self::transition_reason(&transition),
8121 ),
8122 ))
8123 }
8124
8125 async fn redispatch_current_state(&self, processed_input: &str) -> Result<AgentResponse> {
8127 const MAX_REDISPATCH_DEPTH: u32 = 3;
8128 let current_depth = *self.redispatch_depth.read();
8129 if current_depth >= MAX_REDISPATCH_DEPTH {
8130 warn!(depth = current_depth, "Re-dispatch depth limit reached");
8131 let response = AgentResponse::new("");
8132 self.finish_turn_if_root(&response).await?;
8133 return Ok(response);
8134 }
8135 *self.redispatch_depth.write() += 1;
8136 if let Some(context) = self.active_turn_context.write().as_mut() {
8137 context.enter_redispatch();
8138 }
8139 let result = Box::pin(self.run_loop_internal(processed_input)).await;
8140 *self.redispatch_depth.write() -= 1;
8141 if let Some(context) = self.active_turn_context.write().as_mut() {
8142 context.exit_redispatch();
8143 }
8144 let response = result?;
8145 self.finish_turn_if_root(&response).await?;
8146 Ok(response)
8147 }
8148
8149 async fn finish_turn_if_root(&self, response: &AgentResponse) -> Result<()> {
8151 if *self.redispatch_depth.read() == 0 {
8152 self.post_turn_session_lifecycle().await?;
8153 if let Some(context) = self.active_turn_context.write().as_mut() {
8154 context.mark_post_turn_lifecycle_completed();
8155 }
8156 self.hooks.on_response(response).await;
8157 self.end_root_turn();
8158 }
8159 Ok(())
8160 }
8161
8162 async fn execute_state_exit_actions(&self, state_path: &str) {
8164 if let Some(ref sm) = self.state_machine
8165 && let Some(def) = sm.get_definition(state_path)
8166 && !def.on_exit.is_empty()
8167 {
8168 debug!(state = %state_path, count = def.on_exit.len(), "Executing on_exit actions");
8169 self.execute_state_actions(&def.on_exit).await;
8170 }
8171 }
8172
8173 fn state_was_previously_entered(
8175 state_path: &str,
8176 from_state: &str,
8177 history_before: &[StateTransitionEvent],
8178 ) -> bool {
8179 state_path == from_state
8180 || history_before
8181 .iter()
8182 .any(|event| event.from == state_path || event.to == state_path)
8183 }
8184
8185 async fn execute_state_enter_actions(&self, state_path: &str, is_reentry: bool) {
8187 if let Some(ref sm) = self.state_machine
8188 && let Some(def) = sm.get_definition(state_path)
8189 {
8190 if is_reentry && !def.on_reenter.is_empty() {
8191 debug!(state = %state_path, count = def.on_reenter.len(), "Executing on_reenter actions");
8192 self.execute_state_actions(&def.on_reenter).await;
8193 } else if !def.on_enter.is_empty() {
8194 debug!(state = %state_path, count = def.on_enter.len(), "Executing on_enter actions");
8195 self.execute_state_actions(&def.on_enter).await;
8196 }
8197 }
8198 }
8199
8200 async fn execute_state_actions(&self, actions: &[StateAction]) {
8202 for (action_index, action) in actions.iter().enumerate() {
8203 match action {
8204 StateAction::Tool { tool, args } => {
8205 let raw_args = args.clone().unwrap_or(Value::Object(Default::default()));
8206 let args_value = self.render_action_args(&raw_args);
8207 let state = self.state_machine.as_ref().map(|sm| sm.current());
8208 let request = ToolExecutionRequest::new(
8209 uuid::Uuid::new_v4().to_string(),
8210 tool.clone(),
8211 args_value,
8212 ToolCallSource::StateAction {
8213 state,
8214 action_index,
8215 },
8216 );
8217 match self.execute_tool_record(request).await {
8218 Ok(record) if record.success => {
8219 debug!(tool = %record.canonical_id, "State action: tool executed");
8220 let _ = self.context_manager.set(
8221 "last_tool_result",
8222 serde_json::Value::String(record.model_output_string()),
8223 );
8224 let _ = self.context_manager.set(
8225 "last_tool_record",
8226 serde_json::to_value(record).unwrap_or(Value::Null),
8227 );
8228 }
8229 Ok(record) => {
8230 warn!(tool = %record.canonical_id, error = %record.output, "State action: tool failed");
8231 }
8232 Err(e) => {
8233 warn!(tool = %tool, error = %e, "State action: tool failed")
8234 }
8235 }
8236 }
8237 StateAction::Skill { skill } => {
8238 if let Some(ref executor) = self.skill_executor {
8239 if let Some(def) = self.skills.iter().find(|s| s.id == *skill) {
8240 match executor
8241 .execute_with_invoker(def, "", serde_json::json!({}), self)
8242 .await
8243 {
8244 Ok(_) => debug!(skill = %skill, "State action: skill executed"),
8245 Err(e) => {
8246 warn!(skill = %skill, error = %e, "State action: skill failed")
8247 }
8248 }
8249 } else {
8250 warn!(skill = %skill, "State action: skill not found");
8251 }
8252 }
8253 }
8254 StateAction::SetContext { set_context } => {
8255 for (key, value) in set_context {
8256 if let Err(e) = self.context_manager.set(key, value.clone()) {
8257 warn!(key = %key, error = %e, "State action: set_context failed");
8258 } else {
8259 debug!(key = %key, "State action: context set");
8260 }
8261 }
8262 }
8263 StateAction::Prompt {
8264 prompt,
8265 llm,
8266 store_as,
8267 } => {
8268 let llm_result = if let Some(alias) = llm {
8269 self.llm_registry.get(alias)
8270 } else {
8271 self.llm_registry.default()
8272 };
8273 match llm_result {
8274 Ok(llm_provider) => {
8275 let context = self.build_context_with_overlays();
8277 let rendered_prompt = self
8278 .template_renderer
8279 .render(prompt, &context)
8280 .unwrap_or_else(|_| prompt.clone());
8281 let recent =
8282 self.memory.get_messages(Some(5)).await.unwrap_or_default();
8283 let mut messages: Vec<ChatMessage> = recent;
8284 messages.push(ChatMessage::user(&rendered_prompt));
8285 match self
8286 .observe_purpose(
8287 ObservationPurpose::StateAction,
8288 llm_provider.complete(&messages, None),
8289 )
8290 .await
8291 {
8292 Ok(response) => {
8293 if let Some(key) = store_as {
8294 let _ = self
8295 .context_manager
8296 .set(key, Value::String(response.content));
8297 debug!(key = %key, "State action: prompt result stored");
8298 }
8299 }
8300 Err(e) => {
8301 warn!(error = %e, "State action: prompt LLM call failed");
8302 }
8303 }
8304 }
8305 Err(e) => {
8306 warn!(error = %e, "State action: LLM not found for prompt");
8307 }
8308 }
8309 }
8310 }
8311 }
8312 }
8313
8314 async fn run_context_extractors_staged(&self, user_message: &str) -> HashMap<String, Value> {
8315 let extractors = match &self.state_machine {
8316 Some(sm) => match sm.current_definition() {
8317 Some(def) if !def.extract.is_empty() => def.extract.clone(),
8318 _ => return HashMap::new(),
8319 },
8320 None => return HashMap::new(),
8321 };
8322
8323 let mut staged = HashMap::new();
8324 for extractor in &extractors {
8325 let prompt = if let Some(ref custom) = extractor.llm_extract {
8326 format!(
8327 "User message:\n\"{}\"\n\nInstruction:\n{}",
8328 user_message, custom
8329 )
8330 } else if let Some(ref desc) = extractor.description {
8331 format!(
8332 "From the following message, extract: {}\n\n\
8333 Message: \"{}\"\n\n\
8334 If the information is present, return ONLY the extracted value.\n\
8335 If NOT present, return exactly: __NONE__",
8336 desc, user_message
8337 )
8338 } else {
8339 continue;
8340 };
8341
8342 let llm = match self
8343 .llm_registry
8344 .get(&extractor.llm)
8345 .or_else(|_| self.llm_registry.get("router"))
8346 .or_else(|_| self.llm_registry.get("default"))
8347 {
8348 Ok(llm) => llm,
8349 Err(e) => {
8350 warn!(key = %extractor.key, error = %e, "Extractor LLM not found");
8351 continue;
8352 }
8353 };
8354
8355 let messages = vec![ChatMessage::user(&prompt)];
8356 match self
8357 .observe_purpose(
8358 ObservationPurpose::ContextExtraction,
8359 llm.complete(&messages, None),
8360 )
8361 .await
8362 {
8363 Ok(response) => {
8364 let value = response.content.trim().to_string();
8365 if value != "__NONE__" && !value.is_empty() {
8366 staged.insert(
8367 extractor.key.clone(),
8368 serde_json::Value::String(value.clone()),
8369 );
8370 debug!(key = %extractor.key, value = %value, "Context extracted");
8371 } else if extractor.required {
8372 warn!(key = %extractor.key, "Required extraction returned no value");
8373 }
8374 }
8375 Err(e) => {
8376 warn!(key = %extractor.key, error = %e, "Context extraction LLM call failed");
8377 }
8378 }
8379 }
8380 staged
8381 }
8382
8383 fn commit_staged_context_writes(&self, staged: &HashMap<String, Value>) {
8384 for (key, value) in staged {
8385 if let Err(error) = self.context_manager.update(key, value.clone()) {
8386 warn!(key = %key, error = %error, "staged context write failed");
8387 }
8388 }
8389 }
8390
8391 async fn run_context_extractors(&self, user_message: &str) {
8393 let staged = self.run_context_extractors_staged(user_message).await;
8394 self.commit_staged_context_writes(&staged);
8395 }
8396
8397 async fn check_memory_compression(&self) -> Result<()> {
8398 if self.memory.needs_compression() {
8399 let result = self.memory.compress(None).await?;
8400 if let CompressResult::Compressed {
8401 messages_summarized,
8402 new_summary_length,
8403 tokens_saved,
8404 } = result
8405 {
8406 let event = MemoryCompressEvent::new(
8407 messages_summarized,
8408 tokens_saved,
8409 new_summary_length as u32,
8410 );
8411 self.hooks.on_memory_compress(&event).await;
8412 debug!(
8413 messages = messages_summarized,
8414 tokens_saved = tokens_saved,
8415 "Memory compressed"
8416 );
8417 }
8418 }
8419
8420 self.handle_memory_overflow().await?;
8422 self.check_memory_budget().await;
8423
8424 Ok(())
8425 }
8426
8427 async fn check_memory_budget(&self) {
8428 let Some(ref budget) = self.memory_token_budget else {
8429 return;
8430 };
8431
8432 let context = match self.memory.get_context().await {
8433 Ok(ctx) => ctx,
8434 Err(_) => return,
8435 };
8436
8437 let used_tokens = context.estimated_tokens();
8439 if budget.is_over_warn_threshold(used_tokens) {
8440 let event = MemoryBudgetEvent::new("memory", used_tokens, budget.total);
8441 self.hooks.on_memory_budget_warning(&event).await;
8442 debug!(
8443 used = used_tokens,
8444 total = budget.total,
8445 percent = event.usage_percent,
8446 "Memory budget warning"
8447 );
8448 }
8449
8450 if let Some(ref summary) = context.summary {
8452 let summary_tokens = ai_agents_memory::estimate_tokens(summary);
8453 let summary_budget = budget.allocation.summary;
8454 if summary_budget > 0 {
8455 let warn_threshold =
8456 (summary_budget as f64 * budget.warn_at_percent as f64 / 100.0) as u32;
8457 if summary_tokens >= warn_threshold {
8458 let event = MemoryBudgetEvent::new("summary", summary_tokens, summary_budget);
8459 self.hooks.on_memory_budget_warning(&event).await;
8460 }
8461 }
8462 }
8463
8464 let recent_tokens: u32 = context
8466 .messages
8467 .iter()
8468 .map(ai_agents_memory::estimate_message_tokens)
8469 .sum();
8470 let recent_budget = budget.allocation.recent_messages;
8471 if recent_budget > 0 {
8472 let warn_threshold =
8473 (recent_budget as f64 * budget.warn_at_percent as f64 / 100.0) as u32;
8474 if recent_tokens >= warn_threshold {
8475 let event = MemoryBudgetEvent::new("recent_messages", recent_tokens, recent_budget);
8476 self.hooks.on_memory_budget_warning(&event).await;
8477 }
8478 }
8479
8480 let relationship_budget = budget.allocation.relationships;
8481 if relationship_budget > 0 {
8482 let relationship_tokens = self
8483 .relationship_memory_text()
8484 .map(|text| ai_agents_memory::estimate_tokens(&text))
8485 .unwrap_or(0);
8486 let warn_threshold =
8487 (relationship_budget as f64 * budget.warn_at_percent as f64 / 100.0) as u32;
8488 if relationship_tokens >= warn_threshold {
8489 let event = MemoryBudgetEvent::new(
8490 "relationships",
8491 relationship_tokens,
8492 relationship_budget,
8493 );
8494 self.hooks.on_memory_budget_warning(&event).await;
8495 }
8496 }
8497 }
8498
8499 async fn handle_memory_overflow(&self) -> Result<()> {
8500 let Some(ref budget) = self.memory_token_budget else {
8501 return Ok(());
8502 };
8503
8504 let context = self.memory.get_context().await?;
8505 let used_tokens = context.estimated_tokens();
8506
8507 if used_tokens <= budget.total {
8508 return Ok(());
8509 }
8510
8511 match budget.overflow_strategy {
8512 OverflowStrategy::TruncateOldest => {
8513 let tokens_to_free = used_tokens - budget.total;
8514 let messages_to_evict = self.calculate_eviction_count(tokens_to_free);
8515 if messages_to_evict > 0 {
8516 self.evict_messages(messages_to_evict, EvictionReason::TokenBudgetExceeded)
8517 .await?;
8518 }
8519 }
8520 OverflowStrategy::SummarizeMore => {
8521 let max_attempts = context.total_messages.max(1);
8522 for _ in 0..max_attempts {
8523 match self.memory.compress(None).await? {
8524 CompressResult::Compressed {
8525 messages_summarized,
8526 ..
8527 } if messages_summarized > 0 => {
8528 let context = self.memory.get_context().await?;
8529 if context.estimated_tokens() <= budget.total {
8530 return Ok(());
8531 }
8532 }
8533 _ => break,
8534 }
8535 }
8536 let context = self.memory.get_context().await?;
8537 let used_tokens = context.estimated_tokens();
8538 if used_tokens > budget.total {
8539 return Err(AgentError::MemoryBudgetExceeded {
8540 used: used_tokens,
8541 budget: budget.total,
8542 });
8543 }
8544 }
8545 OverflowStrategy::Error => {
8546 return Err(AgentError::MemoryBudgetExceeded {
8547 used: used_tokens,
8548 budget: budget.total,
8549 });
8550 }
8551 }
8552 Ok(())
8553 }
8554
8555 fn calculate_eviction_count(&self, tokens_to_free: u32) -> usize {
8556 ((tokens_to_free as f64 / 50.0).ceil() as usize).max(1)
8558 }
8559
8560 async fn evict_messages(&self, count: usize, reason: EvictionReason) -> Result<()> {
8561 let evicted = self.memory.evict_oldest(count).await?;
8562 if !evicted.is_empty() {
8563 let event = MemoryEvictEvent {
8564 reason,
8565 messages_evicted: evicted.len(),
8566 importance_scores: vec![],
8567 };
8568 self.hooks.on_memory_evict(&event).await;
8569 debug!(count = evicted.len(), "Messages evicted from memory");
8570 }
8571 Ok(())
8572 }
8573
8574 #[instrument(skip(self, input), fields(agent = %self.info.name))]
8575 async fn determine_reasoning_mode(&self, input: &str) -> Result<ReasoningMode> {
8576 match self.determine_reasoning_mode_strict(input).await {
8577 Ok(mode) => Ok(mode),
8578 Err(_) => Ok(ReasoningMode::None),
8579 }
8580 }
8581
8582 async fn determine_reasoning_mode_strict(&self, input: &str) -> Result<ReasoningMode> {
8583 let effective_config = self.get_effective_reasoning_config();
8584
8585 if !matches!(effective_config.mode, ReasoningMode::Auto) {
8586 return Ok(effective_config.mode.clone());
8587 }
8588
8589 let judge_llm = effective_config
8590 .judge_llm
8591 .as_ref()
8592 .and_then(|alias| self.llm_registry.get(alias).ok())
8593 .or_else(|| self.llm_registry.router().ok())
8594 .or_else(|| self.llm_registry.default().ok());
8595
8596 let Some(llm) = judge_llm else {
8597 return Ok(ReasoningMode::None);
8598 };
8599
8600 let prompt = format!(
8601 r#"Analyze this user request and determine the appropriate reasoning mode.
8602
8603User request: "{}"
8604
8605Choose ONE of these modes:
8606- none: Simple queries, greetings, direct answers (fastest)
8607- cot: Complex analysis, multi-step reasoning, math problems
8608- react: Tasks requiring multiple tool calls with observation
8609- plan_and_execute: Complex multi-step tasks requiring coordination
8610
8611Respond with ONLY the mode name (none, cot, react, or plan_and_execute)."#,
8612 input
8613 );
8614
8615 let messages = vec![ChatMessage::user(&prompt)];
8616 let response = self
8617 .observe_purpose(
8618 ObservationPurpose::ReflectionDecision,
8619 llm.complete(&messages, None),
8620 )
8621 .await
8622 .map_err(|e| AgentError::LLM(e.to_string()))?;
8623
8624 let mode_str = response.content.trim().to_lowercase();
8625 Ok(match mode_str.as_str() {
8626 "cot" => ReasoningMode::CoT,
8627 "react" => ReasoningMode::React,
8628 "plan_and_execute" => ReasoningMode::PlanAndExecute,
8629 _ => ReasoningMode::None,
8630 })
8631 }
8632
8633 fn build_cot_system_prompt(&self, base_prompt: &str) -> String {
8634 format!(
8635 "{}\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>",
8636 base_prompt
8637 )
8638 }
8639
8640 fn build_react_system_prompt(&self, base_prompt: &str) -> String {
8641 format!(
8642 "{}\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>",
8643 base_prompt
8644 )
8645 }
8646
8647 async fn generate_plan(&self, input: &str) -> Result<Plan> {
8648 let effective = self.get_effective_reasoning_config();
8649 let planning_config = effective.get_planning();
8650
8651 let planner_llm = planning_config
8652 .and_then(|c| c.planner_llm.as_ref())
8653 .and_then(|alias| self.llm_registry.get(alias).ok())
8654 .or_else(|| self.llm_registry.router().ok())
8655 .or_else(|| self.llm_registry.default().ok())
8656 .ok_or_else(|| AgentError::Config("No LLM available for planning".into()))?;
8657
8658 let mut available_tool_ids: Vec<String> = self
8659 .get_available_tool_ids()
8660 .await
8661 .unwrap_or_else(|_| self.tools.list_ids());
8662 let mut available_skills: Vec<String> = self.skills.iter().map(|s| s.id.clone()).collect();
8663
8664 if let Some(config) = planning_config {
8666 if !config.available.tools.is_all() {
8667 available_tool_ids.retain(|t| config.available.tools.allows(t));
8668 }
8669 if !config.available.skills.is_all() {
8670 available_skills.retain(|s| config.available.skills.allows(s));
8671 }
8672 }
8673
8674 let tool_descriptions: Vec<String> = available_tool_ids
8677 .iter()
8678 .filter_map(|id| {
8679 self.tools.get(id).map(|tool| {
8680 let schema = tool.input_schema();
8681 let args_desc = schema
8682 .get("properties")
8683 .and_then(|p| serde_json::to_string(p).ok())
8684 .unwrap_or_else(|| "{}".to_string());
8685 format!(
8686 "- {} ({}): {}\n Arguments: {}",
8687 id,
8688 tool.name(),
8689 tool.description(),
8690 args_desc
8691 )
8692 })
8693 })
8694 .collect();
8695
8696 let tools_section = if tool_descriptions.is_empty() {
8697 "Available tools: none".to_string()
8698 } else {
8699 format!("Available tools:\n{}", tool_descriptions.join("\n"))
8700 };
8701
8702 let skills_section = if available_skills.is_empty() {
8703 "Available skills: none".to_string()
8704 } else {
8705 format!("Available skills: {}", available_skills.join(", "))
8706 };
8707
8708 let prompt = format!(
8709 r#"Create a step-by-step plan to accomplish this goal.
8710
8711Goal: "{}"
8712
8713{}
8714
8715{}
8716
8717Create a plan with clear steps. For each step, specify:
8718- description: What this step accomplishes
8719- action_type: "tool", "skill", "think", or "respond"
8720- action_target: The tool/skill id (if applicable)
8721- args: The arguments object matching the tool's schema (if action_type is "tool")
8722- dependencies: List of step IDs this depends on (empty if none)
8723
8724Respond in JSON format:
8725{{
8726 "steps": [
8727 {{"id": "step1", "description": "...", "action_type": "tool", "action_target": "tool_id", "args": {{"required_field": "value"}}, "dependencies": []}},
8728 {{"id": "step2", "description": "...", "action_type": "think", "action_target": "...", "dependencies": ["step1"]}}
8729 ]
8730}}"#,
8731 input, tools_section, skills_section,
8732 );
8733
8734 let messages = vec![ChatMessage::user(&prompt)];
8735 let response = self
8736 .observe_purpose(
8737 ObservationPurpose::PlanGeneration,
8738 planner_llm.complete(&messages, None),
8739 )
8740 .await
8741 .map_err(|e| AgentError::LLM(format!("Planning failed: {}", e)))?;
8742
8743 let mut plan = Plan::new(input);
8744
8745 if let Some(json_start) = response.content.find('{')
8746 && let Some(json_end) = response.content.rfind('}')
8747 {
8748 let json_str = &response.content[json_start..=json_end];
8749 if let Ok(parsed) = serde_json::from_str::<serde_json::Value>(json_str)
8750 && let Some(steps) = parsed.get("steps").and_then(|s| s.as_array())
8751 {
8752 for step_value in steps {
8753 let id = step_value
8754 .get("id")
8755 .and_then(|v| v.as_str())
8756 .unwrap_or("step");
8757 let desc = step_value
8758 .get("description")
8759 .and_then(|v| v.as_str())
8760 .unwrap_or("");
8761 let action_type = step_value
8762 .get("action_type")
8763 .and_then(|v| v.as_str())
8764 .unwrap_or("think");
8765 let action_target = step_value
8766 .get("action_target")
8767 .and_then(|v| v.as_str())
8768 .unwrap_or("");
8769 let args = step_value
8770 .get("args")
8771 .cloned()
8772 .unwrap_or(serde_json::json!({}));
8773 let deps: Vec<String> = step_value
8774 .get("dependencies")
8775 .and_then(|v| v.as_array())
8776 .map(|arr| {
8777 arr.iter()
8778 .filter_map(|v| v.as_str().map(String::from))
8779 .collect()
8780 })
8781 .unwrap_or_default();
8782
8783 let action = match action_type {
8784 "tool" => PlanAction::tool(action_target, args),
8785 "skill" => PlanAction::skill(action_target),
8786 "respond" => PlanAction::respond(action_target),
8787 _ => PlanAction::think(desc),
8788 };
8789
8790 let step = PlanStep::new(desc, action)
8791 .with_id(id)
8792 .with_dependencies(deps);
8793 plan.add_step(step);
8794 }
8795 }
8796 }
8797
8798 if plan.steps.is_empty() {
8799 plan.add_step(PlanStep::new(
8800 "Process the request",
8801 PlanAction::think(input),
8802 ));
8803 plan.add_step(PlanStep::new(
8804 "Provide response",
8805 PlanAction::respond("Answer based on analysis"),
8806 ));
8807 }
8808
8809 Ok(plan)
8810 }
8811
8812 async fn execute_plan(&self, plan: &mut Plan) -> Result<String> {
8813 let llm = self.get_state_llm()?;
8814 let mut results: HashMap<String, serde_json::Value> = HashMap::new();
8815 let effective = self.get_effective_reasoning_config();
8816 let max_steps = effective.get_planning().map(|c| c.max_steps).unwrap_or(10);
8817
8818 plan.status = PlanStatus::InProgress;
8819
8820 for step_idx in 0..plan.steps.len().min(max_steps as usize) {
8821 let step = &plan.steps[step_idx];
8822
8823 let deps_satisfied = step.dependencies.iter().all(|dep| {
8824 plan.steps
8825 .iter()
8826 .find(|s| &s.id == dep)
8827 .map(|s| s.status.is_completed())
8828 .unwrap_or(false)
8829 });
8830
8831 if !deps_satisfied {
8832 continue;
8833 }
8834
8835 plan.steps[step_idx].mark_running();
8836
8837 let result = match &plan.steps[step_idx].action {
8838 PlanAction::Tool { tool, args } => {
8839 let has_dep_results = plan.steps[step_idx]
8845 .dependencies
8846 .iter()
8847 .any(|dep| results.contains_key(dep));
8848
8849 let final_args = if has_dep_results {
8850 let dep_context: String = plan.steps[step_idx]
8851 .dependencies
8852 .iter()
8853 .filter_map(|dep| results.get(dep).map(|r| format!("{}: {}", dep, r)))
8854 .collect::<Vec<_>>()
8855 .join("\n");
8856
8857 let tool_schema = self
8858 .tools
8859 .get(tool)
8860 .map(|t| {
8861 let schema = t.input_schema();
8862 let props = schema
8863 .get("properties")
8864 .and_then(|p| serde_json::to_string(p).ok())
8865 .unwrap_or_else(|| "{}".to_string());
8866 format!(
8867 "{}: {}\nArguments schema: {}",
8868 t.id(),
8869 t.description(),
8870 props
8871 )
8872 })
8873 .unwrap_or_default();
8874
8875 let step_desc = &plan.steps[step_idx].description;
8876 let arg_prompt = format!(
8877 "Generate the JSON arguments for a tool call.\n\n\
8878 Tool: {}\n\n\
8879 Task: {}\n\n\
8880 Previous step results:\n{}\n\n\
8881 Planner's draft arguments: {}\n\n\
8882 Produce ONLY a valid JSON object with the correct argument values.\n\
8883 Use actual values from the previous step results, not template references.",
8884 tool_schema,
8885 step_desc,
8886 dep_context,
8887 serde_json::to_string(args).unwrap_or_default()
8888 );
8889 let messages = vec![ChatMessage::user(&arg_prompt)];
8890 match self
8891 .observe_purpose(
8892 ObservationPurpose::PlanStep,
8893 llm.complete(&messages, None),
8894 )
8895 .await
8896 {
8897 Ok(resp) => {
8898 let content = resp.content.trim();
8899 let json_start = content.find('{');
8901 let json_end = content.rfind('}');
8902 if let (Some(start), Some(end)) = (json_start, json_end) {
8903 serde_json::from_str(&content[start..=end])
8904 .unwrap_or_else(|_| args.clone())
8905 } else {
8906 args.clone()
8907 }
8908 }
8909 Err(_) => args.clone(),
8910 }
8911 } else {
8912 args.clone()
8913 };
8914
8915 let request = ToolExecutionRequest::new(
8916 uuid::Uuid::new_v4().to_string(),
8917 tool.clone(),
8918 final_args,
8919 ToolCallSource::Plan {
8920 step_index: step_idx,
8921 },
8922 );
8923 match self.execute_tool_record(request).await {
8924 Ok(record) if record.success => {
8925 serde_json::json!({ "output": record.model_output_string() })
8926 }
8927 Ok(record) => {
8928 plan.steps[step_idx].mark_failed(record.model_output_string());
8929 continue;
8930 }
8931 Err(e) => {
8932 plan.steps[step_idx].mark_failed(e.to_string());
8933 continue;
8934 }
8935 }
8936 }
8937 PlanAction::Skill { skill } => {
8938 if let Some(skill_def) = self.skills.iter().find(|s| &s.id == skill) {
8939 if let Some(ref executor) = self.skill_executor {
8940 match executor
8941 .execute_with_invoker(skill_def, "", serde_json::json!({}), self)
8942 .await
8943 {
8944 Ok(output) => serde_json::json!({ "output": output }),
8945 Err(e) => {
8946 plan.steps[step_idx].mark_failed(e.to_string());
8947 continue;
8948 }
8949 }
8950 } else {
8951 serde_json::json!({ "output": "Skill executor not available" })
8952 }
8953 } else {
8954 plan.steps[step_idx].mark_failed("Skill not found");
8955 continue;
8956 }
8957 }
8958 PlanAction::Think { prompt } => {
8959 let context: String = results
8960 .iter()
8961 .map(|(k, v)| format!("{}: {}", k, v))
8962 .collect::<Vec<_>>()
8963 .join("\n");
8964
8965 let think_prompt = format!("Context:\n{}\n\nTask: {}", context, prompt);
8966 let messages = vec![ChatMessage::user(&think_prompt)];
8967
8968 match self
8969 .observe_purpose(
8970 ObservationPurpose::PlanStep,
8971 llm.complete(&messages, None),
8972 )
8973 .await
8974 {
8975 Ok(resp) => serde_json::json!({ "output": resp.content }),
8976 Err(e) => {
8977 plan.steps[step_idx].mark_failed(e.to_string());
8978 continue;
8979 }
8980 }
8981 }
8982 PlanAction::Respond { template } => {
8983 let context: String = results
8984 .iter()
8985 .map(|(k, v)| format!("{}: {}", k, v))
8986 .collect::<Vec<_>>()
8987 .join("\n");
8988
8989 let respond_prompt = format!(
8990 "Based on this context:\n{}\n\nGenerate a response following this template/instruction: {}",
8991 context, template
8992 );
8993 let messages = vec![ChatMessage::user(&respond_prompt)];
8994
8995 match self
8996 .observe_purpose(
8997 ObservationPurpose::PlanStep,
8998 llm.complete(&messages, None),
8999 )
9000 .await
9001 {
9002 Ok(resp) => serde_json::json!({ "output": resp.content }),
9003 Err(e) => {
9004 plan.steps[step_idx].mark_failed(e.to_string());
9005 continue;
9006 }
9007 }
9008 }
9009 };
9010
9011 results.insert(plan.steps[step_idx].id.clone(), result.clone());
9012 plan.steps[step_idx].mark_completed(Some(result));
9013 }
9014
9015 let has_failures = plan.steps.iter().any(|s| s.status.is_failed());
9017 if has_failures {
9018 let failed_ids: Vec<String> = plan
9019 .steps
9020 .iter()
9021 .filter(|s| s.status.is_failed())
9022 .map(|s| s.id.clone())
9023 .collect();
9024 plan.status = PlanStatus::Failed {
9025 error: format!("Steps failed: {}", failed_ids.join(", ")),
9026 };
9027 } else {
9028 plan.status = PlanStatus::Completed;
9029 }
9030
9031 let all_outputs: Vec<String> = plan
9033 .steps
9034 .iter()
9035 .filter(|s| s.status.is_completed())
9036 .filter_map(|s| {
9037 s.result
9038 .as_ref()
9039 .and_then(|r| r.get("output"))
9040 .and_then(|o| o.as_str())
9041 .map(|o| format!("{}: {}", s.description, o))
9042 })
9043 .collect();
9044
9045 if all_outputs.is_empty() {
9046 return Ok("Plan execution completed but produced no results.".to_string());
9047 }
9048
9049 if all_outputs.len() == 1 {
9050 return Ok(all_outputs.into_iter().next().unwrap());
9051 }
9052
9053 let context = all_outputs.join("\n\n");
9055 let prompt = format!(
9056 "You completed a multi-step plan for: \"{}\"\n\nStep results:\n{}\n\nProvide a coherent final response that synthesizes these results.",
9057 plan.goal, context
9058 );
9059 let messages = vec![ChatMessage::user(&prompt)];
9060 match self
9061 .observe_purpose(ObservationPurpose::PlanStep, llm.complete(&messages, None))
9062 .await
9063 {
9064 Ok(resp) => Ok(resp.content.trim().to_string()),
9065 Err(_) => Ok(context),
9066 }
9067 }
9068
9069 fn extract_thinking(&self, content: &str) -> (Option<String>, String) {
9070 if let Some(start) = content.find("<thinking>")
9071 && let Some(end) = content.find("</thinking>")
9072 {
9073 let thinking = content[start + 10..end].trim().to_string();
9074 let answer = content[end + 11..].trim().to_string();
9075 return (Some(thinking), answer);
9076 }
9077 (None, content.to_string())
9078 }
9079
9080 fn format_response_with_thinking(&self, thinking: Option<&str>, answer: &str) -> String {
9081 match self.get_effective_reasoning_config().output {
9082 ReasoningOutput::Hidden => answer.to_string(),
9083 ReasoningOutput::Visible => {
9084 if let Some(t) = thinking {
9085 format!("Thinking:\n{}\n\nAnswer:\n{}", t, answer)
9086 } else {
9087 answer.to_string()
9088 }
9089 }
9090 ReasoningOutput::Tagged => {
9091 if let Some(t) = thinking {
9092 format!("<thinking>{}</thinking>\n{}", t, answer)
9093 } else {
9094 answer.to_string()
9095 }
9096 }
9097 }
9098 }
9099
9100 fn disambiguation_question_response(
9103 question: &ClarificationQuestion,
9104 detection: &AmbiguityDetectionResult,
9105 awaiting_confirmation: bool,
9106 ) -> AgentResponse {
9107 let status = if awaiting_confirmation {
9108 "awaiting_confirmation"
9109 } else {
9110 "awaiting_clarification"
9111 };
9112 AgentResponse::new(&question.question).with_metadata(
9113 "disambiguation",
9114 serde_json::json!({
9115 "status": status,
9116 "options": question.options,
9117 "clarifying": question.clarifying,
9118 "detection": {
9119 "type": detection.ambiguity_type,
9120 "confidence": detection.confidence,
9121 "what_is_unclear": detection.what_is_unclear,
9122 }
9123 }),
9124 )
9125 }
9126
9127 async fn resolve_disambiguation(&self, input: &str) -> Result<DisambiguationDispatch> {
9139 let Some(ref disambiguator) = self.disambiguation_manager else {
9140 return Ok(DisambiguationDispatch::Proceed(input.to_string()));
9141 };
9142 let disambiguation_context = self.build_disambiguation_context().await?;
9143
9144 let state_override = self
9146 .state_machine
9147 .as_ref()
9148 .and_then(|sm| sm.current_definition())
9149 .and_then(|def| def.disambiguation.clone());
9150
9151 let state_generation = self
9152 .state_machine
9153 .as_ref()
9154 .map(|state_machine| state_machine.generation());
9155 let disambiguation_epoch = self.disambiguation_epoch.load(Ordering::SeqCst);
9156 let mut disambiguation_result = self
9157 .observe_purpose(
9158 ObservationPurpose::DisambiguationDetection,
9159 disambiguator.process_input_with_override(
9160 input,
9161 &disambiguation_context,
9162 state_override.as_ref(),
9163 None,
9164 ),
9165 )
9166 .await?;
9167 let current_state_generation = self
9168 .state_machine
9169 .as_ref()
9170 .map(|state_machine| state_machine.generation());
9171 if current_state_generation != state_generation
9172 || self.disambiguation_epoch.load(Ordering::SeqCst) != disambiguation_epoch
9173 {
9174 disambiguator.clear_pending().await;
9175 *self.pending_skill_id.write() = None;
9176 disambiguation_result = DisambiguationResult::Abandoned { new_input: None };
9177 info!(
9178 confirmation_event = "invalidated",
9179 invalidation_reason = "state_generation_changed",
9180 "Disambiguation result invalidated before redispatch"
9181 );
9182 }
9183 match disambiguation_result {
9184 DisambiguationResult::Clear => {
9185 debug!("Input is clear, proceeding normally");
9186 Ok(DisambiguationDispatch::Proceed(input.to_string()))
9187 }
9188 DisambiguationResult::NeedsClarification {
9189 question,
9190 detection,
9191 } => {
9192 let admission = match self
9193 .admit_disambiguation_redispatch(disambiguation_epoch, state_generation)
9194 .await
9195 {
9196 Ok(admission) => admission,
9197 Err(error) => {
9198 *self.pending_skill_id.write() = None;
9199 return Err(error);
9200 }
9201 };
9202 let awaiting_confirmation = disambiguator.has_pending_confirmation().await;
9203 info!(
9204 ambiguity_type = ?detection.ambiguity_type,
9205 confidence = detection.confidence,
9206 "Input requires clarification"
9207 );
9208
9209 self.commit_root_user_message(input).await?;
9212 self.memory
9213 .add_message(ChatMessage::assistant(&question.question))
9214 .await?;
9215
9216 let response = Self::disambiguation_question_response(
9217 &question,
9218 &detection,
9219 awaiting_confirmation,
9220 );
9221 drop(admission);
9222 self.finish_turn_if_root(&response).await?;
9223 Ok(DisambiguationDispatch::Terminal(response))
9224 }
9225 DisambiguationResult::Clarified {
9226 enriched_input,
9227 resolved,
9228 ..
9229 } => {
9230 let admission = match self
9231 .admit_disambiguation_redispatch(disambiguation_epoch, state_generation)
9232 .await
9233 {
9234 Ok(admission) => admission,
9235 Err(error) => {
9236 *self.pending_skill_id.write() = None;
9237 return Err(error);
9238 }
9239 };
9240 info!(
9241 resolved_count = resolved.len(),
9242 enriched = %enriched_input,
9243 "Input clarified, injecting resolved intent into context"
9244 );
9245
9246 for (key, value) in &resolved {
9249 let context_key = format!("disambiguation.{}", key);
9250 let _ = self.context_manager.set(&context_key, value.clone());
9251 }
9252
9253 if let Some(intent) = resolved.get("intent") {
9254 let _ = self.context_manager.set("resolved_intent", intent.clone());
9255 }
9256
9257 let _ = self
9258 .context_manager
9259 .set("disambiguation.resolved", serde_json::Value::Bool(true));
9260
9261 let skill_id = self.pending_skill_id.read().clone();
9265 drop(admission);
9266 if let Some(skill_id) = skill_id {
9267 info!(skill_id = %skill_id, "Re-checking skill disambiguation on clarified input");
9268 return Ok(DisambiguationDispatch::RecheckSkill {
9269 skill_id,
9270 enriched_input,
9271 disambiguation_epoch,
9272 state_generation,
9273 });
9274 }
9275 Ok(DisambiguationDispatch::Proceed(enriched_input))
9276 }
9277 DisambiguationResult::ProceedWithBestGuess { enriched_input } => {
9278 info!("Proceeding with best guess interpretation");
9279
9280 let skill_id = self.pending_skill_id.read().clone();
9282 if let Some(skill_id) = skill_id {
9283 info!(skill_id = %skill_id, "Re-checking skill disambiguation on best-guess input");
9284 return Ok(DisambiguationDispatch::RecheckSkill {
9285 skill_id,
9286 enriched_input,
9287 disambiguation_epoch,
9288 state_generation,
9289 });
9290 }
9291 Ok(DisambiguationDispatch::Proceed(enriched_input))
9292 }
9293 DisambiguationResult::GiveUp { reason } => {
9294 *self.pending_skill_id.write() = None;
9295 warn!(reason = %reason, "Disambiguation gave up");
9296 let apology = self
9297 .generate_localized_apology(
9298 "Generate a brief, polite apology saying you couldn't understand the request. Be concise.",
9299 &reason,
9300 )
9301 .await
9302 .unwrap_or_else(|_| {
9303 format!("I'm sorry, I couldn't understand your request: {}", reason)
9304 });
9305 let response = AgentResponse::new(&apology);
9306 self.finish_turn_if_root(&response).await?;
9307 Ok(DisambiguationDispatch::Terminal(response))
9308 }
9309 DisambiguationResult::Escalate { reason } => {
9310 *self.pending_skill_id.write() = None;
9311 info!(reason = %reason, "Escalating to human");
9312 if let Some(ref hitl) = self.hitl_engine {
9313 let trigger =
9314 ApprovalTrigger::condition("disambiguation_escalation", reason.clone());
9315 let mut context_map = HashMap::new();
9316 context_map.insert("original_input".to_string(), serde_json::json!(input));
9317 context_map.insert("reason".to_string(), serde_json::json!(&reason));
9318 let check_result = HITLCheckResult::required(
9319 trigger,
9320 context_map,
9321 format!("User request needs human assistance: {}", reason),
9322 Some(hitl.config().default_timeout_seconds),
9323 );
9324 let result = self.request_hitl_approval(check_result).await?;
9325 if matches!(
9326 result,
9327 ApprovalResult::Approved | ApprovalResult::Modified { .. }
9328 ) {
9329 return Ok(DisambiguationDispatch::Proceed(input.to_string()));
9331 }
9332 }
9333 let apology = self
9334 .generate_localized_apology(
9335 "Explain briefly that you're transferring the user to a human agent for help.",
9336 &reason,
9337 )
9338 .await
9339 .unwrap_or_else(|_| {
9340 format!("I need human assistance to help with your request: {}", reason)
9341 });
9342 let response = AgentResponse::new(&apology);
9343 self.finish_turn_if_root(&response).await?;
9344 Ok(DisambiguationDispatch::Terminal(response))
9345 }
9346 DisambiguationResult::Abandoned { new_input } => {
9347 *self.pending_skill_id.write() = None;
9348
9349 info!(
9350 has_new_input = new_input.is_some(),
9351 "Clarification abandoned by user"
9352 );
9353
9354 self.commit_root_user_message(input).await?;
9355
9356 match new_input {
9357 Some(fresh_input) => {
9358 Ok(DisambiguationDispatch::Proceed(fresh_input))
9361 }
9362 None => {
9363 let ack = self
9365 .generate_localized_apology(
9366 "The user changed their mind about their previous request. \
9367 Generate a brief, friendly acknowledgment (e.g. 'OK, no problem. What else can I help with?'). \
9368 Do NOT apologize excessively. Be concise.",
9369 "User abandoned clarification",
9370 )
9371 .await
9372 .unwrap_or_else(|_| {
9373 "OK, no problem. What else can I help with?".to_string()
9374 });
9375
9376 self.memory
9377 .add_message(ChatMessage::assistant(&ack))
9378 .await?;
9379
9380 let response = AgentResponse::new(&ack);
9381 self.finish_turn_if_root(&response).await?;
9382 Ok(DisambiguationDispatch::Terminal(response))
9383 }
9384 }
9385 }
9386 }
9387 }
9388
9389 async fn prepare_turn_context(&self) -> Result<()> {
9393 if !self.context_initialized.load(Ordering::SeqCst) {
9394 self.context_manager.initialize().await?;
9395 self.context_initialized.store(true, Ordering::SeqCst);
9396 debug!("Context manager initialized (defaults, env, builtins)");
9397 }
9398
9399 self.check_turn_timeout().await?;
9400 self.context_manager.refresh_per_turn().await?;
9401 self.context_manager.validate()
9402 }
9403
9404 async fn run_loop(&self, input: &str) -> Result<AgentResponse> {
9407 self.init_storage().await?;
9411 self.begin_root_turn();
9412 let _root_cleanup = RootTurnCleanup::new(self);
9413 info!(input_len = input.len(), "Starting chat");
9414
9415 self.hooks.on_message_received(input).await;
9416
9417 self.prepare_turn_context().await?;
9418
9419 self.clear_disambiguation_context();
9422
9423 let input_to_run = match self.resolve_disambiguation(input).await? {
9426 DisambiguationDispatch::Terminal(response) => return Ok(response),
9427 DisambiguationDispatch::RecheckSkill {
9428 skill_id,
9429 enriched_input,
9430 disambiguation_epoch,
9431 state_generation,
9432 } => {
9433 return self
9434 .recheck_skill_disambiguation(
9435 &skill_id,
9436 &enriched_input,
9437 disambiguation_epoch,
9438 state_generation,
9439 )
9440 .await;
9441 }
9442 DisambiguationDispatch::Proceed(input) => input,
9443 };
9444
9445 self.run_loop_internal(&input_to_run).await
9446 }
9447
9448 async fn generate_localized_apology(&self, instruction: &str, reason: &str) -> Result<String> {
9450 let llm = self.llm_registry.router().map_err(|e| {
9451 AgentError::LLM(format!(
9452 "Router LLM not available for localized response: {}",
9453 e
9454 ))
9455 })?;
9456
9457 let recent: Vec<String> = self
9458 .memory
9459 .get_messages(Some(3))
9460 .await?
9461 .iter()
9462 .map(|m| m.content.clone())
9463 .collect();
9464
9465 let context_hint = if recent.is_empty() {
9466 String::new()
9467 } else {
9468 format!(
9469 "\nRecent conversation (detect the user's language from this):\n{}\n",
9470 recent.join("\n")
9471 )
9472 };
9473
9474 let prompt = format!(
9475 "{}\nReason: {}\n{}Respond in the same language as the user. Output ONLY the message, nothing else.",
9476 instruction, reason, context_hint
9477 );
9478
9479 let messages = vec![ChatMessage::user(&prompt)];
9480 let response = self
9481 .observe_purpose(
9482 ObservationPurpose::DisambiguationClarification,
9483 llm.complete(&messages, None),
9484 )
9485 .await
9486 .map_err(|e| AgentError::LLM(format!("Localized response generation failed: {}", e)))?;
9487
9488 Ok(response.content.trim().to_string())
9489 }
9490
9491 fn render_action_args(&self, args: &Value) -> Value {
9495 let context = self.build_context_with_overlays();
9496 match args {
9497 Value::Object(map) => {
9498 let mut rendered = serde_json::Map::new();
9499 for (k, v) in map {
9500 match v {
9501 Value::String(s) if s.contains("{{") => {
9502 match self.template_renderer.render(s, &context) {
9503 Ok(rendered_str) => {
9504 rendered.insert(k.clone(), Value::String(rendered_str));
9505 }
9506 Err(_) => {
9507 rendered.insert(k.clone(), v.clone());
9508 }
9509 }
9510 }
9511 _ => {
9512 rendered.insert(k.clone(), v.clone());
9513 }
9514 }
9515 }
9516 Value::Object(rendered)
9517 }
9518 _ => args.clone(),
9519 }
9520 }
9521
9522 fn clear_disambiguation_context(&self) {
9524 let _ = self
9525 .context_manager
9526 .set("resolved_intent", serde_json::Value::Null);
9527
9528 let all = self.context_manager.get_all();
9529 for key in all.keys() {
9530 if key.starts_with("disambiguation.") {
9531 let _ = self.context_manager.set(key, serde_json::Value::Null);
9532 }
9533 }
9534 }
9535
9536 async fn recheck_skill_disambiguation(
9542 &self,
9543 skill_id: &str,
9544 enriched_input: &str,
9545 expected_disambiguation_epoch: u64,
9546 expected_state_generation: Option<u64>,
9547 ) -> Result<AgentResponse> {
9548 let skill = self
9549 .skill_router
9550 .as_ref()
9551 .and_then(|r| r.get_skill(skill_id).cloned());
9552
9553 if let Some(ref skill) = skill
9555 && let Some(ref skill_disambig) = skill.disambiguation
9556 && skill_disambig.enabled.unwrap_or(false)
9557 && let Some(ref disambiguator) = self.disambiguation_manager
9558 {
9559 let context = self.build_disambiguation_context().await?;
9560 let state_override = self
9561 .state_machine
9562 .as_ref()
9563 .and_then(|sm| sm.current_definition())
9564 .and_then(|def| def.disambiguation.clone());
9565
9566 let disambiguation_result = self
9567 .observe_purpose(
9568 ObservationPurpose::DisambiguationDetection,
9569 disambiguator.process_input_with_override(
9570 enriched_input,
9571 &context,
9572 state_override.as_ref(),
9573 Some(skill_disambig),
9574 ),
9575 )
9576 .await?;
9577 let current_state_generation = self
9578 .state_machine
9579 .as_ref()
9580 .map(|state_machine| state_machine.generation());
9581 if current_state_generation != expected_state_generation
9582 || self.disambiguation_epoch.load(Ordering::SeqCst) != expected_disambiguation_epoch
9583 {
9584 disambiguator.clear_pending().await;
9585 *self.pending_skill_id.write() = None;
9586 return Err(AgentError::Other(
9587 "State or reset ownership changed during skill disambiguation recheck"
9588 .to_string(),
9589 ));
9590 }
9591 match disambiguation_result {
9592 DisambiguationResult::Clear => {
9593 debug!(skill_id = %skill_id, "Skill re-check: all fields present");
9594 }
9595 DisambiguationResult::NeedsClarification {
9596 question,
9597 detection,
9598 } => {
9599 let admission = self
9600 .admit_disambiguation_redispatch(
9601 expected_disambiguation_epoch,
9602 expected_state_generation,
9603 )
9604 .await?;
9605 let awaiting_confirmation = disambiguator.has_pending_confirmation().await;
9606 info!(
9607 skill_id = %skill_id,
9608 ambiguity_type = ?detection.ambiguity_type,
9609 what_is_unclear = ?detection.what_is_unclear,
9610 "Skill re-check: still missing fields, asking again"
9611 );
9612 self.memory
9616 .add_message(ChatMessage::user(enriched_input))
9617 .await?;
9618 self.memory
9619 .add_message(ChatMessage::assistant(&question.question))
9620 .await?;
9621
9622 let response = AgentResponse::new(&question.question).with_metadata(
9623 "disambiguation",
9624 serde_json::json!({
9625 "status": if awaiting_confirmation { "awaiting_confirmation" } else { "awaiting_clarification" },
9626 "skill_id": skill_id,
9627 "options": question.options,
9628 "clarifying": question.clarifying,
9629 "detection": {
9630 "type": detection.ambiguity_type,
9631 "confidence": detection.confidence,
9632 "what_is_unclear": detection.what_is_unclear,
9633 }
9634 }),
9635 );
9636 drop(admission);
9637 self.finish_turn_if_root(&response).await?;
9638 return Ok(response);
9639 }
9640 DisambiguationResult::Clarified {
9641 enriched_input: re_enriched,
9642 ..
9643 } => {
9644 debug!(skill_id = %skill_id, "Skill re-check: clarified immediately, executing");
9645 let admission = self
9646 .admit_disambiguation_redispatch(
9647 expected_disambiguation_epoch,
9648 expected_state_generation,
9649 )
9650 .await?;
9651 *self.pending_skill_id.write() = None;
9652 drop(admission);
9653 let skill_response = self.execute_skill_by_id(skill_id, &re_enriched).await?;
9654 self.memory
9655 .add_message(ChatMessage::user(&re_enriched))
9656 .await?;
9657 return self
9658 .handle_skill_response(
9659 &re_enriched,
9660 skill_id,
9661 skill_response,
9662 &HashMap::new(),
9663 )
9664 .await;
9665 }
9666 DisambiguationResult::ProceedWithBestGuess {
9667 enriched_input: re_enriched,
9668 } => {
9669 debug!(skill_id = %skill_id, "Skill re-check: proceeding with best guess");
9670 let admission = self
9671 .admit_disambiguation_redispatch(
9672 expected_disambiguation_epoch,
9673 expected_state_generation,
9674 )
9675 .await?;
9676 *self.pending_skill_id.write() = None;
9677 drop(admission);
9678 let skill_response = self.execute_skill_by_id(skill_id, &re_enriched).await?;
9679 self.memory
9680 .add_message(ChatMessage::user(&re_enriched))
9681 .await?;
9682 return self
9683 .handle_skill_response(
9684 &re_enriched,
9685 skill_id,
9686 skill_response,
9687 &HashMap::new(),
9688 )
9689 .await;
9690 }
9691 DisambiguationResult::GiveUp { reason } => {
9692 *self.pending_skill_id.write() = None;
9693 let apology = self
9694 .generate_localized_apology(
9695 "Generate a brief, polite apology saying you couldn't understand the request. Be concise.",
9696 &reason,
9697 )
9698 .await
9699 .unwrap_or_else(|_| {
9700 format!("I'm sorry, I couldn't understand your request: {}", reason)
9701 });
9702 let response = AgentResponse::new(&apology);
9703 self.finish_turn_if_root(&response).await?;
9704 return Ok(response);
9705 }
9706 DisambiguationResult::Escalate { reason } => {
9707 *self.pending_skill_id.write() = None;
9708 let apology = self
9709 .generate_localized_apology(
9710 "Explain briefly that you're transferring the user to a human agent for help.",
9711 &reason,
9712 )
9713 .await
9714 .unwrap_or_else(|_| {
9715 format!("I need human assistance to help with your request: {}", reason)
9716 });
9717 let response = AgentResponse::new(&apology);
9718 self.finish_turn_if_root(&response).await?;
9719 return Ok(response);
9720 }
9721 DisambiguationResult::Abandoned { new_input } => {
9722 *self.pending_skill_id.write() = None;
9725 debug!(skill_id = %skill_id, "Skill re-check: abandoned by user");
9726 if let Some(fresh) = new_input {
9727 return self.run_loop_internal(&fresh).await;
9728 }
9729 let ack = self
9730 .generate_localized_apology(
9731 "The user changed their mind about their previous request. \
9732 Generate a brief, friendly acknowledgment (e.g. 'OK, no problem. What else can I help with?'). \
9733 Do NOT apologize excessively. Be concise.",
9734 "User abandoned clarification",
9735 )
9736 .await
9737 .unwrap_or_else(|_| {
9738 "OK, no problem. What else can I help with?".to_string()
9739 });
9740 self.memory
9741 .add_message(ChatMessage::assistant(&ack))
9742 .await?;
9743 let response = AgentResponse::new(&ack);
9744 self.finish_turn_if_root(&response).await?;
9745 return Ok(response);
9746 }
9747 }
9748 }
9749
9750 let admission = self
9752 .admit_disambiguation_redispatch(
9753 expected_disambiguation_epoch,
9754 expected_state_generation,
9755 )
9756 .await?;
9757 *self.pending_skill_id.write() = None;
9758 drop(admission);
9759 let skill_response = self.execute_skill_by_id(skill_id, enriched_input).await?;
9760 self.memory
9761 .add_message(ChatMessage::user(enriched_input))
9762 .await?;
9763 self.handle_skill_response(enriched_input, skill_id, skill_response, &HashMap::new())
9764 .await
9765 }
9766
9767 async fn handle_skill_response(
9770 &self,
9771 processed_input: &str,
9772 skill_id: &str,
9773 skill_response: String,
9774 input_context: &HashMap<String, Value>,
9775 ) -> Result<AgentResponse> {
9776 let output_data = self.process_output(&skill_response, input_context).await?;
9777 let final_response = output_data.content;
9778
9779 self.memory
9780 .add_message(ChatMessage::assistant(&final_response))
9781 .await?;
9782
9783 self.check_memory_compression().await?;
9784
9785 self.increment_turn();
9786 self.evaluate_transitions(processed_input, &final_response)
9787 .await?;
9788
9789 let response = AgentResponse::new(final_response)
9790 .with_metadata("skill_id", serde_json::json!(skill_id));
9791 self.finish_turn_if_root(&response).await?;
9792 Ok(response)
9793 }
9794
9795 async fn handle_plan_and_execute(
9798 &self,
9799 processed_input: &str,
9800 input_context: &HashMap<String, Value>,
9801 auto_detected: bool,
9802 ) -> Result<AgentResponse> {
9803 let effective = self.get_effective_reasoning_config();
9804 let plan_reflection = effective
9805 .get_planning()
9806 .map(|c| c.reflection.clone())
9807 .unwrap_or_default();
9808
9809 let max_attempts = if plan_reflection.enabled {
9810 1 + plan_reflection.max_replans
9811 } else {
9812 1
9813 };
9814
9815 let mut plan = self.generate_plan(processed_input).await?;
9816 info!(
9817 plan_id = %plan.id,
9818 steps = plan.steps.len(),
9819 "Plan generated"
9820 );
9821
9822 let mut plan_result = String::new();
9823
9824 for attempt in 0..max_attempts {
9825 *self.current_plan.write() = Some(plan.clone());
9826 plan_result = self.execute_plan(&mut plan).await?;
9827
9828 info!(
9829 plan_status = ?plan.status,
9830 completed_steps = plan.completed_steps().count(),
9831 attempt = attempt + 1,
9832 "Plan execution completed"
9833 );
9834
9835 if !plan_reflection.enabled {
9836 break;
9837 }
9838
9839 let has_failures = plan.steps.iter().any(|s| s.status.is_failed());
9840 if !has_failures {
9841 break;
9842 }
9843
9844 if attempt + 1 >= max_attempts {
9845 break;
9846 }
9847
9848 match plan_reflection.on_step_failure {
9849 StepFailureAction::Replan => {
9850 info!(attempt = attempt + 1, "Plan had failures, replanning");
9851 plan = self.generate_plan(processed_input).await?;
9852 }
9853 StepFailureAction::Abort => {
9854 warn!("Plan step failed, aborting");
9855 break;
9856 }
9857 StepFailureAction::Skip | StepFailureAction::Continue => {
9858 break;
9859 }
9860 }
9861 }
9862
9863 *self.current_plan.write() = Some(plan);
9864
9865 let output_data = self.process_output(&plan_result, input_context).await?;
9866 let final_content = output_data.content;
9867
9868 self.memory
9869 .add_message(ChatMessage::assistant(&final_content))
9870 .await?;
9871
9872 self.check_memory_compression().await?;
9873 self.increment_turn();
9874 self.evaluate_transitions(processed_input, &final_content)
9875 .await?;
9876
9877 let reasoning_metadata =
9878 ReasoningMetadata::new(ReasoningMode::PlanAndExecute).with_auto_detected(auto_detected);
9879
9880 let response = AgentResponse::new(&final_content).with_metadata(
9881 "reasoning",
9882 serde_json::to_value(&reasoning_metadata).unwrap_or_default(),
9883 );
9884
9885 self.finish_turn_if_root(&response).await?;
9886 Ok(response)
9887 }
9888
9889 fn inject_reasoning_prompt(
9891 &self,
9892 messages: &mut [ChatMessage],
9893 reasoning_mode: &ReasoningMode,
9894 is_first_iteration: bool,
9895 ) {
9896 if !is_first_iteration {
9897 return;
9898 }
9899 match reasoning_mode {
9900 ReasoningMode::CoT => {
9901 if let Some(msg) = messages.first_mut()
9902 && matches!(msg.role, ai_agents_core::Role::System)
9903 {
9904 msg.content = self.build_cot_system_prompt(&msg.content);
9905 debug!("Applied Chain-of-Thought system prompt");
9906 }
9907 }
9908 ReasoningMode::React => {
9909 if let Some(msg) = messages.first_mut()
9910 && matches!(msg.role, ai_agents_core::Role::System)
9911 {
9912 msg.content = self.build_react_system_prompt(&msg.content);
9913 debug!("Applied ReAct system prompt");
9914 }
9915 }
9916 _ => {}
9917 }
9918 }
9919
9920 async fn generate_main_response_draft(
9925 &self,
9926 processed_input: &str,
9927 reasoning_mode: &ReasoningMode,
9928 ) -> Result<MainResponseDraft> {
9929 let llm = self.get_state_llm()?;
9930 let protocol = self.main_tool_protocol(llm.as_ref(), true).await?;
9931 let mut messages = self
9932 .build_messages_internal(false, Some(processed_input), protocol.choice.is_none())
9933 .await?;
9934 self.inject_reasoning_prompt(&mut messages, reasoning_mode, true);
9935 let response = self
9936 .complete_main_llm_with_recovery(llm, &messages, &protocol)
9937 .await?;
9938 let content = response.content.trim().to_string();
9939 let (thinking, answer) = self.extract_thinking(&content);
9940 if let Some(calls) = self.parse_main_tool_calls(&content, &protocol)? {
9941 return Ok(MainResponseDraft::ToolCalls {
9942 raw_content: content,
9943 calls,
9944 thinking,
9945 });
9946 }
9947 Ok(MainResponseDraft::Text {
9948 raw_content: answer,
9949 thinking,
9950 })
9951 }
9952
9953 async fn commit_main_response_draft(
9958 &self,
9959 processed_input: &str,
9960 input_context: &HashMap<String, Value>,
9961 draft: MainResponseDraft,
9962 reasoning_mode: ReasoningMode,
9963 auto_detected: bool,
9964 ) -> Result<AgentResponse> {
9965 self.commit_root_user_message(processed_input).await?;
9966 match draft {
9967 MainResponseDraft::Text {
9968 raw_content,
9969 thinking,
9970 } => {
9971 self.finish_text_response_from_model(CommittedTextResponse {
9972 processed_input,
9973 input_context,
9974 answer: raw_content,
9975 reasoning_mode,
9976 auto_detected,
9977 iterations: 1,
9978 thinking_content: thinking,
9979 all_tool_calls: Vec::new(),
9980 })
9981 .await
9982 }
9983 MainResponseDraft::ToolCalls {
9984 raw_content,
9985 calls,
9986 thinking: _,
9987 } => {
9988 let mut all_tool_calls = Vec::new();
9989 match self
9990 .handle_tool_calls(
9991 processed_input,
9992 &raw_content,
9993 calls,
9994 &mut all_tool_calls,
9995 None,
9996 )
9997 .await?
9998 {
9999 ToolCallOutcome::Rejected(response) => {
10000 self.finish_turn_if_root(&response).await?;
10001 Ok(response)
10002 }
10003 ToolCallOutcome::Continue | ToolCallOutcome::TransitionFired => {
10004 self.continue_after_committed_tool_draft(processed_input)
10005 .await
10006 }
10007 }
10008 }
10009 }
10010 }
10011
10012 async fn continue_after_committed_tool_draft(
10017 &self,
10018 processed_input: &str,
10019 ) -> Result<AgentResponse> {
10020 *self.redispatch_depth.write() += 1;
10021 if let Some(context) = self.active_turn_context.write().as_mut() {
10022 context.enter_redispatch();
10023 }
10024 let result = Box::pin(self.run_loop_internal(processed_input)).await;
10025 *self.redispatch_depth.write() -= 1;
10026 if let Some(context) = self.active_turn_context.write().as_mut() {
10027 context.exit_redispatch();
10028 }
10029 let response = result?;
10030 self.finish_turn_if_root(&response).await?;
10031 Ok(response)
10032 }
10033
10034 async fn finish_text_response_from_model(
10039 &self,
10040 response: CommittedTextResponse<'_>,
10041 ) -> Result<AgentResponse> {
10042 let CommittedTextResponse {
10043 processed_input,
10044 input_context,
10045 answer,
10046 reasoning_mode,
10047 auto_detected,
10048 iterations,
10049 thinking_content,
10050 all_tool_calls,
10051 } = response;
10052 let output_data = self.process_output(&answer, input_context).await?;
10053 let mut final_content = if output_data.metadata.rejected {
10054 output_data
10055 .metadata
10056 .rejection_reason
10057 .unwrap_or_else(|| answer.to_string())
10058 } else {
10059 output_data.content
10060 };
10061 let llm = self.get_state_llm()?;
10062 let reflection_metadata;
10063 (final_content, reflection_metadata) = self
10064 .run_reflection(&*llm, processed_input, final_content)
10065 .await?;
10066 final_content =
10067 self.format_response_with_thinking(thinking_content.as_deref(), &final_content);
10068 let final_content = {
10069 let result = self
10070 .post_loop_processing(processed_input, final_content)
10071 .await?;
10072 self.apply_post_loop_result(processed_input, result)
10073 .await?
10074 .content
10075 };
10076 let response = self.build_agent_response(AgentResponseParts {
10077 content: final_content,
10078 all_tool_calls,
10079 reasoning_mode,
10080 auto_detected,
10081 iterations,
10082 thinking: thinking_content,
10083 reflection_metadata,
10084 });
10085 self.finish_turn_if_root(&response).await?;
10086 Ok(response)
10087 }
10088
10089 async fn run_committed_response_loop_with_reasoning(
10094 &self,
10095 processed_input: &str,
10096 input_context: &HashMap<String, Value>,
10097 reasoning_mode: ReasoningMode,
10098 auto_detected: bool,
10099 ) -> Result<AgentResponse> {
10100 self.commit_root_user_message(processed_input).await?;
10101 let llm = self.get_state_llm()?;
10102 let mut iterations = 0u32;
10103 let mut all_tool_calls = Vec::new();
10104 let mut thinking_content = None;
10105 loop {
10106 let effective_max = if reasoning_mode != ReasoningMode::None {
10107 let rc = self.get_effective_reasoning_config();
10108 self.max_iterations.min(rc.max_iterations)
10109 } else {
10110 self.max_iterations
10111 };
10112 if iterations >= effective_max {
10113 return Err(AgentError::Other(format!(
10114 "Max iterations ({}) exceeded",
10115 effective_max
10116 )));
10117 }
10118 iterations += 1;
10119 *self.iteration_count.write() = iterations;
10120 let protocol = self.main_tool_protocol(llm.as_ref(), false).await?;
10121 let mut messages = self
10122 .build_messages_internal(true, None, protocol.choice.is_none())
10123 .await?;
10124 self.inject_reasoning_prompt(&mut messages, &reasoning_mode, iterations == 1);
10125 self.hooks.on_llm_start(&messages).await;
10126 let llm_start = Instant::now();
10127 let response = self
10128 .complete_main_llm_with_recovery(Arc::clone(&llm), &messages, &protocol)
10129 .await?;
10130 let llm_duration_ms = llm_start.elapsed().as_millis() as u64;
10131 self.hooks.on_llm_complete(&response, llm_duration_ms).await;
10132 let content = response.content.trim();
10133 if let Some(tool_calls) = self.parse_main_tool_calls(content, &protocol)? {
10134 match self
10135 .handle_tool_calls(
10136 processed_input,
10137 content,
10138 tool_calls,
10139 &mut all_tool_calls,
10140 None,
10141 )
10142 .await?
10143 {
10144 ToolCallOutcome::Continue | ToolCallOutcome::TransitionFired => continue,
10145 ToolCallOutcome::Rejected(resp) => {
10146 self.finish_turn_if_root(&resp).await?;
10147 return Ok(resp);
10148 }
10149 }
10150 }
10151 let (extracted_thinking, answer) = self.extract_thinking(content);
10152 if extracted_thinking.is_some() {
10153 thinking_content = extracted_thinking;
10154 }
10155 return self
10156 .finish_text_response_from_model(CommittedTextResponse {
10157 processed_input,
10158 input_context,
10159 answer,
10160 reasoning_mode,
10161 auto_detected,
10162 iterations,
10163 thinking_content,
10164 all_tool_calls,
10165 })
10166 .await;
10167 }
10168 }
10169
10170 async fn handle_tool_calls(
10176 &self,
10177 processed_input: &str,
10178 content: &str,
10179 tool_calls: Vec<ToolCall>,
10180 all_tool_calls: &mut Vec<ToolCall>,
10181 mut events: Option<&mut Vec<StreamChunk>>,
10182 ) -> Result<ToolCallOutcome> {
10183 let include_tool_events = self.streaming.include_tool_events;
10184 let transition_content = native_readable_projection(content)
10188 .map_err(|error| AgentError::LLM(error.to_string()))?;
10189 let transition_fired = self
10190 .evaluate_transitions(processed_input, &transition_content)
10191 .await?;
10192 if transition_fired {
10193 self.memory
10194 .add_message(ChatMessage::assistant(
10195 "(Transitioned to new state — tool call handled by workflow)",
10196 ))
10197 .await?;
10198 if let Some(events) = events.as_deref_mut()
10199 && self.streaming.include_state_events
10200 && let Some(state) = self.current_state()
10201 {
10202 events.push(StreamChunk::state_transition(None, state));
10203 }
10204 return Ok(ToolCallOutcome::TransitionFired);
10205 }
10206
10207 self.memory
10209 .add_message(ChatMessage::assistant(content))
10210 .await?;
10211 self.remember_committed_native_exchange(content).await?;
10212 let native_tool_call = Self::is_native_tool_call_content(content)?;
10213
10214 if let Some(events) = events.as_deref_mut()
10215 && include_tool_events
10216 {
10217 for tool_call in &tool_calls {
10218 events.push(StreamChunk::tool_start(&tool_call.id, &tool_call.name));
10219 }
10220 }
10221 let results = self.execute_tools_parallel(&tool_calls).await;
10222 let mut rejection = None;
10223
10224 for ((_id, result), tool_call) in results.into_iter().zip(tool_calls.iter()) {
10225 match result {
10226 Ok(output) => {
10227 if let Some(events) = events.as_deref_mut()
10228 && include_tool_events
10229 {
10230 events.push(StreamChunk::tool_result(
10231 &tool_call.id,
10232 &tool_call.name,
10233 &output,
10234 true,
10235 ));
10236 }
10237 self.memory
10238 .add_message(Self::tool_result_message(
10239 tool_call,
10240 &output,
10241 native_tool_call,
10242 )?)
10243 .await?;
10244 }
10245 Err(e) => {
10246 if matches!(e, AgentError::HITLRejected(_)) {
10247 if !native_tool_call {
10248 self.memory
10249 .add_message(ChatMessage::assistant(format!(
10250 "The operation was rejected by the approver: {e}"
10251 )))
10252 .await?;
10253 return Ok(ToolCallOutcome::Rejected(AgentResponse {
10254 content: format!("Operation cancelled: {e}"),
10255 metadata: None,
10256 tool_calls: Some(all_tool_calls.clone()),
10257 }));
10258 }
10259 if rejection.is_none() {
10260 rejection = Some(e.to_string());
10261 }
10262 }
10263 if let Some(events) = events.as_deref_mut()
10264 && include_tool_events
10265 {
10266 events.push(StreamChunk::tool_result(
10267 &tool_call.id,
10268 &tool_call.name,
10269 e.to_string(),
10270 false,
10271 ));
10272 }
10273 self.memory
10274 .add_message(Self::tool_result_message(
10275 tool_call,
10276 &format!("Error: {}", e),
10277 native_tool_call,
10278 )?)
10279 .await?;
10280 }
10281 }
10282 all_tool_calls.push(tool_call.clone());
10283 if let Some(events) = events.as_deref_mut()
10284 && include_tool_events
10285 {
10286 events.push(StreamChunk::tool_end(&tool_call.id));
10287 }
10288 }
10289 if let Some(rejection) = rejection {
10290 self.memory
10291 .add_message(ChatMessage::assistant(format!(
10292 "The operation was rejected by the approver: {rejection}"
10293 )))
10294 .await?;
10295 return Ok(ToolCallOutcome::Rejected(AgentResponse {
10296 content: format!("Operation cancelled: {rejection}"),
10297 metadata: None,
10298 tool_calls: Some(all_tool_calls.clone()),
10299 }));
10300 }
10301 Ok(ToolCallOutcome::Continue)
10302 }
10303
10304 async fn run_reflection(
10306 &self,
10307 llm: &dyn LLMProvider,
10308 processed_input: &str,
10309 mut content: String,
10310 ) -> Result<(String, Option<ReflectionMetadata>)> {
10311 let config = self.get_effective_reflection_config();
10312 let should_reflect = self
10313 .should_reflect_with_config(processed_input, &content, &config)
10314 .await?;
10315 if !should_reflect {
10316 return Ok((content, None));
10317 }
10318
10319 info!("Starting response reflection evaluation");
10320 let mut attempts = 0u32;
10321 let max_retries = config.max_retries;
10322 let mut history: Vec<ReflectionAttempt> = Vec::new();
10323
10324 loop {
10325 let evaluation = self
10326 .evaluate_response_with_config(processed_input, &content, &config)
10327 .await?;
10328
10329 if evaluation.passed || attempts >= max_retries {
10330 info!(
10331 passed = evaluation.passed,
10332 confidence = evaluation.confidence,
10333 attempts = attempts + 1,
10334 "Reflection evaluation complete"
10335 );
10336 let reflection_metadata = Some(
10337 ReflectionMetadata::new(evaluation)
10338 .with_attempts(attempts + 1)
10339 .with_history(history),
10340 );
10341 return Ok((content, reflection_metadata));
10342 }
10343
10344 debug!(
10345 attempt = attempts + 1,
10346 failed_criteria = evaluation.failed_criteria().count(),
10347 "Response did not meet criteria, retrying"
10348 );
10349
10350 history.push(
10351 ReflectionAttempt::new(&content, evaluation.clone())
10352 .with_feedback("Response did not meet quality criteria"),
10353 );
10354
10355 let feedback: Vec<String> = evaluation
10356 .failed_criteria()
10357 .map(|c| format!("- {}", c.criterion))
10358 .collect();
10359
10360 let retry_prompt = format!(
10361 "Your previous response did not meet these criteria:\n{}\n\nPlease provide an improved response.",
10362 feedback.join("\n")
10363 );
10364
10365 self.memory
10366 .add_message(ChatMessage::user(&retry_prompt))
10367 .await?;
10368
10369 let retry_messages = self.build_messages().await?;
10370 let retry_response = self
10371 .observe_purpose(
10372 ObservationPurpose::ReflectionEvaluation,
10373 llm.complete(&retry_messages, None),
10374 )
10375 .await
10376 .map_err(|e| AgentError::LLM(e.to_string()))?;
10377
10378 content = retry_response.content.trim().to_string();
10379 attempts += 1;
10380 }
10381 }
10382
10383 async fn post_loop_processing(
10386 &self,
10387 processed_input: &str,
10388 content: String,
10389 ) -> Result<PostLoopResult> {
10390 self.increment_turn();
10395
10396 self.run_context_extractors(processed_input).await;
10398
10399 let transitioned = self.evaluate_transitions(processed_input, &content).await?;
10400
10401 if !transitioned {
10402 self.memory
10403 .add_message(ChatMessage::assistant(&content))
10404 .await?;
10405 self.check_memory_compression().await?;
10406 return Ok(PostLoopResult::NoTransition(content));
10407 }
10408
10409 if !self.should_regenerate_after_transition() {
10411 self.memory
10412 .add_message(ChatMessage::assistant(&content))
10413 .await?;
10414 self.check_memory_compression().await?;
10415 return Ok(PostLoopResult::Transitioned {
10416 content,
10417 regenerated: false,
10418 });
10419 }
10420
10421 if self.needs_redispatch_for_new_state() {
10425 info!("Post-transition NeedsRedispatch: new state requires full dispatch");
10426 return Ok(PostLoopResult::NeedsRedispatch);
10429 }
10430
10431 self.memory
10434 .add_message(ChatMessage::assistant(&content))
10435 .await?;
10436 self.check_memory_compression().await?;
10437
10438 let new_llm = self.get_state_llm()?;
10444 let mut final_content;
10445
10446 for post_iter in 0..self.max_iterations {
10447 let protocol = self.main_tool_protocol(new_llm.as_ref(), false).await?;
10448 let new_messages = self
10449 .build_messages_internal(true, None, protocol.choice.is_none())
10450 .await?;
10451 if post_iter == 0
10452 && let Some(system_msg) = new_messages.first()
10453 && system_msg.role == ai_agents_core::Role::System
10454 {
10455 debug!(
10456 prompt_preview =
10457 &system_msg.content[system_msg.content.len().saturating_sub(200)..],
10458 "Post-transition system prompt (last 200 chars)"
10459 );
10460 }
10461
10462 let new_response = self
10463 .complete_main_llm_with_recovery(Arc::clone(&new_llm), &new_messages, &protocol)
10464 .await?;
10465 final_content = new_response.content.trim().to_string();
10466
10467 if let Some(tool_calls) = self.parse_main_tool_calls(&final_content, &protocol)? {
10470 let native_tool_call = Self::is_native_tool_call_content(&final_content)?;
10471 debug!(
10472 post_iter = post_iter,
10473 tools = tool_calls.len(),
10474 "Post-transition tool call detected, executing"
10475 );
10476
10477 self.memory
10478 .add_message(ChatMessage::assistant(&final_content))
10479 .await?;
10480 self.remember_committed_native_exchange(&final_content)
10481 .await?;
10482
10483 let results = self.execute_tools_parallel(&tool_calls).await;
10484 let mut rejection = None;
10485 for ((_id, result), tool_call) in results.into_iter().zip(tool_calls.iter()) {
10486 match result {
10487 Ok(output) => {
10488 self.memory
10489 .add_message(Self::tool_result_message(
10490 tool_call,
10491 &output,
10492 native_tool_call,
10493 )?)
10494 .await?;
10495 }
10496 Err(e) => {
10497 if native_tool_call
10498 && rejection.is_none()
10499 && matches!(e, AgentError::HITLRejected(_))
10500 {
10501 rejection = Some(e.to_string());
10502 }
10503 self.memory
10504 .add_message(Self::tool_result_message(
10505 tool_call,
10506 &format!("Error: {}", e),
10507 native_tool_call,
10508 )?)
10509 .await?;
10510 }
10511 }
10512 }
10513 if let Some(rejection) = rejection {
10514 self.memory
10515 .add_message(ChatMessage::assistant(format!(
10516 "The operation was rejected by the approver: {rejection}"
10517 )))
10518 .await?;
10519 return Err(AgentError::HITLRejected(rejection));
10520 }
10521 continue;
10523 }
10524
10525 self.memory
10527 .add_message(ChatMessage::assistant(&final_content))
10528 .await?;
10529 return Ok(PostLoopResult::Transitioned {
10530 content: final_content,
10531 regenerated: true,
10532 });
10533 }
10534
10535 final_content = "Post-transition processing completed.".to_string();
10537 self.memory
10538 .add_message(ChatMessage::assistant(&final_content))
10539 .await?;
10540
10541 Ok(PostLoopResult::Transitioned {
10542 content: final_content,
10543 regenerated: true,
10544 })
10545 }
10546
10547 fn should_regenerate_after_transition(&self) -> bool {
10550 if let Some(ref sm) = self.state_machine {
10551 if !sm.config().regenerate_on_transition {
10553 return false;
10554 }
10555 if let Some(def) = sm.current_definition()
10557 && let Some(regen) = def.regenerate_on_enter
10558 {
10559 return regen;
10560 }
10561 }
10562 true
10563 }
10564
10565 fn needs_redispatch_for_new_state(&self) -> bool {
10568 if let Some(ref sm) = self.state_machine
10569 && let Some(def) = sm.current_definition()
10570 {
10571 if def.concurrent.is_some()
10572 || def.group_chat.is_some()
10573 || def.pipeline.is_some()
10574 || def.handoff.is_some()
10575 || def.delegate.is_some()
10576 {
10577 return true;
10578 }
10579 let effective = self.get_effective_reasoning_config();
10581 if !matches!(effective.mode, ReasoningMode::None) {
10582 return true;
10583 }
10584 }
10585 false
10586 }
10587
10588 async fn apply_post_loop_result(
10594 &self,
10595 processed_input: &str,
10596 result: PostLoopResult,
10597 ) -> Result<AppliedPostLoop> {
10598 match result {
10599 PostLoopResult::NoTransition(content) => Ok(AppliedPostLoop {
10600 content,
10601 transitioned: false,
10602 regenerated: false,
10603 }),
10604 PostLoopResult::Transitioned {
10605 content,
10606 regenerated,
10607 } => Ok(AppliedPostLoop {
10608 content,
10609 transitioned: true,
10610 regenerated,
10611 }),
10612 PostLoopResult::NeedsRedispatch => {
10613 const MAX_REDISPATCH_DEPTH: u32 = 3;
10614 let current_depth = *self.redispatch_depth.read();
10615 if current_depth >= MAX_REDISPATCH_DEPTH {
10616 warn!(
10617 depth = current_depth,
10618 "Post-transition re-dispatch depth limit reached, returning empty response"
10619 );
10620 let content = String::new();
10621 self.memory
10622 .add_message(ChatMessage::assistant(&content))
10623 .await?;
10624 return Ok(AppliedPostLoop {
10626 content,
10627 transitioned: true,
10628 regenerated: false,
10629 });
10630 }
10631 *self.redispatch_depth.write() += 1;
10632 if let Some(context) = self.active_turn_context.write().as_mut() {
10633 context.enter_redispatch();
10634 }
10635 info!(
10636 depth = current_depth + 1,
10637 "Re-dispatching for new state after transition"
10638 );
10639 let resp = Box::pin(self.run_loop_internal(processed_input)).await;
10640 *self.redispatch_depth.write() -= 1;
10641 if let Some(context) = self.active_turn_context.write().as_mut() {
10642 context.exit_redispatch();
10643 }
10644 resp.map(|r| AppliedPostLoop {
10645 content: r.content,
10646 transitioned: true,
10647 regenerated: true,
10648 })
10649 }
10650 }
10651 }
10652
10653 fn build_agent_response(&self, parts: AgentResponseParts) -> AgentResponse {
10655 let AgentResponseParts {
10656 content,
10657 all_tool_calls,
10658 reasoning_mode,
10659 auto_detected,
10660 iterations,
10661 thinking,
10662 reflection_metadata,
10663 } = parts;
10664 let reasoning_metadata = ReasoningMetadata::new(reasoning_mode.clone())
10665 .with_thinking(thinking.clone().unwrap_or_default())
10666 .with_iterations(iterations)
10667 .with_auto_detected(auto_detected);
10668
10669 let mut response = AgentResponse::new(&content);
10670 if !all_tool_calls.is_empty() {
10671 response = response.with_tool_calls(all_tool_calls);
10672 }
10673
10674 if let Some(state) = self.current_state() {
10675 response = response.with_metadata("current_state", serde_json::json!(state));
10676 }
10677
10678 response = response.with_metadata(
10679 "reasoning",
10680 serde_json::to_value(&reasoning_metadata).unwrap_or_default(),
10681 );
10682
10683 if let Some(ref refl_meta) = reflection_metadata {
10684 response = response.with_metadata(
10685 "reflection",
10686 serde_json::to_value(refl_meta).unwrap_or_default(),
10687 );
10688 }
10689
10690 response
10691 }
10692
10693 async fn handle_delegated_state(
10695 &self,
10696 input: &str,
10697 delegate_id: &str,
10698 state_def: &ai_agents_state::StateDefinition,
10699 ) -> Result<AgentResponse> {
10700 use std::time::Instant;
10701
10702 let registry = self.spawner_registry.as_ref().ok_or_else(|| {
10703 AgentError::Config(format!(
10704 "State delegates to '{}' but no agent registry is configured. \
10705 Add a spawner section with auto_spawn to your YAML.",
10706 delegate_id
10707 ))
10708 })?;
10709
10710 let state_name = self
10711 .state_machine
10712 .as_ref()
10713 .map(|sm| sm.current())
10714 .unwrap_or_else(|| "unknown".to_string());
10715
10716 self.hooks.on_delegate_start(delegate_id, &state_name).await;
10717 let start = Instant::now();
10718
10719 let delegate = registry.get(delegate_id).ok_or_else(|| {
10720 AgentError::Other(format!(
10721 "State '{}' delegates to '{}' but no agent with that ID exists in the registry.",
10722 state_name, delegate_id
10723 ))
10724 })?;
10725
10726 let context_mode = state_def.delegate_context.clone().unwrap_or_default();
10728 let effective_input = self
10729 .observe_purpose(
10730 ObservationPurpose::OrchestrationRouting,
10731 crate::orchestration::context::prepare_delegate_input(
10732 input,
10733 &context_mode,
10734 &*self.memory,
10735 self.llm_registry.get("router").ok().as_deref(),
10736 ),
10737 )
10738 .await?;
10739
10740 let response = delegate
10741 .chat_with_actor_context(&effective_input, self.outbound_actor_context())
10742 .await?;
10743
10744 let duration_ms = start.elapsed().as_millis() as u64;
10745 self.hooks
10746 .on_delegate_complete(delegate_id, &state_name, duration_ms)
10747 .await;
10748
10749 let ctx_key = format!("delegation.{}.last_response", delegate_id);
10751 let _ = self.context_manager.set(
10752 &ctx_key,
10753 serde_json::Value::String(response.content.clone()),
10754 );
10755
10756 let _ = self.context_manager.set(
10758 "orchestration",
10759 serde_json::json!({
10760 "type": "delegate",
10761 "agent": delegate_id,
10762 "state": state_name,
10763 "response": response.content,
10764 "duration_ms": duration_ms,
10765 }),
10766 );
10767
10768 self.commit_root_user_message(input).await?;
10769
10770 let post_result = self
10773 .post_loop_processing(
10774 input,
10775 format!("[Delegated to {}]: {}", delegate_id, response.content),
10776 )
10777 .await?;
10778 let final_content = self
10779 .apply_post_loop_result(input, post_result)
10780 .await?
10781 .content;
10782
10783 let mut result = AgentResponse::new(final_content);
10784
10785 let metadata = serde_json::json!({
10786 "orchestration": {
10787 "type": "delegate",
10788 "agent": delegate_id,
10789 "state": state_name,
10790 "response": response.content,
10791 "duration_ms": duration_ms,
10792 }
10793 });
10794 result.metadata = Some(
10795 serde_json::from_value::<std::collections::HashMap<String, serde_json::Value>>(
10796 metadata,
10797 )
10798 .unwrap_or_default(),
10799 );
10800
10801 self.finish_turn_if_root(&result).await?;
10802 Ok(result)
10803 }
10804
10805 async fn handle_concurrent_state(
10807 &self,
10808 input: &str,
10809 config: &ai_agents_state::ConcurrentStateConfig,
10810 ) -> Result<AgentResponse> {
10811 use std::time::Instant;
10812
10813 let registry = self.spawner_registry.as_ref().ok_or_else(|| {
10814 AgentError::Config(
10815 "Concurrent state requires an agent registry. Add a spawner section.".into(),
10816 )
10817 })?;
10818
10819 let context_mode = config.context_mode.clone().unwrap_or_default();
10824 let context_input = self
10825 .observe_purpose(
10826 ObservationPurpose::OrchestrationRouting,
10827 crate::orchestration::context::prepare_delegate_input(
10828 input,
10829 &context_mode,
10830 &*self.memory,
10831 self.llm_registry.get("router").ok().as_deref(),
10832 ),
10833 )
10834 .await?;
10835
10836 let effective_input = if let Some(ref tmpl) = config.input {
10837 render_concurrent_template(tmpl, &context_input, &self.build_context_with_overlays())
10838 .unwrap_or_else(|_| context_input.clone())
10839 } else {
10840 context_input
10841 };
10842
10843 let start = Instant::now();
10844
10845 let llm_name = config
10846 .aggregation
10847 .synthesizer_llm
10848 .as_deref()
10849 .unwrap_or("router");
10850 let llm_provider = self.llm_registry.get(llm_name).ok();
10851
10852 let vote_parallelism = if self.runtime_config.optimization.enabled
10853 && self
10854 .runtime_config
10855 .optimization
10856 .parallel_orchestration_vote_extraction
10857 {
10858 Some(self.runtime_config.optimization.max_parallel_runtime_tasks)
10859 } else {
10860 None
10861 };
10862
10863 let result = self
10864 .observe_purpose(
10865 ObservationPurpose::OrchestrationAggregation,
10866 scope_actor_context(
10867 self.outbound_actor_context(),
10868 crate::orchestration::concurrent(
10869 registry,
10870 &effective_input,
10871 &config.agents,
10872 &config.aggregation,
10873 llm_provider.as_deref(),
10874 config.min_required,
10875 config.timeout_ms,
10876 config.on_partial_failure.clone(),
10877 vote_parallelism,
10878 ),
10879 ),
10880 )
10881 .await?;
10882
10883 let duration_ms = start.elapsed().as_millis() as u64;
10884 let agent_ids: Vec<String> = config.agents.iter().map(|a| a.id().to_string()).collect();
10885 let strategy = format!("{:?}", config.aggregation.strategy);
10886 self.hooks
10887 .on_concurrent_complete(&agent_ids, &strategy, duration_ms)
10888 .await;
10889
10890 let _ = self.context_manager.set(
10892 "concurrent.result",
10893 serde_json::Value::String(result.response.content.clone()),
10894 );
10895
10896 let agents_json: Vec<serde_json::Value> = result
10898 .agent_results
10899 .iter()
10900 .map(|ar| {
10901 serde_json::json!({
10902 "id": ar.agent_id,
10903 "response": ar.response.as_ref().map(|r| r.content.as_str()),
10904 "success": ar.success,
10905 "error": ar.error,
10906 "duration_ms": ar.duration_ms,
10907 })
10908 })
10909 .collect();
10910
10911 let _ = self.context_manager.set(
10913 "orchestration",
10914 serde_json::json!({
10915 "type": "concurrent",
10916 "result": result.response.content,
10917 "strategy": strategy,
10918 "agents": agents_json,
10919 "duration_ms": duration_ms,
10920 }),
10921 );
10922
10923 self.commit_root_user_message(input).await?;
10924
10925 let post_result = self
10926 .post_loop_processing(input, result.response.content.clone())
10927 .await?;
10928 let final_content = self
10929 .apply_post_loop_result(input, post_result)
10930 .await?
10931 .content;
10932
10933 let mut response = AgentResponse::new(final_content);
10934 let metadata = serde_json::json!({
10935 "orchestration": {
10936 "type": "concurrent",
10937 "result": result.response.content,
10938 "strategy": strategy,
10939 "agents": agents_json,
10940 "duration_ms": duration_ms,
10941 }
10942 });
10943 response.metadata = Some(
10944 serde_json::from_value::<std::collections::HashMap<String, serde_json::Value>>(
10945 metadata,
10946 )
10947 .unwrap_or_default(),
10948 );
10949
10950 self.finish_turn_if_root(&response).await?;
10951 Ok(response)
10952 }
10953
10954 async fn handle_group_chat_state(
10956 &self,
10957 input: &str,
10958 config: &ai_agents_state::GroupChatStateConfig,
10959 ) -> Result<AgentResponse> {
10960 use std::time::Instant;
10961
10962 let registry = self.spawner_registry.as_ref().ok_or_else(|| {
10963 AgentError::Config(
10964 "Group chat state requires an agent registry. Add a spawner section.".into(),
10965 )
10966 })?;
10967
10968 let start = Instant::now();
10969
10970 let llm_provider = self.llm_registry.get("router").ok();
10971
10972 let context_mode = config.context_mode.clone().unwrap_or_default();
10974 let context_input = self
10975 .observe_purpose(
10976 ObservationPurpose::OrchestrationRouting,
10977 crate::orchestration::context::prepare_delegate_input(
10978 input,
10979 &context_mode,
10980 &*self.memory,
10981 self.llm_registry.get("router").ok().as_deref(),
10982 ),
10983 )
10984 .await?;
10985
10986 let effective_topic = if let Some(ref tmpl) = config.input {
10988 render_concurrent_template(tmpl, &context_input, &self.build_context_with_overlays())
10989 .unwrap_or_else(|_| context_input.clone())
10990 } else {
10991 context_input
10992 };
10993
10994 let result = self
10995 .observe_purpose(
10996 ObservationPurpose::OrchestrationConversation,
10997 scope_actor_context(
10998 self.outbound_actor_context(),
10999 crate::orchestration::group_chat(
11000 registry,
11001 &effective_topic,
11002 config,
11003 llm_provider.as_deref(),
11004 Some(&*self.hooks),
11005 ),
11006 ),
11007 )
11008 .await?;
11009
11010 let duration_ms = start.elapsed().as_millis() as u64;
11011
11012 let _ = self.context_manager.set(
11014 "group_chat.conclusion",
11015 serde_json::Value::String(result.response.content.clone()),
11016 );
11017
11018 let transcript_json: Vec<serde_json::Value> = result
11020 .transcript
11021 .iter()
11022 .map(|t| {
11023 serde_json::json!({
11024 "speaker": t.speaker,
11025 "round": t.round,
11026 "content": t.content,
11027 })
11028 })
11029 .collect();
11030
11031 let _ = self.context_manager.set(
11033 "orchestration",
11034 serde_json::json!({
11035 "type": "group_chat",
11036 "conclusion": result.response.content,
11037 "transcript": transcript_json,
11038 "rounds": result.rounds_completed,
11039 "termination": result.termination_reason,
11040 "duration_ms": duration_ms,
11041 }),
11042 );
11043
11044 self.commit_root_user_message(input).await?;
11045
11046 let post_result = self
11047 .post_loop_processing(input, result.response.content.clone())
11048 .await?;
11049 let final_content = self
11050 .apply_post_loop_result(input, post_result)
11051 .await?
11052 .content;
11053
11054 let mut response = AgentResponse::new(final_content);
11055 let metadata = serde_json::json!({
11056 "orchestration": {
11057 "type": "group_chat",
11058 "conclusion": result.response.content,
11059 "transcript": transcript_json,
11060 "rounds": result.rounds_completed,
11061 "termination": result.termination_reason,
11062 "duration_ms": duration_ms,
11063 }
11064 });
11065 response.metadata = Some(
11066 serde_json::from_value::<std::collections::HashMap<String, serde_json::Value>>(
11067 metadata,
11068 )
11069 .unwrap_or_default(),
11070 );
11071
11072 self.finish_turn_if_root(&response).await?;
11073 Ok(response)
11074 }
11075
11076 async fn handle_pipeline_state(
11078 &self,
11079 input: &str,
11080 config: &ai_agents_state::PipelineStateConfig,
11081 ) -> Result<AgentResponse> {
11082 use std::time::Instant;
11083
11084 let registry = self.spawner_registry.as_ref().ok_or_else(|| {
11085 AgentError::Config(
11086 "Pipeline state requires an agent registry. Add a spawner section.".into(),
11087 )
11088 })?;
11089
11090 let start = Instant::now();
11091
11092 let stages: Vec<crate::orchestration::PipelineStage> = config
11093 .stages
11094 .iter()
11095 .map(|entry| {
11096 let mut stage = crate::orchestration::PipelineStage::id(entry.id());
11097 if let Some(tmpl) = entry.input() {
11098 stage = stage.with_input(tmpl);
11099 }
11100 stage
11101 })
11102 .collect();
11103
11104 let context_mode = config.context_mode.clone().unwrap_or_default();
11106 let context_input = self
11107 .observe_purpose(
11108 ObservationPurpose::OrchestrationRouting,
11109 crate::orchestration::context::prepare_delegate_input(
11110 input,
11111 &context_mode,
11112 &*self.memory,
11113 self.llm_registry.get("router").ok().as_deref(),
11114 ),
11115 )
11116 .await?;
11117
11118 let context_values = self.build_context_with_overlays();
11119 let result = self
11120 .observe_purpose(
11121 ObservationPurpose::OrchestrationRouting,
11122 scope_actor_context(
11123 self.outbound_actor_context(),
11124 crate::orchestration::pipeline(
11125 registry,
11126 &context_input,
11127 &stages,
11128 config.timeout_ms,
11129 Some(&*self.hooks),
11130 Some(&context_values),
11131 ),
11132 ),
11133 )
11134 .await?;
11135
11136 let duration_ms = start.elapsed().as_millis() as u64;
11137
11138 let _ = self.context_manager.set(
11140 "pipeline.result",
11141 serde_json::Value::String(result.response.content.clone()),
11142 );
11143
11144 let stages_json: Vec<serde_json::Value> = result
11146 .stage_outputs
11147 .iter()
11148 .map(|s| {
11149 serde_json::json!({
11150 "agent_id": s.agent_id,
11151 "output": s.output,
11152 "duration_ms": s.duration_ms,
11153 "skipped": s.skipped,
11154 })
11155 })
11156 .collect();
11157
11158 let _ = self.context_manager.set(
11160 "orchestration",
11161 serde_json::json!({
11162 "type": "pipeline",
11163 "result": result.response.content,
11164 "stages": stages_json,
11165 "duration_ms": duration_ms,
11166 }),
11167 );
11168
11169 self.commit_root_user_message(input).await?;
11170
11171 let post_result = self
11172 .post_loop_processing(input, result.response.content.clone())
11173 .await?;
11174 let final_content = self
11175 .apply_post_loop_result(input, post_result)
11176 .await?
11177 .content;
11178
11179 let mut response = AgentResponse::new(final_content);
11180 let metadata = serde_json::json!({
11181 "orchestration": {
11182 "type": "pipeline",
11183 "result": result.response.content,
11184 "stages": stages_json,
11185 "duration_ms": duration_ms,
11186 }
11187 });
11188 response.metadata = Some(
11189 serde_json::from_value::<std::collections::HashMap<String, serde_json::Value>>(
11190 metadata,
11191 )
11192 .unwrap_or_default(),
11193 );
11194
11195 self.finish_turn_if_root(&response).await?;
11196 Ok(response)
11197 }
11198
11199 async fn handle_handoff_state(
11201 &self,
11202 input: &str,
11203 config: &ai_agents_state::HandoffStateConfig,
11204 ) -> Result<AgentResponse> {
11205 use std::time::Instant;
11206
11207 let registry = self.spawner_registry.as_ref().ok_or_else(|| {
11208 AgentError::Config(
11209 "Handoff state requires an agent registry. Add a spawner section.".into(),
11210 )
11211 })?;
11212
11213 let llm = self
11214 .llm_registry
11215 .get("router")
11216 .map_err(|_| AgentError::Config("Handoff state requires a router LLM.".into()))?;
11217
11218 let start = Instant::now();
11219
11220 let context_mode = config.context_mode.clone().unwrap_or_default();
11222 let context_input = self
11223 .observe_purpose(
11224 ObservationPurpose::OrchestrationRouting,
11225 crate::orchestration::context::prepare_delegate_input(
11226 input,
11227 &context_mode,
11228 &*self.memory,
11229 self.llm_registry.get("router").ok().as_deref(),
11230 ),
11231 )
11232 .await?;
11233
11234 let effective_input = if let Some(ref tmpl) = config.input {
11236 render_concurrent_template(tmpl, &context_input, &self.build_context_with_overlays())
11237 .unwrap_or_else(|_| context_input.clone())
11238 } else {
11239 context_input
11240 };
11241
11242 let result = self
11243 .observe_purpose(
11244 ObservationPurpose::OrchestrationRouting,
11245 scope_actor_context(
11246 self.outbound_actor_context(),
11247 crate::orchestration::handoff(
11248 registry,
11249 &effective_input,
11250 &config.initial_agent,
11251 &config.available_agents,
11252 config.max_handoffs,
11253 llm.as_ref(),
11254 Some(&*self.hooks),
11255 ),
11256 ),
11257 )
11258 .await?;
11259
11260 let duration_ms = start.elapsed().as_millis() as u64;
11261
11262 let _ = self.context_manager.set(
11264 "handoff.result",
11265 serde_json::Value::String(result.response.content.clone()),
11266 );
11267
11268 let chain_json: Vec<serde_json::Value> = result
11270 .handoff_chain
11271 .iter()
11272 .map(|h| {
11273 serde_json::json!({
11274 "from": h.from_agent,
11275 "to": h.to_agent,
11276 "reason": h.reason,
11277 })
11278 })
11279 .collect();
11280
11281 let _ = self.context_manager.set(
11283 "orchestration",
11284 serde_json::json!({
11285 "type": "handoff",
11286 "result": result.response.content,
11287 "final_agent": result.final_agent,
11288 "handoff_chain": chain_json,
11289 "duration_ms": duration_ms,
11290 }),
11291 );
11292
11293 self.commit_root_user_message(input).await?;
11294
11295 let post_result = self
11296 .post_loop_processing(input, result.response.content.clone())
11297 .await?;
11298 let final_content = self
11299 .apply_post_loop_result(input, post_result)
11300 .await?
11301 .content;
11302
11303 let mut response = AgentResponse::new(final_content);
11304 let metadata = serde_json::json!({
11305 "orchestration": {
11306 "type": "handoff",
11307 "result": result.response.content,
11308 "final_agent": result.final_agent,
11309 "handoff_chain": chain_json,
11310 "duration_ms": duration_ms,
11311 }
11312 });
11313 response.metadata = Some(
11314 serde_json::from_value::<std::collections::HashMap<String, serde_json::Value>>(
11315 metadata,
11316 )
11317 .unwrap_or_default(),
11318 );
11319
11320 self.finish_turn_if_root(&response).await?;
11321 Ok(response)
11322 }
11323
11324 async fn run_loop_internal(&self, input: &str) -> Result<AgentResponse> {
11326 self.begin_root_turn();
11327 self.pre_turn_session_lifecycle().await;
11329
11330 let input_data = self.process_input(input).await?;
11331 self.update_active_turn_context(&input_data.content, input_data.context.clone());
11332
11333 for (key, value) in &input_data.context {
11336 let _ = self.context_manager.set(key, value.clone());
11337 }
11338
11339 if input_data.metadata.rejected {
11340 let reason = input_data
11341 .metadata
11342 .rejection_reason
11343 .unwrap_or_else(|| "Input rejected".to_string());
11344 warn!(reason = %reason, "Input rejected");
11345 let response = AgentResponse::new(reason);
11346 self.finish_turn_if_root(&response).await?;
11347 return Ok(response);
11348 }
11349
11350 let processed_input = &input_data.content;
11351
11352 if let Some(response) = self.try_pre_response_transition(processed_input).await? {
11353 return Ok(response);
11354 }
11355
11356 if let Some(ref sm) = self.state_machine
11358 && let Some(def) = sm.current_definition()
11359 {
11360 if let Some(ref delegate_id) = def.delegate {
11361 return self
11362 .handle_delegated_state(processed_input, delegate_id, &def)
11363 .await;
11364 }
11365 if let Some(ref concurrent_config) = def.concurrent {
11366 return self
11367 .handle_concurrent_state(processed_input, concurrent_config)
11368 .await;
11369 }
11370 if let Some(ref group_chat_config) = def.group_chat {
11371 return self
11372 .handle_group_chat_state(processed_input, group_chat_config)
11373 .await;
11374 }
11375 if let Some(ref pipeline_config) = def.pipeline {
11376 return self
11377 .handle_pipeline_state(processed_input, pipeline_config)
11378 .await;
11379 }
11380 if let Some(ref handoff_config) = def.handoff {
11381 return self
11382 .handle_handoff_state(processed_input, handoff_config)
11383 .await;
11384 }
11385 }
11386
11387 if let Some(response) =
11392 Box::pin(self.try_speculative_branches(processed_input, &input_data.context)).await?
11393 {
11394 return Ok(response);
11395 }
11396
11397 match self.try_skill_route(processed_input).await? {
11398 SkillRouteResult::Response { skill_id, content } => {
11399 self.commit_root_user_message(processed_input).await?;
11400 return self
11401 .handle_skill_response(processed_input, &skill_id, content, &input_data.context)
11402 .await;
11403 }
11404 SkillRouteResult::NeedsClarification {
11405 response,
11406 ownership,
11407 } => {
11408 let admission = self
11409 .admit_optional_disambiguation_ownership(ownership)
11410 .await?;
11411 self.commit_root_user_message(processed_input).await?;
11412 if Self::skill_clarification_needs_memory_record(&response) {
11413 self.memory
11416 .add_message(ChatMessage::assistant(&response.content))
11417 .await?;
11418 }
11419 drop(admission);
11420 self.finish_turn_if_root(&response).await?;
11421 return Ok(response);
11422 }
11423 SkillRouteResult::NoMatch => {} }
11425
11426 let effective_reasoning = self.get_effective_reasoning_config();
11427 let reasoning_mode = self.determine_reasoning_mode(processed_input).await?;
11428 let auto_detected = matches!(effective_reasoning.mode, ReasoningMode::Auto);
11429
11430 info!(
11431 reasoning_mode = ?reasoning_mode,
11432 auto_detected = auto_detected,
11433 reflection_enabled = ?self.reflection_config.enabled,
11434 "Reasoning mode determined"
11435 );
11436
11437 if matches!(reasoning_mode, ReasoningMode::PlanAndExecute) {
11438 self.commit_root_user_message(processed_input).await?;
11439 return self
11440 .handle_plan_and_execute(processed_input, &input_data.context, auto_detected)
11441 .await;
11442 }
11443
11444 self.commit_root_user_message(processed_input).await?;
11445
11446 let mut iterations = 0u32;
11447 let mut all_tool_calls: Vec<ToolCall> = Vec::new();
11448 let mut thinking_content: Option<String> = None;
11449
11450 let llm = self.get_state_llm()?;
11451
11452 loop {
11453 let effective_max = if reasoning_mode != ReasoningMode::None {
11455 let rc = self.get_effective_reasoning_config();
11456 self.max_iterations.min(rc.max_iterations)
11457 } else {
11458 self.max_iterations
11459 };
11460
11461 if iterations >= effective_max {
11462 let err = AgentError::Other(format!("Max iterations ({}) exceeded", effective_max));
11463 self.hooks.on_error(&err).await;
11464 error!(iterations = iterations, "Max iterations exceeded");
11465 return Err(err);
11466 }
11467 iterations += 1;
11468 *self.iteration_count.write() = iterations;
11469
11470 debug!(iteration = iterations, max = effective_max, "LLM call");
11471
11472 let protocol = self.main_tool_protocol(llm.as_ref(), false).await?;
11473 let mut messages = self
11474 .build_messages_internal(true, None, protocol.choice.is_none())
11475 .await?;
11476 self.inject_reasoning_prompt(&mut messages, &reasoning_mode, iterations == 1);
11477
11478 self.hooks.on_llm_start(&messages).await;
11479 let llm_start = Instant::now();
11480 let response = self
11481 .complete_main_llm_with_recovery(Arc::clone(&llm), &messages, &protocol)
11482 .await?;
11483
11484 let llm_duration_ms = llm_start.elapsed().as_millis() as u64;
11485 self.hooks.on_llm_complete(&response, llm_duration_ms).await;
11486
11487 let content = response.content.trim();
11488
11489 if let Some(tool_calls) = self.parse_main_tool_calls(content, &protocol)? {
11490 match self
11491 .handle_tool_calls(
11492 processed_input,
11493 content,
11494 tool_calls,
11495 &mut all_tool_calls,
11496 None,
11497 )
11498 .await?
11499 {
11500 ToolCallOutcome::Continue | ToolCallOutcome::TransitionFired => continue,
11501 ToolCallOutcome::Rejected(resp) => {
11502 self.finish_turn_if_root(&resp).await?;
11503 return Ok(resp);
11504 }
11505 }
11506 }
11507
11508 let (extracted_thinking, answer) = self.extract_thinking(content);
11509 if extracted_thinking.is_some() {
11510 thinking_content = extracted_thinking;
11511 }
11512
11513 let output_data = self.process_output(&answer, &input_data.context).await?;
11514
11515 let mut final_content = if output_data.metadata.rejected {
11516 output_data
11517 .metadata
11518 .rejection_reason
11519 .unwrap_or_else(|| answer.to_string())
11520 } else {
11521 output_data.content
11522 };
11523
11524 let reflection_metadata;
11526 (final_content, reflection_metadata) = self
11527 .run_reflection(&*llm, processed_input, final_content)
11528 .await?;
11529
11530 final_content =
11531 self.format_response_with_thinking(thinking_content.as_deref(), &final_content);
11532
11533 let final_content = {
11537 let result = self
11538 .post_loop_processing(processed_input, final_content)
11539 .await?;
11540 self.apply_post_loop_result(processed_input, result)
11541 .await?
11542 .content
11543 };
11544
11545 let reflected = reflection_metadata.is_some();
11546 let reasoning_mode_debug = format!("{:?}", reasoning_mode);
11547
11548 let response = self.build_agent_response(AgentResponseParts {
11549 content: final_content,
11550 all_tool_calls,
11551 reasoning_mode,
11552 auto_detected,
11553 iterations,
11554 thinking: thinking_content,
11555 reflection_metadata,
11556 });
11557
11558 self.finish_turn_if_root(&response).await?;
11559
11560 let tool_call_count = response.tool_calls.as_ref().map(|tc| tc.len()).unwrap_or(0);
11561 info!(
11562 tool_calls = tool_call_count,
11563 response_len = response.content.len(),
11564 reasoning_mode = %reasoning_mode_debug,
11565 reflected = reflected,
11566 "Chat completed"
11567 );
11568 return Ok(response);
11569 }
11570 }
11571
11572 async fn generate_buffered_streaming_draft(
11573 &self,
11574 processed_input: &str,
11575 routing_resolved: Arc<AtomicBool>,
11576 ) -> Result<StreamingDraftResult> {
11577 let llm = self.get_state_llm()?;
11578 if llm.configured_tool_choice().is_some() {
11579 let draft = self
11580 .generate_main_response_draft(processed_input, &ReasoningMode::None)
11581 .await?;
11582 return Ok(StreamingDraftResult::new(draft, Vec::new()));
11583 }
11584 let protocol = self.main_tool_protocol(llm.as_ref(), true).await?;
11586 let messages = self.build_messages_for_draft(processed_input).await?;
11587 let source = self
11588 .open_main_stream_with_recovery(Arc::clone(&llm), &messages, &protocol)
11589 .await?;
11590 let mut buffer = crate::optimization::StreamBranchBuffer::new(self.streaming.buffer_size)?;
11591 let mut chunks = Vec::new();
11592 let mut accumulated = String::new();
11593 match source {
11594 MainStreamSource::StaticResponse(text) => {
11595 accumulated.push_str(&text);
11596 let stream_chunk = StreamChunk::content(text);
11597 if routing_resolved.load(Ordering::SeqCst) {
11598 chunks.push(stream_chunk);
11599 } else {
11600 buffer.push(stream_chunk)?;
11601 }
11602 }
11603 MainStreamSource::Stream(mut stream) => {
11604 while let Some(chunk_result) = stream.next().await {
11605 let chunk = chunk_result.map_err(|e| AgentError::LLM(e.to_string()))?;
11606 accumulated.push_str(&chunk.delta);
11607 let stream_chunk = StreamChunk::content(chunk.delta);
11608 if routing_resolved.load(Ordering::SeqCst) {
11609 chunks.push(stream_chunk);
11610 } else {
11611 buffer.push(stream_chunk)?;
11612 }
11613 }
11614 }
11615 }
11616 chunks.splice(0..0, buffer.drain());
11617 let content = accumulated.trim().to_string();
11618 let draft = if let Some(calls) = self.parse_tool_calls(&content)? {
11619 MainResponseDraft::ToolCalls {
11620 raw_content: content,
11621 calls,
11622 thinking: None,
11623 }
11624 } else {
11625 MainResponseDraft::Text {
11626 raw_content: content,
11627 thinking: None,
11628 }
11629 };
11630 Ok(StreamingDraftResult::new(draft, chunks))
11631 }
11632
11633 async fn try_buffered_streaming_branches(
11634 &self,
11635 processed_input: &str,
11636 input_context: &HashMap<String, Value>,
11637 ) -> Result<Option<(AgentResponse, Vec<StreamChunk>)>> {
11638 let optimization = &self.runtime_config.optimization;
11639 if !optimization.enabled {
11640 return Ok(None);
11641 }
11642 if !matches!(
11648 self.get_effective_reasoning_config().mode,
11649 ReasoningMode::None
11650 ) {
11651 return Ok(None);
11652 }
11653 let transition_enabled =
11654 optimization.speculative_state_transitions && self.has_parallel_transition_candidates();
11655 if !transition_enabled {
11656 return Ok(None);
11657 }
11658 let mut branch_scheduler =
11659 TurnBranchScheduler::new(optimization.max_parallel_runtime_tasks)?;
11660 if !branch_scheduler.reserve_task() {
11661 return Ok(None);
11662 }
11663 if !self
11664 .reserve_active_speculative_llm_call(RuntimeOptimizationKind::BufferedStreamingRouting)
11665 {
11666 branch_scheduler.release_task();
11667 return Ok(None);
11668 }
11669 if !branch_scheduler.reserve_task() {
11670 branch_scheduler.release_task();
11671 return Ok(None);
11672 }
11673 let mut main_branch = RuntimeBranch::new(
11674 RuntimeTaskPurpose::MainResponse,
11675 RuntimeOptimizationKind::BufferedStreamingRouting,
11676 RuntimeTaskPriority::Normal,
11677 RuntimeCommitBehavior::FinalResponse,
11678 );
11679 let mut transition_branch = RuntimeBranch::new(
11680 RuntimeTaskPurpose::StateTransition,
11681 RuntimeOptimizationKind::ParallelStateTransition,
11682 RuntimeTaskPriority::Critical,
11683 RuntimeCommitBehavior::TransitionDecision,
11684 );
11685 let main_id = main_branch.branch_id();
11686 let transition_id = transition_branch.branch_id();
11687 let routing_resolved = Arc::new(AtomicBool::new(false));
11688 let mut main_future =
11689 Box::pin(crate::optimization::observability::with_branch_observation(
11690 &main_id,
11691 RuntimeOptimizationKind::BufferedStreamingRouting,
11692 RuntimeCommitBehavior::FinalResponse,
11693 self.generate_buffered_streaming_draft(
11694 processed_input,
11695 Arc::clone(&routing_resolved),
11696 ),
11697 ));
11698 let mut transition_future =
11699 Box::pin(crate::optimization::observability::with_branch_observation(
11700 &transition_id,
11701 RuntimeOptimizationKind::ParallelStateTransition,
11702 RuntimeCommitBehavior::TransitionDecision,
11703 self.select_parallel_transition_candidate(processed_input),
11704 ));
11705 let mut main_pending = true;
11706 let mut transition_pending = true;
11707 let mut main_result: Option<Result<StreamingDraftResult>> = None;
11708 let mut transition_finalized = false;
11709 let mut transition_candidate: Option<TransitionCandidate> = None;
11710 loop {
11711 if let Some(candidate) = transition_candidate.take() {
11712 if self
11713 .approve_transition_target(&candidate.from_state, candidate.target())
11714 .await?
11715 {
11716 drop(main_future);
11718 drop(transition_future);
11719 self.finalize_branch_loss(
11720 &main_id,
11721 RuntimeOptimizationKind::BufferedStreamingRouting,
11722 RuntimeCommitBehavior::FinalResponse,
11723 main_pending,
11724 main_result.as_ref().map(|result| result.is_err()),
11725 );
11726 if !self
11727 .apply_pre_response_transition_candidate(
11728 &candidate,
11729 &HashMap::new(),
11730 processed_input,
11731 )
11732 .await?
11733 {
11734 self.finalize_optional_branch(
11735 &transition_id,
11736 RuntimeOptimizationKind::ParallelStateTransition,
11737 RuntimeCommitBehavior::TransitionDecision,
11738 "discarded",
11739 false,
11740 );
11741 return Ok(None);
11742 }
11743 self.finalize_optional_branch(
11744 &transition_id,
11745 RuntimeOptimizationKind::ParallelStateTransition,
11746 RuntimeCommitBehavior::TransitionDecision,
11747 "committed",
11748 true,
11749 );
11750 let response = self.redispatch_current_state(processed_input).await?;
11751 return Ok(Some((
11752 response.clone(),
11753 vec![StreamChunk::content(response.content)],
11754 )));
11755 }
11756 self.finalize_optional_branch(
11757 &transition_id,
11758 RuntimeOptimizationKind::ParallelStateTransition,
11759 RuntimeCommitBehavior::TransitionDecision,
11760 "discarded",
11761 false,
11762 );
11763 transition_finalized = true;
11764 }
11765 if transition_finalized && !routing_resolved.load(Ordering::SeqCst) {
11771 match self
11772 .resolve_buffered_skill_after_transition(processed_input, &routing_resolved)
11773 .await
11774 {
11775 Ok(Some(candidate)) => {
11776 drop(main_future);
11778 drop(transition_future);
11779 self.finalize_branch_loss(
11780 &main_id,
11781 RuntimeOptimizationKind::BufferedStreamingRouting,
11782 RuntimeCommitBehavior::FinalResponse,
11783 main_pending,
11784 main_result.as_ref().map(|result| result.is_err()),
11785 );
11786 return match self
11787 .commit_winning_skill_candidate(
11788 candidate,
11789 processed_input,
11790 input_context,
11791 )
11792 .await?
11793 {
11794 Some(response) => Ok(Some((
11795 response.clone(),
11796 vec![StreamChunk::content(response.content)],
11797 ))),
11798 None => Ok(None),
11799 };
11800 }
11801 Ok(None) => {}
11802 Err(error) => {
11803 drop(main_future);
11804 drop(transition_future);
11805 self.finalize_branch_loss(
11806 &main_id,
11807 RuntimeOptimizationKind::BufferedStreamingRouting,
11808 RuntimeCommitBehavior::FinalResponse,
11809 main_pending,
11810 main_result.as_ref().map(|result| result.is_err()),
11811 );
11812 return Err(error);
11813 }
11814 }
11815 }
11816 if transition_finalized
11817 && routing_resolved.load(Ordering::SeqCst)
11818 && let Some(result) = main_result.take()
11819 {
11820 let stream_draft = match result {
11821 Ok(stream_draft) => stream_draft,
11822 Err(error) => {
11823 self.finalize_optional_branch(
11824 &main_id,
11825 RuntimeOptimizationKind::BufferedStreamingRouting,
11826 RuntimeCommitBehavior::FinalResponse,
11827 "failed",
11828 false,
11829 );
11830 return Err(error);
11831 }
11832 };
11833 let raw_draft_content = stream_draft.draft.raw_content().to_string();
11834 let buffered_chunks = stream_draft.chunks;
11835 self.finalize_optional_branch(
11836 &main_id,
11837 RuntimeOptimizationKind::BufferedStreamingRouting,
11838 RuntimeCommitBehavior::FinalResponse,
11839 "committed",
11840 true,
11841 );
11842 let response = self
11843 .commit_main_response_draft(
11844 processed_input,
11845 input_context,
11846 stream_draft.draft,
11847 ReasoningMode::None,
11848 false,
11849 )
11850 .await?;
11851 let chunks = if response.content == raw_draft_content {
11852 buffered_chunks
11853 } else {
11854 vec![StreamChunk::content(response.content.clone())]
11855 };
11856 return Ok(Some((response, chunks)));
11857 }
11858 tokio::select! {
11859 result = &mut main_future, if main_pending => {
11860 main_pending = false;
11861 main_branch.transition_to(RuntimeBranchStatus::Completed)?;
11862 main_result = Some(result);
11863 }
11864 result = &mut transition_future, if transition_pending => {
11865 transition_pending = false;
11866 transition_branch.transition_to(RuntimeBranchStatus::Completed)?;
11867 match result {
11868 Ok(ParallelTransitionSelection::Candidate(candidate)) => {
11869 transition_candidate = Some(candidate)
11870 }
11871 Ok(ParallelTransitionSelection::NoMatch) => {
11872 self.finalize_optional_branch(
11873 &transition_id,
11874 RuntimeOptimizationKind::ParallelStateTransition,
11875 RuntimeCommitBehavior::TransitionDecision,
11876 "discarded",
11877 false,
11878 );
11879 transition_finalized = true;
11880 }
11881 Ok(ParallelTransitionSelection::ReservationExhausted) => {
11882 self.finalize_optional_branch(
11883 &transition_id,
11884 RuntimeOptimizationKind::ParallelStateTransition,
11885 RuntimeCommitBehavior::TransitionDecision,
11886 "cancelled",
11887 false,
11888 );
11889 routing_resolved.store(true, Ordering::SeqCst);
11890 self.finalize_branch_loss(
11891 &main_id,
11892 RuntimeOptimizationKind::BufferedStreamingRouting,
11893 RuntimeCommitBehavior::FinalResponse,
11894 main_pending,
11895 main_result.as_ref().map(|result| result.is_err()),
11896 );
11897 return Ok(None);
11898 }
11899 Err(_) => {
11900 self.finalize_optional_branch(
11901 &transition_id,
11902 RuntimeOptimizationKind::ParallelStateTransition,
11903 RuntimeCommitBehavior::TransitionDecision,
11904 "failed",
11905 false,
11906 );
11907 transition_finalized = true;
11908 }
11909 }
11910 }
11911 }
11912 }
11913 }
11914
11915 async fn resolve_buffered_skill_after_transition(
11921 &self,
11922 processed_input: &str,
11923 routing_resolved: &AtomicBool,
11924 ) -> Result<Option<SkillCandidate>> {
11925 let candidate = if self.skill_router.is_some() {
11926 self.select_skill_candidate(processed_input).await?
11927 } else {
11928 None
11929 };
11930 if candidate.is_none() {
11931 routing_resolved.store(true, Ordering::SeqCst);
11932 }
11933 Ok(candidate)
11934 }
11935
11936 fn run_loop_internal_stream<'a>(
11940 &'a self,
11941 input: &'a str,
11942 terminal: RuntimeStreamTerminalSlot,
11943 ) -> Pin<Box<dyn Stream<Item = StreamChunk> + Send + 'a>> {
11944 let include_state_events = self.streaming.include_state_events;
11945
11946 Box::pin(async_stream::stream! {
11947 self.begin_root_turn();
11948 self.pre_turn_session_lifecycle().await;
11950
11951 let input_data = match self.process_input(input).await {
11952 Ok(data) => data,
11953 Err(e) => {
11954 yield StreamChunk::error(e.to_string());
11955 return;
11956 }
11957 };
11958 self.update_active_turn_context(&input_data.content, input_data.context.clone());
11959
11960 for (key, value) in &input_data.context {
11962 let _ = self.context_manager.set(key, value.clone());
11963 }
11964
11965 if input_data.metadata.rejected {
11966 let reason = input_data
11967 .metadata
11968 .rejection_reason
11969 .unwrap_or_else(|| "Input rejected".to_string());
11970 warn!(reason = %reason, "Input rejected (stream)");
11971 let response = AgentResponse::new(&reason);
11974 if let Err(e) = self.finish_turn_if_root(&response).await {
11975 yield StreamChunk::error(e.to_string());
11976 return;
11977 }
11978 yield StreamChunk::content(&reason);
11979 record_runtime_stream_final(&terminal, response);
11980 yield StreamChunk::Done {};
11981 return;
11982 }
11983
11984 let processed_input = &input_data.content;
11985
11986 let streaming_policy = self.runtime_config.optimization.streaming_policy;
11987
11988 if self.runtime_config.optimization.enabled
11995 && !matches!(
11996 streaming_policy,
11997 crate::optimization::StreamingOptimizationPolicy::Disabled
11998 )
11999 {
12000 match self.try_pre_response_transition(processed_input).await {
12001 Ok(Some(response)) => {
12002 yield StreamChunk::content(&response.content);
12003 record_runtime_stream_final(&terminal, response);
12004 yield StreamChunk::Done {};
12005 return;
12006 }
12007 Ok(None) => {}
12008 Err(e) => {
12009 yield StreamChunk::error(e.to_string());
12010 return;
12011 }
12012 }
12013 }
12014
12015 if self.runtime_config.optimization.enabled
12016 && matches!(
12017 streaming_policy,
12018 crate::optimization::StreamingOptimizationPolicy::BufferUntilRoutingDone
12019 )
12020 {
12021 match Box::pin(self.try_buffered_streaming_branches(processed_input, &input_data.context)).await {
12026 Ok(Some((response, chunks))) => {
12027 for chunk in chunks {
12028 yield chunk;
12029 }
12030 record_runtime_stream_final(&terminal, response);
12031 yield StreamChunk::Done {};
12032 return;
12033 }
12034 Ok(None) => {}
12035 Err(e) => {
12036 yield StreamChunk::error(e.to_string());
12037 return;
12038 }
12039 }
12040 }
12041
12042 if let Some(ref sm) = self.state_machine
12044 && let Some(def) = sm.current_definition()
12045 {
12046 let orchestration_result = if let Some(ref delegate_id) = def.delegate {
12047 Some(self.handle_delegated_state(processed_input, delegate_id, &def).await)
12048 } else if let Some(ref concurrent_config) = def.concurrent {
12049 Some(self.handle_concurrent_state(processed_input, concurrent_config).await)
12050 } else if let Some(ref group_chat_config) = def.group_chat {
12051 Some(self.handle_group_chat_state(processed_input, group_chat_config).await)
12052 } else if let Some(ref pipeline_config) = def.pipeline {
12053 Some(self.handle_pipeline_state(processed_input, pipeline_config).await)
12054 } else if let Some(ref handoff_config) = def.handoff {
12055 Some(self.handle_handoff_state(processed_input, handoff_config).await)
12056 } else {
12057 None
12058 };
12059
12060 if let Some(result) = orchestration_result {
12061 match result {
12062 Ok(response) => {
12063 yield StreamChunk::content(&response.content);
12064 record_runtime_stream_final(&terminal, response);
12065 yield StreamChunk::Done {};
12066 }
12067 Err(e) => {
12068 yield StreamChunk::error(e.to_string());
12069 }
12070 }
12071 return;
12072 }
12073 }
12074
12075 match self.try_skill_route(processed_input).await {
12077 Ok(SkillRouteResult::Response { skill_id, content }) => {
12078 if let Err(e) = self.commit_root_user_message(processed_input).await {
12079 yield StreamChunk::error(e.to_string());
12080 return;
12081 }
12082 match self.handle_skill_response(processed_input, &skill_id, content, &input_data.context).await {
12083 Ok(resp) => {
12084 yield StreamChunk::content(&resp.content);
12085 record_runtime_stream_final(&terminal, resp);
12086 yield StreamChunk::Done {};
12087 return;
12088 }
12089 Err(e) => {
12090 yield StreamChunk::error(e.to_string());
12091 return;
12092 }
12093 }
12094 }
12095 Ok(SkillRouteResult::NeedsClarification {
12096 response,
12097 ownership,
12098 }) => {
12099 let admission = match self
12100 .admit_optional_disambiguation_ownership(ownership)
12101 .await
12102 {
12103 Ok(admission) => admission,
12104 Err(e) => {
12105 yield StreamChunk::error(e.to_string());
12106 return;
12107 }
12108 };
12109 if let Err(e) = self.commit_root_user_message(processed_input).await {
12110 yield StreamChunk::error(e.to_string());
12111 return;
12112 }
12113 if Self::skill_clarification_needs_memory_record(&response)
12115 && let Err(e) = self.memory.add_message(ChatMessage::assistant(&response.content)).await
12116 {
12117 yield StreamChunk::error(e.to_string());
12118 return;
12119 }
12120 drop(admission);
12121 if let Err(e) = self.finish_turn_if_root(&response).await {
12122 yield StreamChunk::error(e.to_string());
12123 return;
12124 }
12125 yield StreamChunk::content(&response.content);
12126 record_runtime_stream_final(&terminal, response);
12127 yield StreamChunk::Done {};
12128 return;
12129 }
12130 Ok(SkillRouteResult::NoMatch) => {} Err(e) => {
12132 yield StreamChunk::error(e.to_string());
12133 return;
12134 }
12135 }
12136
12137 let effective_reasoning = self.get_effective_reasoning_config();
12139 let reasoning_mode = match self.determine_reasoning_mode(processed_input).await {
12140 Ok(mode) => mode,
12141 Err(e) => {
12142 yield StreamChunk::error(e.to_string());
12143 return;
12144 }
12145 };
12146 let auto_detected = matches!(effective_reasoning.mode, ReasoningMode::Auto);
12147
12148 info!(
12149 reasoning_mode = ?reasoning_mode,
12150 auto_detected = auto_detected,
12151 "Reasoning mode determined (stream)"
12152 );
12153
12154 if matches!(reasoning_mode, ReasoningMode::PlanAndExecute) {
12156 if let Err(e) = self.commit_root_user_message(processed_input).await {
12157 yield StreamChunk::error(e.to_string());
12158 return;
12159 }
12160 match self.handle_plan_and_execute(processed_input, &input_data.context, auto_detected).await {
12161 Ok(resp) => {
12162 yield StreamChunk::content(&resp.content);
12163 record_runtime_stream_final(&terminal, resp);
12164 yield StreamChunk::Done {};
12165 return;
12166 }
12167 Err(e) => {
12168 yield StreamChunk::error(e.to_string());
12169 return;
12170 }
12171 }
12172 }
12173
12174 if let Err(e) = self.commit_root_user_message(processed_input).await {
12175 yield StreamChunk::error(e.to_string());
12176 return;
12177 }
12178
12179 let llm = match self.get_state_llm() {
12180 Ok(llm) => llm,
12181 Err(e) => {
12182 yield StreamChunk::error(e.to_string());
12183 return;
12184 }
12185 };
12186
12187 let mut iterations = 0u32;
12188 let mut all_tool_calls: Vec<ToolCall> = Vec::new();
12189 let mut thinking_content: Option<String> = None;
12190
12191 loop {
12192 let effective_max = if reasoning_mode != ReasoningMode::None {
12194 let rc = self.get_effective_reasoning_config();
12195 self.max_iterations.min(rc.max_iterations)
12196 } else {
12197 self.max_iterations
12198 };
12199
12200 if iterations >= effective_max {
12201 let err_msg = format!("Max iterations ({}) exceeded", effective_max);
12202 let err = AgentError::Other(err_msg.clone());
12203 self.hooks.on_error(&err).await;
12204 error!(iterations = iterations, "Max iterations exceeded (stream)");
12205 yield StreamChunk::error(err_msg);
12206 return;
12207 }
12208 iterations += 1;
12209 *self.iteration_count.write() = iterations;
12210
12211 debug!(iteration = iterations, max = effective_max, "LLM call (stream)");
12212
12213 let protocol = match self.main_tool_protocol(llm.as_ref(), false).await {
12214 Ok(protocol) => protocol,
12215 Err(e) => {
12216 yield StreamChunk::error(e.to_string());
12217 return;
12218 }
12219 };
12220 let mut messages = match self
12221 .build_messages_internal(true, None, protocol.choice.is_none())
12222 .await
12223 {
12224 Ok(m) => m,
12225 Err(e) => {
12226 yield StreamChunk::error(e.to_string());
12227 return;
12228 }
12229 };
12230 self.inject_reasoning_prompt(&mut messages, &reasoning_mode, iterations == 1);
12231
12232 self.hooks.on_llm_start(&messages).await;
12233 let llm_start = Instant::now();
12234
12235 let buffered_decision = self.main_stream_must_buffer(&reasoning_mode, &protocol);
12236 let content = if buffered_decision {
12237 let response = match self
12241 .complete_main_llm_with_recovery(
12242 Arc::clone(&llm),
12243 &messages,
12244 &protocol,
12245 )
12246 .await
12247 {
12248 Ok(r) => r,
12249 Err(e) => {
12250 yield StreamChunk::error(e.to_string());
12251 return;
12252 }
12253 };
12254 let llm_duration_ms = llm_start.elapsed().as_millis() as u64;
12255 self.hooks.on_llm_complete(&response, llm_duration_ms).await;
12256 response.content.trim().to_string()
12257 } else {
12258 let source = match self
12260 .open_main_stream_with_recovery(Arc::clone(&llm), &messages, &protocol)
12261 .await
12262 {
12263 Ok(source) => source,
12264 Err(e) => {
12265 yield StreamChunk::error(e.to_string());
12266 return;
12267 }
12268 };
12269 let mut accumulated = String::new();
12270 match source {
12271 MainStreamSource::StaticResponse(text) => {
12272 accumulated.push_str(&text);
12273 yield StreamChunk::content(text);
12274 }
12275 MainStreamSource::Stream(mut stream_inner) => {
12276 while let Some(chunk_result) = stream_inner.next().await {
12277 match chunk_result {
12278 Ok(chunk) => {
12279 accumulated.push_str(&chunk.delta);
12280 yield StreamChunk::content(chunk.delta);
12281 }
12282 Err(e) => {
12283 yield StreamChunk::error(e.to_string());
12285 return;
12286 }
12287 }
12288 }
12289 }
12290 }
12291 let llm_duration_ms = llm_start.elapsed().as_millis() as u64;
12292 let llm_response = ai_agents_core::LLMResponse::new(
12294 accumulated.trim(),
12295 ai_agents_core::FinishReason::Stop,
12296 );
12297 self.hooks.on_llm_complete(&llm_response, llm_duration_ms).await;
12298 accumulated.trim().to_string()
12299 };
12300
12301 let parsed_tool_calls = match self.parse_main_tool_calls(&content, &protocol) {
12303 Ok(calls) => calls,
12304 Err(error) => {
12305 yield StreamChunk::error(error.to_string());
12306 return;
12307 }
12308 };
12309 if let Some(tool_calls) = parsed_tool_calls {
12310 let mut events = Vec::new();
12313 let outcome = self
12314 .handle_tool_calls(
12315 processed_input,
12316 &content,
12317 tool_calls,
12318 &mut all_tool_calls,
12319 Some(&mut events),
12320 )
12321 .await;
12322 for chunk in events.drain(..) {
12323 yield chunk;
12324 }
12325 match outcome {
12326 Ok(ToolCallOutcome::Continue) | Ok(ToolCallOutcome::TransitionFired) => continue,
12327 Ok(ToolCallOutcome::Rejected(response)) => {
12328 if let Err(finalize_error) = self.finish_turn_if_root(&response).await {
12329 yield StreamChunk::error(finalize_error.to_string());
12330 return;
12331 }
12332 let legacy_error = response.content.clone();
12333 record_runtime_stream_final(&terminal, response);
12334 yield StreamChunk::error(legacy_error);
12335 yield StreamChunk::Done {};
12336 return;
12337 }
12338 Err(e) => {
12339 yield StreamChunk::error(e.to_string());
12340 return;
12341 }
12342 }
12343 }
12344
12345 let (extracted_thinking, answer) = self.extract_thinking(&content);
12347 if extracted_thinking.is_some() {
12348 thinking_content = extracted_thinking;
12349 }
12350
12351 let output_data = match self.process_output(&answer, &input_data.context).await {
12352 Ok(d) => d,
12353 Err(e) => {
12354 yield StreamChunk::error(e.to_string());
12355 return;
12356 }
12357 };
12358
12359 let final_content = if output_data.metadata.rejected {
12360 output_data
12361 .metadata
12362 .rejection_reason
12363 .unwrap_or_else(|| answer.to_string())
12364 } else {
12365 output_data.content
12366 };
12367
12368 let (final_content, reflection_metadata) = match self
12370 .run_reflection(&*llm, processed_input, final_content)
12371 .await
12372 {
12373 Ok(r) => r,
12374 Err(e) => {
12375 yield StreamChunk::error(e.to_string());
12376 return;
12377 }
12378 };
12379
12380 let final_content = self.format_response_with_thinking(
12381 thinking_content.as_deref(),
12382 &final_content,
12383 );
12384
12385 if buffered_decision {
12387 yield StreamChunk::content(&final_content);
12388 }
12389
12390 let post_result = match self
12394 .post_loop_processing(processed_input, final_content)
12395 .await
12396 {
12397 Ok(r) => r,
12398 Err(e) => {
12399 yield StreamChunk::error(e.to_string());
12400 return;
12401 }
12402 };
12403
12404 let applied = match self.apply_post_loop_result(processed_input, post_result).await {
12405 Ok(applied) => applied,
12406 Err(e) => {
12407 yield StreamChunk::error(e.to_string());
12408 return;
12409 }
12410 };
12411
12412 if applied.transitioned {
12413 if include_state_events
12414 && let Some(state) = self.current_state()
12415 {
12416 yield StreamChunk::state_transition(None, state);
12417 }
12418 if applied.regenerated {
12424 yield StreamChunk::content(&applied.content);
12425 }
12426 }
12427 let final_content = applied.content;
12428
12429 let final_response = self.build_agent_response(AgentResponseParts {
12431 content: final_content,
12432 all_tool_calls,
12433 reasoning_mode,
12434 auto_detected,
12435 iterations,
12436 thinking: thinking_content,
12437 reflection_metadata,
12438 });
12439 if let Err(e) = self.finish_turn_if_root(&final_response).await {
12440 yield StreamChunk::error(e.to_string());
12441 return;
12442 }
12443
12444 record_runtime_stream_final(&terminal, final_response);
12445 yield StreamChunk::Done {};
12446 return;
12447 }
12448 })
12449 }
12450
12451 fn run_loop_stream<'a>(
12454 &'a self,
12455 input: &'a str,
12456 terminal: RuntimeStreamTerminalSlot,
12457 ) -> Pin<Box<dyn Stream<Item = StreamChunk> + Send + 'a>> {
12458 Box::pin(async_stream::stream! {
12459 self.begin_root_turn();
12460 let _root_cleanup = RootTurnCleanup::new(self);
12461 self.hooks.on_message_received(input).await;
12462
12463 if let Err(e) = self.prepare_turn_context().await {
12465 yield StreamChunk::error(e.to_string());
12466 return;
12467 }
12468
12469 self.clear_disambiguation_context();
12471
12472 let input_to_run = match self.resolve_disambiguation(input).await {
12476 Err(e) => {
12477 yield StreamChunk::error(e.to_string());
12478 return;
12479 }
12480 Ok(DisambiguationDispatch::Terminal(response)) => {
12481 yield StreamChunk::content(&response.content);
12482 record_runtime_stream_final(&terminal, response);
12483 yield StreamChunk::Done {};
12484 return;
12485 }
12486 Ok(DisambiguationDispatch::RecheckSkill {
12487 skill_id,
12488 enriched_input,
12489 disambiguation_epoch,
12490 state_generation,
12491 }) => {
12492 match self
12493 .recheck_skill_disambiguation(
12494 &skill_id,
12495 &enriched_input,
12496 disambiguation_epoch,
12497 state_generation,
12498 )
12499 .await
12500 {
12501 Ok(resp) => {
12502 yield StreamChunk::content(&resp.content);
12503 record_runtime_stream_final(&terminal, resp);
12504 yield StreamChunk::Done {};
12505 return;
12506 }
12507 Err(e) => {
12508 yield StreamChunk::error(e.to_string());
12509 return;
12510 }
12511 }
12512 }
12513 Ok(DisambiguationDispatch::Proceed(input)) => input,
12514 };
12515
12516 let mut inner = self.run_loop_internal_stream(&input_to_run, Arc::clone(&terminal));
12517 while let Some(chunk) = inner.next().await {
12518 yield chunk;
12519 }
12520 })
12521 }
12522
12523 pub fn info(&self) -> AgentInfo {
12524 self.info.clone()
12525 }
12526
12527 pub fn skills(&self) -> &[SkillDefinition] {
12528 &self.skills
12529 }
12530
12531 async fn reset_runtime_state(&self) -> Result<()> {
12533 let _admission = self.disambiguation_admission.write().await;
12534 if self.state_transition_reserved.load(Ordering::SeqCst) {
12535 return Err(AgentError::Other(
12536 "Cannot reset while a state transition is in progress".to_string(),
12537 ));
12538 }
12539 self.disambiguation_epoch.fetch_add(1, Ordering::SeqCst);
12540 *self.pending_skill_id.write() = None;
12541 if let Some(disambiguator) = self.disambiguation_manager.as_ref() {
12542 disambiguator.clear_pending().await;
12543 }
12544 self.memory.clear().await?;
12545 self.active_native_exchanges.write().clear();
12546 *self.iteration_count.write() = 0;
12547 self.tool_call_history.write().clear();
12548 if let Some(ref sm) = self.state_machine {
12549 sm.reset();
12550 }
12551 Ok(())
12552 }
12553
12554 pub async fn reset(&self) -> Result<()> {
12556 self.reset_runtime_state().await
12557 }
12558
12559 pub fn max_context_tokens(&self) -> u32 {
12560 self.max_context_tokens
12561 }
12562
12563 pub fn llm_registry(&self) -> &Arc<LLMRegistry> {
12564 &self.llm_registry
12565 }
12566
12567 pub fn state_machine(&self) -> Option<&Arc<StateMachine>> {
12568 self.state_machine.as_ref()
12569 }
12570
12571 pub fn context_manager(&self) -> &Arc<ContextManager> {
12572 &self.context_manager
12573 }
12574
12575 pub fn tool_call_history(&self) -> Vec<ToolCallRecord> {
12576 self.tool_call_history.read().clone()
12577 }
12578
12579 pub fn memory_token_budget(&self) -> Option<&MemoryTokenBudget> {
12580 self.memory_token_budget.as_ref()
12581 }
12582
12583 pub fn parallel_tools_config(&self) -> &ParallelToolsConfig {
12584 &self.parallel_tools
12585 }
12586
12587 pub fn streaming_config(&self) -> &StreamingConfig {
12588 &self.streaming
12589 }
12590
12591 pub fn hooks(&self) -> &Arc<dyn AgentHooks> {
12592 &self.hooks
12593 }
12594
12595 pub fn hitl_engine(&self) -> Option<&HITLEngine> {
12596 self.hitl_engine.as_ref()
12597 }
12598
12599 pub fn approval_handler(&self) -> &Arc<dyn ApprovalHandler> {
12600 &self.approval_handler
12601 }
12602
12603 fn build_hitl_language_context(&self) -> HashMap<String, Value> {
12605 let mut ctx = HashMap::new();
12606 for key in &["user.language", "input.detected.language", "language"] {
12607 if let Some(val) = self.context_manager.get(key) {
12608 ctx.insert(key.to_string(), val);
12609 }
12610 }
12611 ctx
12612 }
12613
12614 async fn request_hitl_approval(&self, check_result: HITLCheckResult) -> Result<ApprovalResult> {
12616 let Some(request) = check_result.into_request() else {
12617 return Ok(ApprovalResult::Approved);
12618 };
12619
12620 self.hooks.on_approval_requested(&request).await;
12621
12622 let timeout = request.timeout;
12623
12624 let raw_result = if let Some(duration) = timeout {
12625 match tokio::time::timeout(
12626 duration,
12627 self.approval_handler.request_approval(request.clone()),
12628 )
12629 .await
12630 {
12631 Ok(result) => result,
12632 Err(_) => ApprovalResult::timeout(),
12633 }
12634 } else {
12635 self.approval_handler
12636 .request_approval(request.clone())
12637 .await
12638 };
12639
12640 self.hooks
12641 .on_approval_result(&request.id, &raw_result)
12642 .await;
12643
12644 let (outcome, effective_result): (ApprovalResolvedOutcome, Result<ApprovalResult>) =
12645 match &raw_result {
12646 ApprovalResult::Approved => (
12647 ApprovalResolvedOutcome::Approved,
12648 Ok(ApprovalResult::Approved),
12649 ),
12650 ApprovalResult::Rejected { reason } => (
12651 ApprovalResolvedOutcome::Rejected {
12652 reason: reason.clone(),
12653 },
12654 Ok(ApprovalResult::Rejected {
12655 reason: reason.clone(),
12656 }),
12657 ),
12658 ApprovalResult::Modified { changes } => (
12659 ApprovalResolvedOutcome::Modified {
12660 changes: changes.clone(),
12661 },
12662 Ok(ApprovalResult::Modified {
12663 changes: changes.clone(),
12664 }),
12665 ),
12666 ApprovalResult::Timeout => {
12667 if let Some(ref engine) = self.hitl_engine {
12668 match engine.config().on_timeout {
12669 TimeoutAction::Approve => (
12670 ApprovalResolvedOutcome::Approved,
12671 Ok(ApprovalResult::Approved),
12672 ),
12673 TimeoutAction::Reject => {
12674 let reason = Some("Timeout".to_string());
12675 (
12676 ApprovalResolvedOutcome::Rejected {
12677 reason: reason.clone(),
12678 },
12679 Ok(ApprovalResult::Rejected { reason }),
12680 )
12681 }
12682 TimeoutAction::Error => {
12683 let message = "HITL approval timeout".to_string();
12684 (
12685 ApprovalResolvedOutcome::Error {
12686 message: message.clone(),
12687 },
12688 Err(AgentError::Other(message)),
12689 )
12690 }
12691 }
12692 } else {
12693 let reason = Some("Timeout (no engine)".to_string());
12694 (
12695 ApprovalResolvedOutcome::Rejected {
12696 reason: reason.clone(),
12697 },
12698 Ok(ApprovalResult::Rejected { reason }),
12699 )
12700 }
12701 }
12702 };
12703
12704 self.hooks
12705 .on_approval_resolved(&request, &raw_result, &outcome)
12706 .await;
12707
12708 effective_result
12709 }
12710
12711 pub async fn check_state_hitl(&self, from: Option<&str>, to: &str) -> Result<bool> {
12712 if let Some(ref hitl_engine) = self.hitl_engine {
12713 let hitl_lang_ctx = self.build_hitl_language_context();
12714 let check_result = self
12715 .observe_purpose(
12716 ObservationPurpose::HitlLocalization,
12717 hitl_engine.check_state_transition_with_localization(
12718 from,
12719 to,
12720 &hitl_lang_ctx,
12721 self.approval_handler.as_ref(),
12722 Some(&self.llm_registry),
12723 ),
12724 )
12725 .await?;
12726 if check_result.is_required() {
12727 let result = self.request_hitl_approval(check_result).await?;
12728 return Ok(matches!(
12729 result,
12730 ApprovalResult::Approved | ApprovalResult::Modified { .. }
12731 ));
12732 }
12733 }
12734 Ok(true)
12735 }
12736
12737 async fn execute_tools_parallel(
12739 &self,
12740 tool_calls: &[ToolCall],
12741 ) -> Vec<(String, Result<String>)> {
12742 let can_run_parallel = tool_calls.iter().all(|tc| {
12743 self.tools
12744 .resolve(&tc.name)
12745 .map(|resolved| resolved.tool.classify_call(&tc.arguments).concurrency_safe)
12746 .unwrap_or(false)
12747 });
12748
12749 if !self.parallel_tools.enabled || tool_calls.len() <= 1 || !can_run_parallel {
12750 let mut results = Vec::new();
12751 for tc in tool_calls {
12752 let result = self
12753 .observe_purpose(
12754 current_observation_context()
12755 .map(|context| context.purpose)
12756 .unwrap_or_default(),
12757 self.execute_tool_smart(tc),
12758 )
12759 .await;
12760 results.push((tc.id.clone(), result));
12761 }
12762 return results;
12763 }
12764
12765 let chunks: Vec<_> = tool_calls
12766 .chunks(self.parallel_tools.max_parallel)
12767 .collect();
12768
12769 let mut all_results = Vec::new();
12770
12771 for chunk in chunks {
12772 let futures: Vec<_> = chunk
12773 .iter()
12774 .map(|tc| {
12775 let tc = tc.clone();
12776 async move {
12777 let result = self.execute_tool_smart(&tc).await;
12778 (tc.id.clone(), result)
12779 }
12780 })
12781 .collect();
12782
12783 let results = futures::future::join_all(futures).await;
12784 all_results.extend(results);
12785 }
12786
12787 all_results
12788 }
12789
12790 pub async fn chat_stream<'a>(
12794 &'a self,
12795 input: &'a str,
12796 ) -> Result<Pin<Box<dyn Stream<Item = StreamChunk> + Send + 'a>>> {
12797 let RootTurnAdmission {
12798 guard: root_turn_guard,
12799 identity_stack,
12800 } = self.acquire_root_turn().await?;
12801 scope_runtime_gate_identity_stack(&identity_stack, self.init_storage()).await?;
12805 info!(input_len = input.len(), "Starting streaming chat");
12806 let terminal = new_runtime_stream_terminal_slot();
12807 let inner = self.run_loop_stream(input, terminal);
12808 let observation_context = self.build_observation_context(None);
12809 let stream: Pin<Box<dyn Stream<Item = StreamChunk> + Send + 'a>> =
12810 Box::pin(async_stream::stream! {
12811 let mut root_turn_guard = Some(root_turn_guard);
12812 let mut inner = inner;
12813 loop {
12814 let next = scope_runtime_gate_identity_stack(&identity_stack, async {
12815 if let Some(context) = observation_context.as_ref() {
12816 with_observation_context(context.clone(), inner.next()).await
12817 } else {
12818 inner.next().await
12819 }
12820 })
12821 .await;
12822 match next {
12823 Some(StreamChunk::Done {}) => {
12824 while scope_runtime_gate_identity_stack(&identity_stack, async {
12825 if let Some(context) = observation_context.as_ref() {
12826 with_observation_context(context.clone(), inner.next())
12827 .await
12828 .is_some()
12829 } else {
12830 inner.next().await.is_some()
12831 }
12832 })
12833 .await
12834 {}
12835 if observation_context.is_some() {
12836 scope_runtime_gate_identity_stack(
12837 &identity_stack,
12838 self.export_observability_if_configured(),
12839 )
12840 .await;
12841 }
12842 drop(root_turn_guard.take());
12843 yield StreamChunk::Done {};
12844 return;
12845 }
12846 Some(chunk) => yield chunk,
12847 None => {
12848 if observation_context.is_some() {
12849 scope_runtime_gate_identity_stack(
12850 &identity_stack,
12851 self.export_observability_if_configured(),
12852 )
12853 .await;
12854 }
12855 drop(root_turn_guard.take());
12856 return;
12857 }
12858 }
12859 }
12860 });
12861 Ok(stream)
12862 }
12863
12864 pub async fn chat_stream_events<'a>(
12868 &'a self,
12869 input: &'a str,
12870 ) -> Result<Pin<Box<dyn Stream<Item = AgentStreamEvent> + Send + 'a>>> {
12871 let RootTurnAdmission {
12872 guard,
12873 identity_stack,
12874 } = self.acquire_root_turn().await?;
12875 scope_runtime_gate_identity_stack(&identity_stack, self.init_storage()).await?;
12879 info!(input_len = input.len(), "Starting streaming chat events");
12880 let terminal = new_runtime_stream_terminal_slot();
12881 let inner = self.run_loop_stream(input, Arc::clone(&terminal));
12882 let observation_context = self.build_observation_context(None);
12883 Ok(self.drive_event_stream(
12884 inner,
12885 terminal,
12886 guard,
12887 identity_stack,
12888 observation_context,
12889 None,
12890 ))
12891 }
12892
12893 pub async fn chat_stream_events_with_actor_context<'a>(
12899 &'a self,
12900 input: &'a str,
12901 actor_context: crate::TurnActorContext,
12902 ) -> Result<Pin<Box<dyn Stream<Item = AgentStreamEvent> + Send + 'a>>> {
12903 let RootTurnAdmission {
12904 guard,
12905 identity_stack,
12906 } = self.acquire_root_turn().await?;
12907 scope_runtime_gate_identity_stack(&identity_stack, self.init_storage()).await?;
12908 info!(
12909 input_len = input.len(),
12910 "Starting streaming chat events with actor context"
12911 );
12912 let actor_id = actor_context.effective_actor_id().map(str::to_string);
12913 let terminal = new_runtime_stream_terminal_slot();
12914 let inner = self.run_loop_stream(input, Arc::clone(&terminal));
12915 let observation_context = self.build_observation_context(actor_id);
12916 Ok(self.drive_event_stream(
12917 inner,
12918 terminal,
12919 guard,
12920 identity_stack,
12921 observation_context,
12922 Some(actor_context),
12923 ))
12924 }
12925
12926 fn drive_event_stream<'a>(
12935 &'a self,
12936 mut inner: Pin<Box<dyn Stream<Item = StreamChunk> + Send + 'a>>,
12937 terminal: RuntimeStreamTerminalSlot,
12938 root_turn_guard: tokio::sync::OwnedMutexGuard<()>,
12939 identity_stack: RootTurnGateIdentityStack,
12940 observation_context: Option<SpanContext>,
12941 actor_context: Option<crate::TurnActorContext>,
12942 ) -> Pin<Box<dyn Stream<Item = AgentStreamEvent> + Send + 'a>> {
12943 Box::pin(async_stream::stream! {
12944 let mut root_turn_guard = Some(root_turn_guard);
12945 loop {
12946 let next = poll_scoped_chunk(
12947 &mut inner,
12948 &identity_stack,
12949 observation_context.as_ref(),
12950 actor_context.as_ref(),
12951 )
12952 .await;
12953 match next {
12954 Some(StreamChunk::Done {}) => {
12955 let terminal_event = { terminal.write().take() };
12956 if let Some(response) = terminal_event {
12957 while poll_scoped_chunk(
12958 &mut inner,
12959 &identity_stack,
12960 observation_context.as_ref(),
12961 actor_context.as_ref(),
12962 )
12963 .await
12964 .is_some()
12965 {}
12966 if observation_context.is_some() {
12967 scope_runtime_gate_identity_stack(
12968 &identity_stack,
12969 self.export_observability_if_configured(),
12970 )
12971 .await;
12972 }
12973 drop(root_turn_guard.take());
12974 yield AgentStreamEvent::Final(response);
12975 return;
12976 }
12977 }
12978 Some(StreamChunk::Error { message }) => {
12979 let finalized = { terminal.read().is_some() };
12980 if finalized {
12981 continue;
12982 }
12983 while poll_scoped_chunk(
12984 &mut inner,
12985 &identity_stack,
12986 observation_context.as_ref(),
12987 actor_context.as_ref(),
12988 )
12989 .await
12990 .is_some()
12991 {}
12992 if observation_context.is_some() {
12993 scope_runtime_gate_identity_stack(
12994 &identity_stack,
12995 self.export_observability_if_configured(),
12996 )
12997 .await;
12998 }
12999 drop(root_turn_guard.take());
13000 yield AgentStreamEvent::Chunk(StreamChunk::Error { message });
13001 return;
13002 }
13003 Some(chunk) => yield AgentStreamEvent::Chunk(chunk),
13004 None => {
13005 if observation_context.is_some() {
13006 scope_runtime_gate_identity_stack(
13007 &identity_stack,
13008 self.export_observability_if_configured(),
13009 )
13010 .await;
13011 }
13012 drop(root_turn_guard.take());
13013 return;
13014 }
13015 }
13016 }
13017 })
13018 }
13019}
13020
13021async fn poll_scoped_chunk<'a>(
13027 inner: &mut Pin<Box<dyn Stream<Item = StreamChunk> + Send + 'a>>,
13028 identity_stack: &RootTurnGateIdentityStack,
13029 observation_context: Option<&SpanContext>,
13030 actor_context: Option<&crate::TurnActorContext>,
13031) -> Option<StreamChunk> {
13032 scope_runtime_gate_identity_stack(identity_stack, async {
13033 let next = inner.next();
13034 match (observation_context, actor_context) {
13035 (Some(observation), Some(actor)) => {
13036 with_observation_context(
13037 observation.clone(),
13038 scope_actor_context(actor.clone(), next),
13039 )
13040 .await
13041 }
13042 (Some(observation), None) => with_observation_context(observation.clone(), next).await,
13043 (None, Some(actor)) => scope_actor_context(actor.clone(), next).await,
13044 (None, None) => next.await,
13045 }
13046 })
13047 .await
13048}
13049
13050#[async_trait]
13051impl ToolInvoker for RuntimeAgent {
13052 async fn invoke_tool(&self, request: ToolExecutionRequest) -> Result<ToolExecutionRecord> {
13053 self.execute_tool_record(request).await
13054 }
13055}
13056
13057#[async_trait]
13058impl Agent for RuntimeAgent {
13059 async fn chat(&self, input: &str) -> Result<AgentResponse> {
13061 let RootTurnAdmission {
13062 guard,
13063 identity_stack,
13064 } = self.acquire_root_turn().await?;
13065 let result = scope_runtime_gate_identity_stack(&identity_stack, async {
13066 let result = if let Some(context) = self.build_observation_context(None) {
13067 with_observation_context(context, self.run_loop(input)).await
13068 } else {
13069 self.run_loop(input).await
13070 };
13071 self.export_observability_if_configured().await;
13072 result
13073 })
13074 .await;
13075 drop(guard);
13076 result
13077 }
13078
13079 fn info(&self) -> AgentInfo {
13080 self.info.clone()
13081 }
13082
13083 async fn reset(&self) -> Result<()> {
13085 self.reset_runtime_state().await
13086 }
13087}
13088
13089fn background_maintenance_tags(
13099 label: &str,
13100 stage: &str,
13101 reason: Option<&str>,
13102 policy: Option<&crate::optimization::config::MaintenanceTaskPolicy>,
13103) -> HashMap<String, String> {
13104 let mut tags = HashMap::new();
13105 tags.insert("runtime.background".to_string(), "true".to_string());
13106 tags.insert("runtime.maintenance".to_string(), label.to_string());
13107 tags.insert("runtime.maintenance_stage".to_string(), stage.to_string());
13108 if let Some(policy) = policy {
13109 tags.insert(
13110 "runtime.await_before_next_turn".to_string(),
13111 await_before_next_turn_label(policy.await_before_next_turn).to_string(),
13112 );
13113 tags.insert(
13114 "runtime.maintenance_mode".to_string(),
13115 maintenance_mode_label(policy.mode).to_string(),
13116 );
13117 }
13118 if let Some(reason) = reason {
13119 tags.insert("runtime.reason".to_string(), reason.to_string());
13120 }
13121 tags
13122}
13123
13124fn await_before_next_turn_label(policy: AwaitBeforeNextTurn) -> &'static str {
13125 match policy {
13126 AwaitBeforeNextTurn::Never => "never",
13127 AwaitBeforeNextTurn::SameActor => "same_actor",
13128 AwaitBeforeNextTurn::Always => "always",
13129 }
13130}
13131
13132fn maintenance_mode_label(mode: MaintenanceMode) -> &'static str {
13133 match mode {
13134 MaintenanceMode::InlineSerial => "inline_serial",
13135 MaintenanceMode::InlineParallel => "inline_parallel",
13136 MaintenanceMode::Background => "background",
13137 }
13138}
13139
13140fn record_background_maintenance_event(
13142 manager: Option<&Arc<ObservabilityManager>>,
13143 label: &str,
13144 status: EventStatus,
13145 duration_ms: u64,
13146 stage: &str,
13147 reason: Option<String>,
13148 policy: Option<&crate::optimization::config::MaintenanceTaskPolicy>,
13149) {
13150 if let Some(manager) = manager {
13151 manager.record_lifecycle_event(
13152 EventType::MemoryOperation {
13153 operation: format!("{}_background_{}", label, stage),
13154 },
13155 ObservationPurpose::Other(format!("{}_maintenance", label)),
13156 status,
13157 duration_ms,
13158 background_maintenance_tags(label, stage, reason.as_deref(), policy),
13159 None,
13160 );
13161 }
13162}
13163
13164fn effective_maintenance_mode(mode: MaintenanceMode, force_parallel: bool) -> MaintenanceMode {
13165 if force_parallel && matches!(mode, MaintenanceMode::InlineSerial) {
13166 MaintenanceMode::InlineParallel
13167 } else {
13168 mode
13169 }
13170}
13171
13172fn observation_purpose_for_process(hint: ProcessPurposeHint) -> ObservationPurpose {
13173 match hint {
13174 ProcessPurposeHint::Detect => ObservationPurpose::ProcessDetect,
13175 ProcessPurposeHint::Extract => ObservationPurpose::ProcessExtract,
13176 ProcessPurposeHint::Validate => ObservationPurpose::ProcessValidate,
13177 ProcessPurposeHint::Transform | ProcessPurposeHint::Other => {
13178 ObservationPurpose::ProcessTransform
13179 }
13180 }
13181}
13182
13183fn new_tool_resource_locks() -> ToolResourceLocks {
13184 Arc::new(RwLock::new(HashMap::new()))
13185}
13186
13187fn tool_resource_lock_keys(
13192 _canonical_id: &str,
13193 args: &Value,
13194 bindings: &ai_agents_core::ToolPolicyBindings,
13195 classification: &ai_agents_core::ToolCallClassification,
13196) -> Vec<String> {
13197 if classification.concurrency_safe {
13198 return Vec::new();
13199 }
13200
13201 let mut keys = Vec::new();
13202 let mut has_path_resource = false;
13203 for binding in &bindings.path_fields {
13204 let value = value_at_argument_path(args, &binding.field)
13205 .cloned()
13206 .or_else(|| {
13207 binding
13208 .default_path
13209 .as_ref()
13210 .map(|path| Value::String(path.clone()))
13211 });
13212 if let Some(value) = value {
13213 collect_resource_strings(&value, |_| {
13214 has_path_resource = true;
13215 });
13216 }
13217 }
13218 for binding in &bindings.domain_fields {
13219 if let Some(value) = value_at_argument_path(args, &binding.field) {
13220 collect_resource_strings(value, |domain| {
13221 let normalized = if binding.is_url {
13222 normalized_url_resource_key(domain)
13223 } else {
13224 domain.trim().trim_end_matches('.').to_ascii_lowercase()
13225 };
13226 keys.push(format!("domain:{}", normalized));
13227 });
13228 }
13229 }
13230 for binding in &bindings.command_fields {
13231 if !matches!(binding.kind, ai_agents_core::CommandBindingKind::Cwd) {
13232 continue;
13233 }
13234 if let Some(value) = value_at_argument_path(args, &binding.field) {
13235 collect_resource_strings(value, |_| {
13236 has_path_resource = true;
13237 });
13238 }
13239 }
13240 if has_path_resource {
13241 keys.push("path-mutation:global".to_string());
13242 }
13243 if keys.is_empty() {
13244 keys.push("side-effect:unbound".to_string());
13245 }
13246 keys.sort();
13247 keys.dedup();
13248 keys
13249}
13250
13251fn value_at_argument_path<'a>(value: &'a Value, field: &str) -> Option<&'a Value> {
13252 let mut current = value;
13253 for segment in field.split('.') {
13254 if segment.is_empty() {
13255 return None;
13256 }
13257 current = current.get(segment)?;
13258 }
13259 Some(current)
13260}
13261
13262fn collect_resource_strings(value: &Value, mut collect: impl FnMut(&str)) {
13263 match value {
13264 Value::String(value) => collect(value),
13265 Value::Array(values) => {
13266 for value in values {
13267 if let Some(value) = value.as_str() {
13268 collect(value);
13269 }
13270 }
13271 }
13272 _ => {}
13273 }
13274}
13275
13276fn normalized_url_resource_key(value: &str) -> String {
13277 let value = value.trim();
13278 let Some((scheme, remainder)) = value.split_once("://") else {
13279 return value.to_ascii_lowercase();
13280 };
13281 let authority_end = remainder.find(['/', '?', '#']).unwrap_or(remainder.len());
13282 let (authority, suffix) = remainder.split_at(authority_end);
13283 format!(
13284 "{}://{}{}",
13285 scheme.to_ascii_lowercase(),
13286 authority.to_ascii_lowercase(),
13287 suffix
13288 )
13289}
13290
13291fn render_concurrent_template(
13292 template: &str,
13293 user_input: &str,
13294 context_values: &std::collections::HashMap<String, serde_json::Value>,
13295) -> Result<String> {
13296 let mut env = minijinja::Environment::new();
13297 env.add_template("concurrent", template)
13298 .map_err(|e| AgentError::Other(format!("Concurrent template parse error: {}", e)))?;
13299
13300 let mut ctx = std::collections::BTreeMap::new();
13301 ctx.insert("user_input".to_string(), minijinja::Value::from(user_input));
13302
13303 let context_obj = minijinja::Value::from_serialize(context_values);
13305 ctx.insert("context".to_string(), context_obj);
13306
13307 let tmpl = env
13308 .get_template("concurrent")
13309 .map_err(|e| AgentError::Other(format!("Concurrent template error: {}", e)))?;
13310
13311 tmpl.render(minijinja::Value::from_serialize(&ctx))
13312 .map_err(|e| AgentError::Other(format!("Concurrent template render error: {}", e)))
13313}
13314
13315#[cfg(test)]
13316mod tests {
13317 use super::*;
13318 use crate::AgentBuilder;
13319 use ai_agents_core::{LLMChunk, LLMConfig, LLMError, LLMFeature, Tool};
13320 use ai_agents_llm::mock::MockLLMProvider;
13321 use ai_agents_skills::{SkillDefinition, SkillStep};
13322 use ai_agents_tools::{
13323 CalculatorTool, CopyPathTool, DeletePathTool, FileWriteTool, MovePathTool, ToolAliases,
13324 ToolDescriptor, ToolProvider, ToolProviderError, ToolProviderType, WebFetchResolver,
13325 WebFetchTool, WebFetchTransport, WebFetchTransportRequest, WebFetchTransportResponse,
13326 };
13327
13328 fn mock_with_response(response: &str) -> MockLLMProvider {
13329 let mut mock = MockLLMProvider::new("test");
13330 mock.set_response(response);
13331 mock
13332 }
13333
13334 fn mock_with_responses(responses: Vec<&str>) -> MockLLMProvider {
13335 let mut mock = MockLLMProvider::new("test");
13336 mock.set_responses(responses.into_iter().map(String::from).collect(), true);
13337 mock
13338 }
13339
13340 async fn collect_stream_events(
13342 agent: &RuntimeAgent,
13343 input: &str,
13344 ) -> (String, Vec<StreamChunk>, Option<AgentResponse>) {
13345 use futures::StreamExt;
13346 let mut events = agent.chat_stream_events(input).await.expect("stream opens");
13347 let mut content = String::new();
13348 let mut chunks = Vec::new();
13349 let mut final_response = None;
13350 while let Some(event) = events.next().await {
13351 match event {
13352 AgentStreamEvent::Chunk(chunk) => {
13353 if let StreamChunk::Content { text } = &chunk {
13354 content.push_str(text);
13355 }
13356 chunks.push(chunk);
13357 }
13358 AgentStreamEvent::Final(response) => final_response = Some(response),
13359 }
13360 }
13361 (content, chunks, final_response)
13362 }
13363
13364 fn metadata_keys(response: &AgentResponse) -> std::collections::BTreeSet<String> {
13365 response
13366 .metadata
13367 .as_ref()
13368 .map(|m| m.keys().cloned().collect())
13369 .unwrap_or_default()
13370 }
13371
13372 async fn assert_blocking_streaming_parity<F>(
13375 build: F,
13376 input: &str,
13377 ) -> (AgentResponse, AgentResponse, Vec<StreamChunk>)
13378 where
13379 F: Fn() -> RuntimeAgent,
13380 {
13381 let blocking_agent = build();
13382 let streaming_agent = build();
13383
13384 let blocking = blocking_agent
13385 .chat(input)
13386 .await
13387 .expect("blocking chat succeeds");
13388 let (_, chunks, final_response) = collect_stream_events(&streaming_agent, input).await;
13389 let streamed = final_response.expect("streaming must emit Final when blocking succeeds");
13390
13391 assert_eq!(
13392 blocking.content, streamed.content,
13393 "committed content differs"
13394 );
13395 assert_eq!(
13396 metadata_keys(&blocking),
13397 metadata_keys(&streamed),
13398 "metadata key sets differ"
13399 );
13400 assert_eq!(
13401 blocking.tool_calls.as_ref().map(Vec::len),
13402 streamed.tool_calls.as_ref().map(Vec::len),
13403 "tool call counts differ"
13404 );
13405 assert_eq!(
13406 blocking_agent.current_state(),
13407 streaming_agent.current_state(),
13408 "final states differ"
13409 );
13410 (blocking, streamed, chunks)
13411 }
13412
13413 fn signed_calculator_response(
13414 exchange_id: &str,
13415 call_id: &str,
13416 expression: &str,
13417 ) -> LLMResponse {
13418 let call = ToolCall {
13419 id: call_id.to_string(),
13420 name: "calculator".to_string(),
13421 arguments: serde_json::json!({"expression": expression}),
13422 };
13423 let state = ai_agents_core::NativeProviderState::new(
13424 exchange_id,
13425 "fixture",
13426 "native-tools",
13427 ai_agents_core::NativeProviderTarget::new("https://fixture.invalid/", "fixture-model")
13428 .unwrap(),
13429 serde_json::json!({
13430 "role": "model",
13431 "parts": [{
13432 "functionCall": {"name": "calculator", "args": {"expression": expression}},
13433 "thoughtSignature": format!("signature-{exchange_id}")
13434 }]
13435 }),
13436 vec![ai_agents_core::NativeCallBinding::new(call_id, 0).unwrap()],
13437 )
13438 .unwrap();
13439 LLMResponse::new("", FinishReason::ToolCall)
13440 .with_provider_state(state)
13441 .unwrap()
13442 .with_tool_calls(vec![call])
13443 .unwrap()
13444 }
13445
13446 struct TerminalHistoryProvider {
13447 calls: Arc<std::sync::atomic::AtomicU32>,
13448 }
13449
13450 struct DroppingSignedAssistantMemory {
13451 messages: RwLock<Vec<ChatMessage>>,
13452 }
13453
13454 struct DroppingEarlierSequentialMemory {
13455 messages: RwLock<Vec<ChatMessage>>,
13456 signed_seen: std::sync::atomic::AtomicUsize,
13457 }
13458
13459 #[async_trait]
13460 impl ai_agents_core::Memory for DroppingSignedAssistantMemory {
13461 async fn add_message(&self, message: ChatMessage) -> Result<()> {
13462 let signed = message.role == ai_agents_core::Role::Assistant
13463 && ai_agents_core::decode_native_tool_call_markers(&message.content)
13464 .map_err(|error| AgentError::LLM(error.to_string()))?
13465 .is_some_and(|batch| batch.provider_state().is_some());
13466 if !signed {
13467 self.messages.write().push(message);
13468 }
13469 Ok(())
13470 }
13471
13472 async fn get_messages(&self, limit: Option<usize>) -> Result<Vec<ChatMessage>> {
13473 let messages = self.messages.read();
13474 let start = limit
13475 .map(|limit| messages.len().saturating_sub(limit))
13476 .unwrap_or(0);
13477 Ok(messages[start..].to_vec())
13478 }
13479
13480 async fn clear(&self) -> Result<()> {
13481 self.messages.write().clear();
13482 Ok(())
13483 }
13484
13485 fn len(&self) -> usize {
13486 self.messages.read().len()
13487 }
13488
13489 async fn restore(&self, snapshot: ai_agents_core::MemorySnapshot) -> Result<()> {
13490 *self.messages.write() = snapshot.messages;
13491 Ok(())
13492 }
13493 }
13494
13495 #[async_trait]
13496 impl ai_agents_memory::Memory for DroppingSignedAssistantMemory {}
13497
13498 #[async_trait]
13499 impl ai_agents_core::Memory for DroppingEarlierSequentialMemory {
13500 async fn add_message(&self, message: ChatMessage) -> Result<()> {
13501 let signed = message.role == ai_agents_core::Role::Assistant
13502 && ai_agents_core::decode_native_tool_call_markers(&message.content)
13503 .map_err(|error| AgentError::LLM(error.to_string()))?
13504 .is_some_and(|batch| batch.provider_state().is_some());
13505 let mut messages = self.messages.write();
13506 if signed && self.signed_seen.fetch_add(1, Ordering::SeqCst) == 1 {
13507 messages.retain(|stored| !stored.content.contains("seq-call-1"));
13508 }
13509 messages.push(message);
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 self.signed_seen.store(0, Ordering::SeqCst);
13524 Ok(())
13525 }
13526
13527 fn len(&self) -> usize {
13528 self.messages.read().len()
13529 }
13530
13531 async fn restore(&self, snapshot: ai_agents_core::MemorySnapshot) -> Result<()> {
13532 *self.messages.write() = snapshot.messages;
13533 self.signed_seen.store(0, Ordering::SeqCst);
13534 Ok(())
13535 }
13536 }
13537
13538 #[async_trait]
13539 impl ai_agents_memory::Memory for DroppingEarlierSequentialMemory {}
13540
13541 #[async_trait]
13542 impl LLMProvider for TerminalHistoryProvider {
13543 async fn complete(
13544 &self,
13545 _messages: &[ChatMessage],
13546 _config: Option<&LLMConfig>,
13547 ) -> std::result::Result<LLMResponse, LLMError> {
13548 self.calls.fetch_add(1, Ordering::SeqCst);
13549 Err(LLMError::Serialization(
13550 "native history integrity failure".to_string(),
13551 ))
13552 }
13553
13554 async fn complete_stream(
13555 &self,
13556 _messages: &[ChatMessage],
13557 _config: Option<&LLMConfig>,
13558 ) -> std::result::Result<
13559 Box<dyn Stream<Item = std::result::Result<LLMChunk, LLMError>> + Unpin + Send>,
13560 LLMError,
13561 > {
13562 Err(LLMError::Serialization(
13563 "native history integrity failure".to_string(),
13564 ))
13565 }
13566
13567 fn provider_name(&self) -> &str {
13568 "terminal-history"
13569 }
13570
13571 fn supports(&self, _feature: LLMFeature) -> bool {
13572 false
13573 }
13574
13575 fn is_terminal_error(&self, error: &LLMError) -> bool {
13576 matches!(error, LLMError::Serialization(_))
13577 }
13578 }
13579
13580 fn disambiguation_state_machine(
13582 state_enabled: Option<bool>,
13583 require_confirmation: bool,
13584 ) -> Arc<StateMachine> {
13585 let definition = ai_agents_state::StateDefinition {
13586 prompt: Some("Handle the resolved request.".to_string()),
13587 disambiguation: Some(ai_agents_disambiguation::StateDisambiguationOverride {
13588 enabled: state_enabled,
13589 require_confirmation,
13590 ..Default::default()
13591 }),
13592 ..Default::default()
13593 };
13594 let review = ai_agents_state::StateDefinition {
13595 prompt: Some("Review a fresh request.".to_string()),
13596 ..Default::default()
13597 };
13598 Arc::new(
13599 StateMachine::new(ai_agents_state::StateConfig {
13600 initial: "active".to_string(),
13601 states: std::collections::HashMap::from([
13602 ("active".to_string(), definition),
13603 ("review".to_string(), review),
13604 ]),
13605 global_transitions: Vec::new(),
13606 fallback: None,
13607 max_no_transition: None,
13608 regenerate_on_transition: true,
13609 })
13610 .unwrap(),
13611 )
13612 }
13613
13614 fn state_disambiguation_agent(
13616 responses: Vec<&str>,
13617 manager_enabled: bool,
13618 state_enabled: Option<bool>,
13619 require_confirmation: bool,
13620 ) -> (RuntimeAgent, MockLLMProvider) {
13621 state_disambiguation_agent_with_skills(
13622 responses,
13623 manager_enabled,
13624 state_enabled,
13625 require_confirmation,
13626 Vec::new(),
13627 )
13628 }
13629
13630 fn state_disambiguation_agent_with_skills(
13632 responses: Vec<&str>,
13633 manager_enabled: bool,
13634 state_enabled: Option<bool>,
13635 require_confirmation: bool,
13636 skills: Vec<SkillDefinition>,
13637 ) -> (RuntimeAgent, MockLLMProvider) {
13638 let mut mock = MockLLMProvider::new("state-confirmation");
13639 mock.set_responses(responses.into_iter().map(String::from).collect(), false);
13640 let observed = mock.clone();
13641 let agent = AgentBuilder::new()
13642 .system_prompt("Handle requests.")
13643 .llm(Arc::new(mock.clone()))
13644 .llm_alias("router", Arc::new(mock))
13645 .state_machine(disambiguation_state_machine(
13646 state_enabled,
13647 require_confirmation,
13648 ))
13649 .skills(skills)
13650 .build()
13651 .unwrap()
13652 .with_disambiguation(DisambiguationConfig {
13653 enabled: manager_enabled,
13654 ..Default::default()
13655 });
13656 (agent, observed)
13657 }
13658
13659 fn confirmation_skill() -> SkillDefinition {
13661 SkillDefinition {
13662 id: "send_report".to_string(),
13663 description: "Send a report after clarification".to_string(),
13664 trigger: "When the user asks to send a report".to_string(),
13665 steps: vec![SkillStep::Prompt {
13666 prompt: "Execute confirmed report skill for: {{ input }}".to_string(),
13667 llm: None,
13668 }],
13669 reasoning: None,
13670 reflection: None,
13671 disambiguation: Some(ai_agents_disambiguation::SkillDisambiguationOverride {
13672 enabled: Some(true),
13673 ..Default::default()
13674 }),
13675 }
13676 }
13677
13678 fn confirmation_skill_call_count(observed: &MockLLMProvider) -> usize {
13680 observed
13681 .call_history()
13682 .iter()
13683 .filter(|call| {
13684 call.messages
13685 .iter()
13686 .any(|message| message.content.contains("Execute confirmed report skill"))
13687 })
13688 .count()
13689 }
13690
13691 struct BlockingRuntimeConfirmationObserver {
13692 entered: tokio::sync::Barrier,
13693 release: tokio::sync::Notify,
13694 }
13695
13696 impl BlockingRuntimeConfirmationObserver {
13697 fn new() -> Self {
13698 Self {
13699 entered: tokio::sync::Barrier::new(2),
13700 release: tokio::sync::Notify::new(),
13701 }
13702 }
13703 }
13704
13705 struct ResetOnTransitionHooks {
13706 agent: parking_lot::Mutex<Option<Weak<RuntimeAgent>>>,
13707 invoked: AtomicBool,
13708 }
13709
13710 #[async_trait]
13711 impl AgentHooks for ResetOnTransitionHooks {
13712 async fn on_state_transition(&self, _from: Option<&str>, _to: &str, _reason: &str) {
13713 if self.invoked.swap(true, Ordering::SeqCst) {
13714 return;
13715 }
13716 let agent = self.agent.lock().as_ref().and_then(Weak::upgrade);
13717 if let Some(agent) = agent {
13718 agent.reset().await.unwrap();
13719 }
13720 }
13721 }
13722
13723 impl ClarificationObserver for BlockingRuntimeConfirmationObserver {
13724 fn observe_question<'a>(
13725 &'a self,
13726 future: ClarificationQuestionFuture<'a>,
13727 ) -> ClarificationQuestionFuture<'a> {
13728 future
13729 }
13730
13731 fn observe_parse<'a>(
13732 &'a self,
13733 future: ClarificationParseFuture<'a>,
13734 ) -> ClarificationParseFuture<'a> {
13735 future
13736 }
13737
13738 fn observe_confirmation_parse<'a>(
13739 &'a self,
13740 future: ConfirmationParseFuture<'a>,
13741 ) -> ConfirmationParseFuture<'a> {
13742 Box::pin(async move {
13743 self.entered.wait().await;
13744 self.release.notified().await;
13745 future.await
13746 })
13747 }
13748 }
13749
13750 #[tokio::test]
13751 async fn state_confirmation_blocks_redispatch_until_explicit_agreement() {
13752 let (agent, observed) = state_disambiguation_agent(
13753 vec![
13754 r#"{"is_ambiguous":true,"confidence":0.2,"ambiguity_type":"missing_target","reasoning":"target missing","what_is_unclear":["target"],"detected_language":"en"}"#,
13755 r#"{"question":"What should I send?","options":null}"#,
13756 r#"{"status":"answered","selected_option":null,"enriched_input":"Send the report to Ada","resolved":{"intent":"send_report"}}"#,
13757 r#"{"question":"Should I send the report to Ada?"}"#,
13758 r#"{"status":"confirmed"}"#,
13759 "Request executed.",
13760 ],
13761 true,
13762 None,
13763 true,
13764 );
13765
13766 let clarification = agent.chat("Send it").await.unwrap();
13767 assert_eq!(clarification.content, "What should I send?");
13768 assert_eq!(observed.call_count(), 2);
13769
13770 let confirmation = agent.chat("The report to Ada").await.unwrap();
13771 assert_eq!(confirmation.content, "Should I send the report to Ada?");
13772 assert_eq!(
13773 confirmation
13774 .metadata
13775 .as_ref()
13776 .and_then(|metadata| metadata.get("disambiguation"))
13777 .and_then(|metadata| metadata.get("status"))
13778 .and_then(Value::as_str),
13779 Some("awaiting_confirmation")
13780 );
13781 assert_eq!(observed.call_count(), 4);
13782
13783 let completed = agent.chat("Yes").await.unwrap();
13784 assert_eq!(completed.content, "Request executed.");
13785 assert_eq!(observed.call_count(), 6);
13786 }
13787
13788 #[tokio::test]
13789 async fn streaming_state_confirmation_ends_the_turn_before_redispatch() {
13790 let (agent, observed) = state_disambiguation_agent(
13791 vec![
13792 r#"{"is_ambiguous":true,"confidence":0.2,"ambiguity_type":"missing_target","reasoning":"target missing","what_is_unclear":["target"],"detected_language":"en"}"#,
13793 r#"{"question":"What should I send?","options":null}"#,
13794 r#"{"status":"answered","selected_option":null,"enriched_input":"Send the report to Ada","resolved":{"intent":"send_report"}}"#,
13795 r#"{"question":"Should I send the report to Ada?"}"#,
13796 r#"{"status":"confirmed"}"#,
13797 "Request executed.",
13798 ],
13799 true,
13800 None,
13801 true,
13802 );
13803
13804 let mut clarification_stream = agent.chat_stream("Send it").await.unwrap();
13805 let mut clarification = String::new();
13806 while let Some(chunk) = clarification_stream.next().await {
13807 match chunk {
13808 StreamChunk::Content { text } => clarification.push_str(&text),
13809 StreamChunk::Done {} => break,
13810 StreamChunk::Error { message } => panic!("unexpected stream error: {message}"),
13811 _ => {}
13812 }
13813 }
13814 assert_eq!(clarification, "What should I send?");
13815 assert_eq!(observed.call_count(), 2);
13816
13817 let mut confirmation_stream = agent.chat_stream_events("The report to Ada").await.unwrap();
13818 let mut confirmation = None;
13819 while let Some(event) = confirmation_stream.next().await {
13820 match event {
13821 AgentStreamEvent::Final(response) => confirmation = Some(response),
13822 AgentStreamEvent::Chunk(StreamChunk::Error { message }) => {
13823 panic!("unexpected stream error: {message}")
13824 }
13825 AgentStreamEvent::Chunk(_) => {}
13826 }
13827 }
13828 let confirmation = confirmation.expect("confirmation must finalize");
13829 assert_eq!(confirmation.content, "Should I send the report to Ada?");
13830 assert_eq!(
13831 confirmation
13832 .metadata
13833 .as_ref()
13834 .and_then(|metadata| metadata.get("disambiguation"))
13835 .and_then(|metadata| metadata.get("status"))
13836 .and_then(Value::as_str),
13837 Some("awaiting_confirmation")
13838 );
13839 assert_eq!(observed.call_count(), 4);
13840
13841 let mut completed_stream = agent.chat_stream("Yes").await.unwrap();
13842 let mut completed = String::new();
13843 while let Some(chunk) = completed_stream.next().await {
13844 match chunk {
13845 StreamChunk::Content { text } => completed.push_str(&text),
13846 StreamChunk::Done {} => break,
13847 StreamChunk::Error { message } => panic!("unexpected stream error: {message}"),
13848 _ => {}
13849 }
13850 }
13851 assert_eq!(completed, "Request executed.");
13852 assert_eq!(observed.call_count(), 6);
13853 }
13854
13855 #[tokio::test]
13857 async fn root_turn_gate_serializes_blocking_and_streaming_entry_points() {
13858 let (complete_entered, mut complete_events) = tokio::sync::mpsc::unbounded_channel();
13859 let agent = Arc::new(
13860 AgentBuilder::new()
13861 .system_prompt("Serialize root turns.")
13862 .llm(Arc::new(RootTurnProbeProvider { complete_entered }))
13863 .build()
13864 .unwrap(),
13865 );
13866 let blocking_agent = Arc::clone(&agent);
13867
13868 let legacy_stream = agent.chat_stream("stream owner").await.unwrap();
13869 assert!(agent.root_turn_gate.try_lock().is_err());
13870 let blocking = tokio::spawn(async move { blocking_agent.chat("blocked").await.unwrap() });
13871 assert!(
13872 tokio::time::timeout(std::time::Duration::from_millis(50), complete_events.recv())
13873 .await
13874 .is_err(),
13875 "blocking turn reached the provider while the legacy stream owned the root gate"
13876 );
13877
13878 drop(legacy_stream);
13879 assert_eq!(
13880 tokio::time::timeout(std::time::Duration::from_secs(2), complete_events.recv())
13881 .await
13882 .expect("blocking turn did not enter after stream drop"),
13883 Some(())
13884 );
13885 let response = tokio::time::timeout(std::time::Duration::from_secs(2), blocking)
13886 .await
13887 .expect("blocking turn did not finish after stream drop")
13888 .unwrap();
13889 assert_eq!(response.content, "blocking complete");
13890
13891 let mut event_stream = agent.chat_stream_events("event terminal").await.unwrap();
13892 assert!(agent.root_turn_gate.try_lock().is_err());
13893 let mut saw_final = false;
13894 while let Some(event) = event_stream.next().await {
13895 if matches!(event, AgentStreamEvent::Final(_)) {
13896 saw_final = true;
13897 break;
13898 }
13899 }
13900 assert!(saw_final);
13901 assert!(
13902 agent.root_turn_gate.try_lock().is_ok(),
13903 "authoritative terminal event retained the root gate"
13904 );
13905 }
13906
13907 #[tokio::test]
13909 async fn response_hook_rejects_same_runtime_chat_reentry() {
13910 let hooks = Arc::new(ResponseChatHooks {
13911 target: parking_lot::Mutex::new(None),
13912 invoked: AtomicBool::new(false),
13913 nested_result: parking_lot::Mutex::new(None),
13914 });
13915 let agent = Arc::new(
13916 AgentBuilder::new()
13917 .system_prompt("Reject response hook reentry.")
13918 .llm(Arc::new(mock_with_response("outer response")))
13919 .hooks(hooks.clone())
13920 .build()
13921 .unwrap(),
13922 );
13923 *hooks.target.lock() = Some(Arc::downgrade(&agent));
13924
13925 let response = tokio::time::timeout(
13926 std::time::Duration::from_secs(2),
13927 agent.chat("outer request"),
13928 )
13929 .await
13930 .expect("same-runtime response hook reentry must fail without deadlocking")
13931 .unwrap();
13932
13933 assert_eq!(response.content, "outer response");
13934 let nested_result = hooks
13935 .nested_result
13936 .lock()
13937 .clone()
13938 .expect("response hook must record its nested call");
13939 let error = nested_result.expect_err("same-runtime nested chat must be rejected");
13940 assert!(error.contains("reentrant root turn ownership"));
13941 }
13942
13943 #[tokio::test]
13945 async fn root_turn_gate_allows_nested_runtime_and_rejects_cycles() {
13946 let agent_a = AgentBuilder::new()
13947 .system_prompt("Runtime A.")
13948 .llm(Arc::new(mock_with_response("response A")))
13949 .build()
13950 .unwrap();
13951 let agent_b = AgentBuilder::new()
13952 .system_prompt("Runtime B.")
13953 .llm(Arc::new(mock_with_response("response B")))
13954 .build()
13955 .unwrap();
13956 let RootTurnAdmission {
13957 guard: guard_a,
13958 identity_stack: stack_a,
13959 } = agent_a.acquire_root_turn().await.unwrap();
13960
13961 let cycle_error = scope_runtime_gate_identity_stack(&stack_a, async {
13962 let RootTurnAdmission {
13963 guard: guard_b,
13964 identity_stack: stack_b,
13965 } = agent_b
13966 .acquire_root_turn()
13967 .await
13968 .expect("runtime B must acquire a different gate");
13969 let result =
13970 scope_runtime_gate_identity_stack(&stack_b, agent_a.acquire_root_turn()).await;
13971 drop(guard_b);
13972 match result {
13973 Err(error) => error,
13974 Ok(_) => panic!("runtime A accepted a repeated gate identity"),
13975 }
13976 })
13977 .await;
13978 drop(guard_a);
13979
13980 assert!(
13981 cycle_error
13982 .to_string()
13983 .contains("reentrant root turn ownership")
13984 );
13985 }
13986
13987 #[tokio::test]
13989 async fn concurrent_orchestration_propagates_root_gate_ancestry() {
13990 let registry = Arc::new(crate::spawner::AgentRegistry::new());
13991 let hooks_a = Arc::new(ConcurrentResponseHooks {
13992 registry: Arc::downgrade(®istry),
13993 child_id: "runtime-b".to_string(),
13994 invoked: AtomicBool::new(false),
13995 nested_result: parking_lot::Mutex::new(None),
13996 });
13997 let hooks_b = Arc::new(ResponseChatHooks {
13998 target: parking_lot::Mutex::new(None),
13999 invoked: AtomicBool::new(false),
14000 nested_result: parking_lot::Mutex::new(None),
14001 });
14002 let agent_a = AgentBuilder::new()
14003 .system_prompt("Runtime A dispatches runtime B concurrently.")
14004 .llm(Arc::new(mock_with_response("response A")))
14005 .hooks(hooks_a.clone())
14006 .build()
14007 .unwrap();
14008 let agent_b = AgentBuilder::new()
14009 .system_prompt("Runtime B attempts to re-enter runtime A.")
14010 .llm(Arc::new(mock_with_response("response B")))
14011 .hooks(hooks_b.clone())
14012 .build()
14013 .unwrap();
14014 let spec_a = crate::spec::AgentSpec {
14015 name: "runtime-a".to_string(),
14016 system_prompt: "Runtime A dispatches runtime B concurrently.".to_string(),
14017 ..crate::spec::AgentSpec::default()
14018 };
14019 let spec_b = crate::spec::AgentSpec {
14020 name: "runtime-b".to_string(),
14021 system_prompt: "Runtime B attempts to re-enter runtime A.".to_string(),
14022 ..crate::spec::AgentSpec::default()
14023 };
14024 registry
14025 .register(crate::spawner::SpawnedAgent::from_runtime(
14026 "runtime-a".to_string(),
14027 agent_a,
14028 spec_a,
14029 ))
14030 .await
14031 .unwrap();
14032 registry
14033 .register(crate::spawner::SpawnedAgent::from_runtime(
14034 "runtime-b".to_string(),
14035 agent_b,
14036 spec_b,
14037 ))
14038 .await
14039 .unwrap();
14040 let runtime_a = registry.get("runtime-a").unwrap();
14041 *hooks_b.target.lock() = Some(Arc::downgrade(&runtime_a));
14042
14043 let response = tokio::time::timeout(
14044 std::time::Duration::from_secs(2),
14045 runtime_a.chat("outer concurrent request"),
14046 )
14047 .await
14048 .expect("concurrent orchestration cycle must fail without deadlocking")
14049 .unwrap();
14050
14051 assert_eq!(response.content, "response A");
14052 let child_result = hooks_a
14053 .nested_result
14054 .lock()
14055 .clone()
14056 .expect("runtime A hook must record runtime B completion");
14057 assert_eq!(child_result.unwrap(), "response B");
14058 let cycle_result = hooks_b
14059 .nested_result
14060 .lock()
14061 .clone()
14062 .expect("runtime B hook must record runtime A reentry");
14063 assert!(
14064 cycle_result
14065 .expect_err("runtime A accepted a repeated gate identity")
14066 .contains("reentrant root turn ownership")
14067 );
14068 }
14069
14070 fn skill_clarification_responses() -> Vec<&'static str> {
14073 vec![
14074 r#"{"is_ambiguous":false,"confidence":0.99,"ambiguity_type":null,"reasoning":"top-level clear","what_is_unclear":[],"detected_language":"en"}"#,
14075 "send_report",
14076 r#"{"is_ambiguous":true,"confidence":0.2,"ambiguity_type":"missing_target","reasoning":"target missing","what_is_unclear":["target"],"detected_language":"en"}"#,
14077 r#"{"question":"What should I send?","options":null}"#,
14078 ]
14079 }
14080
14081 #[tokio::test]
14087 async fn test_stream_skill_clarification_memory_matches_blocking() {
14088 let (blocking_agent, _) = state_disambiguation_agent_with_skills(
14089 skill_clarification_responses(),
14090 true,
14091 None,
14092 true,
14093 vec![confirmation_skill()],
14094 );
14095 let blocking = blocking_agent.chat("Send it").await.unwrap();
14096 let blocking_messages = blocking_agent.memory.get_messages(None).await.unwrap();
14097
14098 let (streaming_agent, _) = state_disambiguation_agent_with_skills(
14099 skill_clarification_responses(),
14100 true,
14101 None,
14102 true,
14103 vec![confirmation_skill()],
14104 );
14105 let (content, chunks, streamed) = collect_stream_events(&streaming_agent, "Send it").await;
14106 assert!(!chunks.iter().any(StreamChunk::is_error), "{chunks:?}");
14107 let streamed = streamed.expect("skill clarification must finalize as Final");
14108 let streaming_messages = streaming_agent.memory.get_messages(None).await.unwrap();
14109
14110 assert_eq!(blocking.content, "What should I send?");
14111 assert_eq!(streamed.content, blocking.content);
14112 assert_eq!(content, streamed.content);
14113 assert_eq!(
14114 blocking
14115 .metadata
14116 .as_ref()
14117 .and_then(|m| m.get("disambiguation")),
14118 streamed
14119 .metadata
14120 .as_ref()
14121 .and_then(|m| m.get("disambiguation")),
14122 );
14123 assert_eq!(
14124 streamed
14125 .metadata
14126 .as_ref()
14127 .and_then(|m| m.get("disambiguation"))
14128 .and_then(|d| d.get("status"))
14129 .and_then(Value::as_str),
14130 Some("awaiting_clarification"),
14131 );
14132 let shape = |messages: &[ChatMessage]| {
14133 messages
14134 .iter()
14135 .map(|m| (format!("{:?}", m.role), m.content.clone()))
14136 .collect::<Vec<_>>()
14137 };
14138 assert_eq!(shape(&blocking_messages), shape(&streaming_messages));
14139 assert_eq!(
14140 shape(&streaming_messages),
14141 vec![
14142 ("User".to_string(), "Send it".to_string()),
14143 ("Assistant".to_string(), "What should I send?".to_string()),
14144 ],
14145 );
14146 assert_eq!(
14147 *streaming_agent.pending_skill_id.read(),
14148 Some("send_report".to_string()),
14149 );
14150 }
14151
14152 #[tokio::test]
14154 async fn test_stream_skill_clarification_memory_failure_surfaces_as_error() {
14155 let build = || {
14157 let mut mock = MockLLMProvider::new("skill-clarification");
14158 mock.set_responses(
14159 skill_clarification_responses()
14160 .into_iter()
14161 .map(String::from)
14162 .collect(),
14163 false,
14164 );
14165 AgentBuilder::new()
14166 .system_prompt("Handle requests.")
14167 .llm(Arc::new(mock.clone()))
14168 .llm_alias("router", Arc::new(mock))
14169 .state_machine(disambiguation_state_machine(None, true))
14170 .skills(vec![confirmation_skill()])
14171 .memory(Arc::new(FailingMemory {
14172 messages: parking_lot::RwLock::new(Vec::new()),
14173 fail_on_add: 2,
14174 adds: std::sync::atomic::AtomicUsize::new(0),
14175 }))
14176 .build()
14177 .unwrap()
14178 .with_disambiguation(DisambiguationConfig {
14179 enabled: true,
14180 ..Default::default()
14181 })
14182 };
14183
14184 let blocking = build().chat("Send it").await;
14185 assert!(
14186 blocking.is_err(),
14187 "blocking must surface the failed clarification write: {blocking:?}"
14188 );
14189
14190 let (_, chunks, streamed) = collect_stream_events(&build(), "Send it").await;
14191 assert!(
14192 streamed.is_none(),
14193 "a failed write must not finalize the turn"
14194 );
14195 assert!(
14196 chunks.iter().any(|chunk| matches!(
14197 chunk,
14198 StreamChunk::Error { message } if message.contains("simulated memory failure")
14199 )),
14200 "streaming must surface the failed clarification write: {chunks:?}"
14201 );
14202 }
14203
14204 #[tokio::test]
14206 async fn confirmed_skill_route_executes_exactly_once() {
14207 let (agent, observed) = state_disambiguation_agent_with_skills(
14208 vec![
14209 r#"{"is_ambiguous":false,"confidence":0.99,"ambiguity_type":null,"reasoning":"top-level clear","what_is_unclear":[],"detected_language":"en"}"#,
14210 "send_report",
14211 r#"{"is_ambiguous":true,"confidence":0.2,"ambiguity_type":"missing_target","reasoning":"target missing","what_is_unclear":["target"],"detected_language":"en"}"#,
14212 r#"{"question":"What should I send?","options":null}"#,
14213 r#"{"status":"answered","selected_option":null,"enriched_input":"Send the report to Ada","resolved":{"intent":"send_report"}}"#,
14214 r#"{"question":"Should I send the report to Ada?"}"#,
14215 r#"{"status":"confirmed"}"#,
14216 r#"{"is_ambiguous":false,"confidence":0.99,"ambiguity_type":null,"reasoning":"resolved","what_is_unclear":[],"detected_language":"en"}"#,
14217 "Report skill executed.",
14218 ],
14219 true,
14220 None,
14221 true,
14222 vec![confirmation_skill()],
14223 );
14224
14225 let clarification = agent.chat("Send it").await.unwrap();
14226 assert_eq!(clarification.content, "What should I send?");
14227 assert_eq!(confirmation_skill_call_count(&observed), 0);
14228
14229 let confirmation = agent.chat("The report to Ada").await.unwrap();
14230 assert_eq!(confirmation.content, "Should I send the report to Ada?");
14231 assert_eq!(
14232 confirmation
14233 .metadata
14234 .as_ref()
14235 .and_then(|metadata| metadata.get("disambiguation"))
14236 .and_then(|metadata| metadata.get("status"))
14237 .and_then(Value::as_str),
14238 Some("awaiting_confirmation")
14239 );
14240 assert_eq!(confirmation_skill_call_count(&observed), 0);
14241
14242 let completed = agent.chat("Yes").await.unwrap();
14243 assert_eq!(completed.content, "Report skill executed.");
14244 assert_eq!(confirmation_skill_call_count(&observed), 1);
14245 assert!(agent.pending_skill_id.read().is_none());
14246 let messages = agent.memory.get_messages(None).await.unwrap();
14247 assert!(!messages.iter().any(|message| message.content == "Yes"));
14248 }
14249
14250 #[tokio::test]
14252 async fn confirmed_skill_recheck_preserves_new_clarification_metadata() {
14253 let (agent, observed) = state_disambiguation_agent_with_skills(
14254 vec![
14255 r#"{"is_ambiguous":false,"confidence":0.99,"ambiguity_type":null,"reasoning":"top-level clear","what_is_unclear":[],"detected_language":"en"}"#,
14256 "send_report",
14257 r#"{"is_ambiguous":true,"confidence":0.2,"ambiguity_type":"missing_target","reasoning":"target missing","what_is_unclear":["target"],"detected_language":"en"}"#,
14258 r#"{"question":"What should I send?","options":null}"#,
14259 r#"{"status":"answered","selected_option":null,"enriched_input":"Send the report to Ada","resolved":{"intent":"send_report"}}"#,
14260 r#"{"question":"Should I send the report to Ada?"}"#,
14261 r#"{"status":"confirmed"}"#,
14262 r#"{"is_ambiguous":true,"confidence":0.3,"ambiguity_type":"missing_parameters","reasoning":"timing missing","what_is_unclear":["timing"],"detected_language":"en"}"#,
14263 r#"{"question":"When should I send it?","options":null}"#,
14264 ],
14265 true,
14266 None,
14267 true,
14268 vec![confirmation_skill()],
14269 );
14270
14271 agent.chat("Send it").await.unwrap();
14272 agent.chat("The report to Ada").await.unwrap();
14273 let follow_up = agent.chat("Yes").await.unwrap();
14274
14275 assert_eq!(follow_up.content, "When should I send it?");
14276 let metadata = follow_up
14277 .metadata
14278 .as_ref()
14279 .and_then(|metadata| metadata.get("disambiguation"))
14280 .unwrap();
14281 assert_eq!(
14282 metadata.get("status").and_then(Value::as_str),
14283 Some("awaiting_clarification")
14284 );
14285 assert_eq!(
14286 metadata.get("skill_id").and_then(Value::as_str),
14287 Some("send_report")
14288 );
14289 assert!(metadata.get("detection").is_some());
14290 assert_eq!(confirmation_skill_call_count(&observed), 0);
14291 }
14292
14293 #[tokio::test]
14295 async fn rejected_skill_confirmation_never_executes() {
14296 let (agent, observed) = state_disambiguation_agent_with_skills(
14297 vec![
14298 r#"{"is_ambiguous":false,"confidence":0.99,"ambiguity_type":null,"reasoning":"top-level clear","what_is_unclear":[],"detected_language":"en"}"#,
14299 "send_report",
14300 r#"{"is_ambiguous":true,"confidence":0.2,"ambiguity_type":"missing_target","reasoning":"target missing","what_is_unclear":["target"],"detected_language":"en"}"#,
14301 r#"{"question":"What should I send?","options":null}"#,
14302 r#"{"status":"answered","selected_option":null,"enriched_input":"Send the report to Ada","resolved":{"intent":"send_report"}}"#,
14303 r#"{"question":"Should I send the report to Ada?"}"#,
14304 r#"{"status":"rejected"}"#,
14305 "Confirmation rejected.",
14306 ],
14307 true,
14308 None,
14309 true,
14310 vec![confirmation_skill()],
14311 );
14312
14313 agent.chat("Send it").await.unwrap();
14314 agent.chat("The report to Ada").await.unwrap();
14315 let rejected = agent.chat("No").await.unwrap();
14316
14317 assert_eq!(rejected.content, "Confirmation rejected.");
14318 assert_eq!(confirmation_skill_call_count(&observed), 0);
14319 assert!(agent.pending_skill_id.read().is_none());
14320 }
14321
14322 #[tokio::test]
14324 async fn reset_invalidates_pending_skill_confirmation_before_streaming_input() {
14325 let (agent, observed) = state_disambiguation_agent_with_skills(
14326 vec![
14327 r#"{"is_ambiguous":false,"confidence":0.99,"ambiguity_type":null,"reasoning":"top-level clear","what_is_unclear":[],"detected_language":"en"}"#,
14328 "send_report",
14329 r#"{"is_ambiguous":true,"confidence":0.2,"ambiguity_type":"missing_target","reasoning":"target missing","what_is_unclear":["target"],"detected_language":"en"}"#,
14330 r#"{"question":"What should I send?","options":null}"#,
14331 r#"{"status":"answered","selected_option":null,"enriched_input":"Send the report to Ada","resolved":{"intent":"send_report"}}"#,
14332 r#"{"question":"Should I send the report to Ada?"}"#,
14333 r#"{"is_ambiguous":false,"confidence":0.99,"ambiguity_type":null,"reasoning":"fresh input","what_is_unclear":[],"detected_language":"en"}"#,
14334 "none",
14335 "Fresh response.",
14336 ],
14337 true,
14338 None,
14339 true,
14340 vec![confirmation_skill()],
14341 );
14342
14343 agent.chat("Send it").await.unwrap();
14344 agent.chat("The report to Ada").await.unwrap();
14345 agent.reset().await.unwrap();
14346 assert!(agent.pending_skill_id.read().is_none());
14347 assert!(
14348 !agent
14349 .disambiguation_manager()
14350 .unwrap()
14351 .has_pending_clarification()
14352 .await
14353 );
14354
14355 let mut stream = agent.chat_stream("Yes").await.unwrap();
14356 let mut content = String::new();
14357 while let Some(chunk) = stream.next().await {
14358 match chunk {
14359 StreamChunk::Content { text } => content.push_str(&text),
14360 StreamChunk::Done {} => break,
14361 StreamChunk::Error { message } => panic!("unexpected stream error: {message}"),
14362 _ => {}
14363 }
14364 }
14365
14366 assert_eq!(content, "Fresh response.");
14367 assert_eq!(confirmation_skill_call_count(&observed), 0);
14368 }
14369
14370 #[tokio::test]
14372 async fn trait_reset_clears_pending_skill_confirmation() {
14373 let (agent, _) = state_disambiguation_agent_with_skills(
14374 vec![
14375 r#"{"is_ambiguous":false,"confidence":0.99,"ambiguity_type":null,"reasoning":"top-level clear","what_is_unclear":[],"detected_language":"en"}"#,
14376 "send_report",
14377 r#"{"is_ambiguous":true,"confidence":0.2,"ambiguity_type":"missing_target","reasoning":"target missing","what_is_unclear":["target"],"detected_language":"en"}"#,
14378 r#"{"question":"What should I send?","options":null}"#,
14379 r#"{"status":"answered","selected_option":null,"enriched_input":"Send the report to Ada","resolved":{"intent":"send_report"}}"#,
14380 r#"{"question":"Should I send the report to Ada?"}"#,
14381 ],
14382 true,
14383 None,
14384 true,
14385 vec![confirmation_skill()],
14386 );
14387
14388 agent.chat("Send it").await.unwrap();
14389 agent.chat("The report to Ada").await.unwrap();
14390 <RuntimeAgent as Agent>::reset(&agent).await.unwrap();
14391
14392 assert!(agent.pending_skill_id.read().is_none());
14393 assert!(
14394 !agent
14395 .disambiguation_manager()
14396 .unwrap()
14397 .has_pending_clarification()
14398 .await
14399 );
14400 }
14401
14402 #[tokio::test]
14404 async fn state_change_invalidates_pending_skill_confirmation() {
14405 let (agent, observed) = state_disambiguation_agent_with_skills(
14406 vec![
14407 r#"{"is_ambiguous":false,"confidence":0.99,"ambiguity_type":null,"reasoning":"top-level clear","what_is_unclear":[],"detected_language":"en"}"#,
14408 "send_report",
14409 r#"{"is_ambiguous":true,"confidence":0.2,"ambiguity_type":"missing_target","reasoning":"target missing","what_is_unclear":["target"],"detected_language":"en"}"#,
14410 r#"{"question":"What should I send?","options":null}"#,
14411 r#"{"status":"answered","selected_option":null,"enriched_input":"Send the report to Ada","resolved":{"intent":"send_report"}}"#,
14412 r#"{"question":"Should I send the report to Ada?"}"#,
14413 r#"{"is_ambiguous":false,"confidence":0.99,"ambiguity_type":null,"reasoning":"fresh input","what_is_unclear":[],"detected_language":"en"}"#,
14414 "none",
14415 "Fresh response.",
14416 ],
14417 true,
14418 None,
14419 true,
14420 vec![confirmation_skill()],
14421 );
14422
14423 agent.chat("Send it").await.unwrap();
14424 agent.chat("The report to Ada").await.unwrap();
14425 agent.transition_to("review").await.unwrap();
14426 let cancelled = agent.chat("Yes").await.unwrap();
14427
14428 assert_eq!(cancelled.content, "Fresh response.");
14429 assert_eq!(confirmation_skill_call_count(&observed), 0);
14430 assert!(agent.pending_skill_id.read().is_none());
14431 }
14432
14433 #[tokio::test]
14435 async fn in_flight_confirmation_cannot_redispatch_after_reset() {
14436 let (mut agent, observed) = state_disambiguation_agent_with_skills(
14437 vec![
14438 r#"{"is_ambiguous":false,"confidence":0.99,"ambiguity_type":null,"reasoning":"top-level clear","what_is_unclear":[],"detected_language":"en"}"#,
14439 "send_report",
14440 r#"{"is_ambiguous":true,"confidence":0.2,"ambiguity_type":"missing_target","reasoning":"target missing","what_is_unclear":["target"],"detected_language":"en"}"#,
14441 r#"{"question":"What should I send?","options":null}"#,
14442 r#"{"status":"answered","selected_option":null,"enriched_input":"Send the report to Ada","resolved":{"intent":"send_report"}}"#,
14443 r#"{"question":"Should I send the report to Ada?"}"#,
14444 r#"{"status":"confirmed"}"#,
14445 "Confirmation cancelled.",
14446 ],
14447 true,
14448 None,
14449 true,
14450 vec![confirmation_skill()],
14451 );
14452 let observer = Arc::new(BlockingRuntimeConfirmationObserver::new());
14453 let manager = agent
14454 .disambiguation_manager
14455 .take()
14456 .unwrap()
14457 .with_clarification_observer(observer.clone());
14458 agent.disambiguation_manager = Some(manager);
14459 let agent = Arc::new(agent);
14460
14461 agent.chat("Send it").await.unwrap();
14462 agent.chat("The report to Ada").await.unwrap();
14463
14464 let confirming_agent = Arc::clone(&agent);
14465 let confirmation = tokio::spawn(async move { confirming_agent.chat("Yes").await });
14466 observer.entered.wait().await;
14467 agent.reset().await.unwrap();
14468 observer.release.notify_one();
14469
14470 let response = confirmation.await.unwrap().unwrap();
14471 assert_eq!(response.content, "Confirmation cancelled.");
14472 assert_eq!(confirmation_skill_call_count(&observed), 0);
14473 assert!(agent.pending_skill_id.read().is_none());
14474 }
14475
14476 #[tokio::test]
14478 async fn queued_reset_prevents_stale_confirmation_question_publication() {
14479 let (agent, observed) = state_disambiguation_agent(
14480 vec![
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 ],
14486 true,
14487 None,
14488 true,
14489 );
14490 let agent = Arc::new(agent);
14491 agent.chat("Send it").await.unwrap();
14492
14493 let admission = agent.disambiguation_admission.write().await;
14494 let (started_tx, started_rx) = tokio::sync::oneshot::channel();
14495 let resetting_agent = Arc::clone(&agent);
14496 let reset = tokio::spawn(async move {
14497 let _ = started_tx.send(());
14498 resetting_agent.reset().await
14499 });
14500 started_rx.await.unwrap();
14501 tokio::task::yield_now().await;
14502
14503 let responding_agent = Arc::clone(&agent);
14504 let response =
14505 tokio::spawn(async move { responding_agent.chat("The report to Ada").await });
14506 tokio::time::timeout(std::time::Duration::from_secs(2), async {
14507 while observed.call_count() < 4 {
14508 tokio::task::yield_now().await;
14509 }
14510 })
14511 .await
14512 .expect("clarification processing must reach terminal publication");
14513 drop(admission);
14514
14515 reset.await.unwrap().unwrap();
14516 let error = response.await.unwrap().unwrap_err();
14517 assert!(error.to_string().contains("ownership changed"));
14518 assert!(
14519 !agent
14520 .disambiguation_manager()
14521 .unwrap()
14522 .has_pending_clarification()
14523 .await
14524 );
14525 assert!(agent.memory.get_messages(None).await.unwrap().is_empty());
14526 }
14527
14528 #[tokio::test]
14530 async fn queued_reset_prevents_stale_skill_clarification_publication() {
14531 let (agent, observed) = state_disambiguation_agent_with_skills(
14532 vec![
14533 r#"{"is_ambiguous":false,"confidence":0.99,"ambiguity_type":null,"reasoning":"top-level clear","what_is_unclear":[],"detected_language":"en"}"#,
14534 "send_report",
14535 r#"{"is_ambiguous":true,"confidence":0.2,"ambiguity_type":"missing_target","reasoning":"target missing","what_is_unclear":["target"],"detected_language":"en"}"#,
14536 r#"{"question":"What should I send?","options":null}"#,
14537 ],
14538 true,
14539 None,
14540 true,
14541 vec![confirmation_skill()],
14542 );
14543 let agent = Arc::new(agent);
14544 let admission = agent.disambiguation_admission.write().await;
14545 let (started_tx, started_rx) = tokio::sync::oneshot::channel();
14546 let resetting_agent = Arc::clone(&agent);
14547 let reset = tokio::spawn(async move {
14548 let _ = started_tx.send(());
14549 resetting_agent.reset().await
14550 });
14551 started_rx.await.unwrap();
14552 tokio::task::yield_now().await;
14553
14554 let responding_agent = Arc::clone(&agent);
14555 let response = tokio::spawn(async move { responding_agent.chat("Send it").await });
14556 tokio::time::timeout(std::time::Duration::from_secs(2), async {
14557 while observed.call_count() < 4 {
14558 tokio::task::yield_now().await;
14559 }
14560 })
14561 .await
14562 .expect("skill clarification must reach terminal publication");
14563 drop(admission);
14564
14565 reset.await.unwrap().unwrap();
14566 let error = response.await.unwrap().unwrap_err();
14567 assert!(error.to_string().contains("ownership changed"));
14568 assert_eq!(confirmation_skill_call_count(&observed), 0);
14569 assert!(agent.pending_skill_id.read().is_none());
14570 assert!(agent.memory.get_messages(None).await.unwrap().is_empty());
14571 }
14572
14573 #[tokio::test]
14575 async fn transition_hook_can_reset_without_admission_deadlock() {
14576 let hooks = Arc::new(ResetOnTransitionHooks {
14577 agent: parking_lot::Mutex::new(None),
14578 invoked: AtomicBool::new(false),
14579 });
14580 let agent = Arc::new(
14581 AgentBuilder::new()
14582 .system_prompt("Test transition hook reentrancy.")
14583 .llm(Arc::new(mock_with_response("done")))
14584 .state_machine(disambiguation_state_machine(None, false))
14585 .build()
14586 .unwrap()
14587 .with_hooks(hooks.clone()),
14588 );
14589 *hooks.agent.lock() = Some(Arc::downgrade(&agent));
14590
14591 let transitioned = tokio::time::timeout(
14592 std::time::Duration::from_secs(2),
14593 agent.apply_transition_target("active", "review", "test transition", None),
14594 )
14595 .await
14596 .expect("transition hook reset must not deadlock")
14597 .unwrap();
14598
14599 assert!(transitioned);
14600 assert!(hooks.invoked.load(Ordering::SeqCst));
14601 assert_eq!(agent.current_state().as_deref(), Some("active"));
14602 }
14603
14604 #[tokio::test]
14606 async fn concurrent_transition_cannot_duplicate_exit_actions() {
14607 let gate = PathMutationGate::new();
14608 let active = ai_agents_state::StateDefinition {
14609 on_exit: vec![StateAction::Tool {
14610 tool: "transition_exit".to_string(),
14611 args: Some(serde_json::json!({"path": "./transition-exit.txt"})),
14612 }],
14613 ..Default::default()
14614 };
14615 let state_machine = Arc::new(
14616 StateMachine::new(ai_agents_state::StateConfig {
14617 initial: "active".to_string(),
14618 states: HashMap::from([
14619 ("active".to_string(), active),
14620 (
14621 "review".to_string(),
14622 ai_agents_state::StateDefinition::default(),
14623 ),
14624 ]),
14625 global_transitions: Vec::new(),
14626 fallback: None,
14627 max_no_transition: None,
14628 regenerate_on_transition: true,
14629 })
14630 .unwrap(),
14631 );
14632 let agent = Arc::new(
14633 AgentBuilder::new()
14634 .system_prompt("Test transition reservation.")
14635 .llm(Arc::new(mock_with_response("done")))
14636 .tool(Arc::new(BlockingPathMutationTool {
14637 id: "transition_exit",
14638 path_fields: vec![ai_agents_core::PathPolicyBinding::write("path")],
14639 gate: gate.clone(),
14640 }))
14641 .state_machine(state_machine)
14642 .build()
14643 .unwrap(),
14644 );
14645
14646 let first_agent = Arc::clone(&agent);
14647 let first = tokio::spawn(async move { first_agent.transition_to("review").await });
14648 tokio::time::timeout(std::time::Duration::from_secs(2), gate.wait_until_entered())
14649 .await
14650 .expect("reserved transition must enter its exit action");
14651
14652 let second = tokio::time::timeout(
14653 std::time::Duration::from_secs(2),
14654 agent.transition_to("review"),
14655 )
14656 .await
14657 .expect("competing transition must fail without waiting for the exit action")
14658 .unwrap_err();
14659 assert!(second.to_string().contains("already in progress"));
14660
14661 gate.release();
14662 first.await.unwrap().unwrap();
14663 assert_eq!(agent.current_state().as_deref(), Some("review"));
14664 }
14665
14666 #[tokio::test]
14668 async fn concurrent_transition_cannot_overtake_enter_actions() {
14669 let gate = PathMutationGate::new();
14670 let review = ai_agents_state::StateDefinition {
14671 on_enter: vec![StateAction::Tool {
14672 tool: "transition_enter".to_string(),
14673 args: Some(serde_json::json!({"path": "./transition-enter.txt"})),
14674 }],
14675 ..Default::default()
14676 };
14677 let state_machine = Arc::new(
14678 StateMachine::new(ai_agents_state::StateConfig {
14679 initial: "active".to_string(),
14680 states: HashMap::from([
14681 (
14682 "active".to_string(),
14683 ai_agents_state::StateDefinition::default(),
14684 ),
14685 ("review".to_string(), review),
14686 ]),
14687 global_transitions: Vec::new(),
14688 fallback: None,
14689 max_no_transition: None,
14690 regenerate_on_transition: true,
14691 })
14692 .unwrap(),
14693 );
14694 let agent = Arc::new(
14695 AgentBuilder::new()
14696 .system_prompt("Test transition lifecycle reservation.")
14697 .llm(Arc::new(mock_with_response("done")))
14698 .tool(Arc::new(BlockingPathMutationTool {
14699 id: "transition_enter",
14700 path_fields: vec![ai_agents_core::PathPolicyBinding::write("path")],
14701 gate: gate.clone(),
14702 }))
14703 .state_machine(state_machine)
14704 .build()
14705 .unwrap(),
14706 );
14707
14708 let first_agent = Arc::clone(&agent);
14709 let first = tokio::spawn(async move { first_agent.transition_to("review").await });
14710 tokio::time::timeout(std::time::Duration::from_secs(2), gate.wait_until_entered())
14711 .await
14712 .expect("committed transition must enter its destination action");
14713
14714 let second = agent.transition_to("active").await.unwrap_err();
14715 assert!(second.to_string().contains("already in progress"));
14716 assert!(agent.reset().await.is_err());
14717
14718 gate.release();
14719 first.await.unwrap().unwrap();
14720 assert_eq!(agent.current_state().as_deref(), Some("review"));
14721 }
14722
14723 #[tokio::test]
14725 async fn same_state_restore_invalidates_pending_skill_confirmation() {
14726 let (agent, observed) = state_disambiguation_agent_with_skills(
14727 vec![
14728 r#"{"is_ambiguous":false,"confidence":0.99,"ambiguity_type":null,"reasoning":"top-level clear","what_is_unclear":[],"detected_language":"en"}"#,
14729 "send_report",
14730 r#"{"is_ambiguous":true,"confidence":0.2,"ambiguity_type":"missing_target","reasoning":"target missing","what_is_unclear":["target"],"detected_language":"en"}"#,
14731 r#"{"question":"What should I send?","options":null}"#,
14732 r#"{"status":"answered","selected_option":null,"enriched_input":"Send the report to Ada","resolved":{"intent":"send_report"}}"#,
14733 r#"{"question":"Should I send the report to Ada?"}"#,
14734 ],
14735 true,
14736 None,
14737 true,
14738 vec![confirmation_skill()],
14739 );
14740
14741 agent.chat("Send it").await.unwrap();
14742 agent.chat("The report to Ada").await.unwrap();
14743 let snapshot = agent.save_state().await.unwrap();
14744 assert_eq!(agent.current_state().as_deref(), Some("active"));
14745
14746 agent.restore_state(snapshot).await.unwrap();
14747
14748 assert_eq!(agent.current_state().as_deref(), Some("active"));
14749 assert!(agent.pending_skill_id.read().is_none());
14750 assert!(
14751 !agent
14752 .disambiguation_manager()
14753 .unwrap()
14754 .has_pending_clarification()
14755 .await
14756 );
14757 assert_eq!(confirmation_skill_call_count(&observed), 0);
14758 }
14759
14760 #[tokio::test]
14762 async fn direct_state_generation_change_invalidates_confirmation() {
14763 let (agent, observed) = state_disambiguation_agent_with_skills(
14764 vec![
14765 r#"{"is_ambiguous":false,"confidence":0.99,"ambiguity_type":null,"reasoning":"top-level clear","what_is_unclear":[],"detected_language":"en"}"#,
14766 "send_report",
14767 r#"{"is_ambiguous":true,"confidence":0.2,"ambiguity_type":"missing_target","reasoning":"target missing","what_is_unclear":["target"],"detected_language":"en"}"#,
14768 r#"{"question":"What should I send?","options":null}"#,
14769 r#"{"status":"answered","selected_option":null,"enriched_input":"Send the report to Ada","resolved":{"intent":"send_report"}}"#,
14770 r#"{"question":"Should I send the report to Ada?"}"#,
14771 "Confirmation cancelled.",
14772 ],
14773 true,
14774 None,
14775 true,
14776 vec![confirmation_skill()],
14777 );
14778
14779 agent.chat("Send it").await.unwrap();
14780 agent.chat("The report to Ada").await.unwrap();
14781 let state_machine = agent.state_machine().unwrap();
14782 state_machine
14783 .transition_to("review", "external test")
14784 .unwrap();
14785 state_machine
14786 .transition_to("active", "external test")
14787 .unwrap();
14788
14789 let response = agent.chat("Yes").await.unwrap();
14790
14791 assert_eq!(response.content, "Confirmation cancelled.");
14792 assert_eq!(confirmation_skill_call_count(&observed), 0);
14793 assert!(agent.pending_skill_id.read().is_none());
14794 }
14795
14796 #[tokio::test]
14797 async fn state_confirmation_does_not_add_a_question_for_clear_input() {
14798 let (agent, observed) = state_disambiguation_agent(
14799 vec![
14800 r#"{"is_ambiguous":false,"confidence":0.99,"ambiguity_type":null,"reasoning":"clear","what_is_unclear":[],"detected_language":"en"}"#,
14801 "Request executed.",
14802 ],
14803 true,
14804 None,
14805 true,
14806 );
14807
14808 let response = agent.chat("Send the report to Ada").await.unwrap();
14809
14810 assert_eq!(response.content, "Request executed.");
14811 assert_eq!(observed.call_count(), 2);
14812 }
14813
14814 #[tokio::test]
14815 async fn state_override_cannot_activate_a_disabled_top_level_manager() {
14816 let (agent, observed) =
14817 state_disambiguation_agent(vec!["Request executed."], false, Some(true), true);
14818
14819 assert!(!agent.has_disambiguation());
14820 let response = agent.chat("Send it").await.unwrap();
14821
14822 assert_eq!(response.content, "Request executed.");
14823 assert_eq!(observed.call_count(), 1);
14824 }
14825
14826 #[tokio::test]
14827 async fn native_required_choice_executes_through_the_shared_tool_path() {
14828 let mut mock = MockLLMProvider::new("native-required");
14829 mock.set_tool_choice(Some(ToolChoice::Required));
14830 let native_call = ToolCall {
14831 id: "provider-call-1".to_string(),
14832 name: "calculator".to_string(),
14833 arguments: serde_json::json!({"expression": "2 + 2"}),
14834 };
14835 let provider_state = ai_agents_core::NativeProviderState::new(
14836 "fixture-exchange-1",
14837 "fixture",
14838 "native-tools",
14839 ai_agents_core::NativeProviderTarget::new("https://fixture.invalid/", "fixture-model")
14840 .unwrap(),
14841 serde_json::json!({
14842 "role": "model",
14843 "parts": [{
14844 "functionCall": {"name": "calculator", "args": {"expression": "2 + 2"}},
14845 "thoughtSignature": "fixture-signature"
14846 }]
14847 }),
14848 vec![ai_agents_core::NativeCallBinding::new("provider-call-1", 0).unwrap()],
14849 )
14850 .unwrap();
14851 mock.add_response(
14852 LLMResponse::new("", FinishReason::ToolCall)
14853 .with_provider_state(provider_state)
14854 .unwrap()
14855 .with_tool_calls(vec![native_call])
14856 .unwrap(),
14857 );
14858 mock.add_response(LLMResponse::new("The answer is 4.", FinishReason::Stop));
14859 let observed = mock.clone();
14860 let agent = AgentBuilder::new()
14861 .system_prompt("Use the calculator when needed.")
14862 .llm(Arc::new(mock))
14863 .tool(Arc::new(CalculatorTool::new()))
14864 .build()
14865 .unwrap();
14866
14867 let response = agent.chat("What is 2 + 2?").await.unwrap();
14868
14869 assert_eq!(response.content, "The answer is 4.");
14870 assert_eq!(
14871 response.tool_calls.as_ref().unwrap()[0].id,
14872 "provider-call-1"
14873 );
14874 let calls = observed.call_history();
14875 assert_eq!(calls.len(), 2);
14876 assert!(matches!(
14877 calls[0].request.as_ref().map(|request| &request.choice),
14878 Some(ToolChoice::Required)
14879 ));
14880 assert!(matches!(
14881 calls[1].request.as_ref().map(|request| &request.choice),
14882 Some(ToolChoice::Auto)
14883 ));
14884 let replay_batch = calls[1]
14885 .messages
14886 .iter()
14887 .find_map(|message| {
14888 ai_agents_core::decode_native_tool_call_markers(&message.content).unwrap()
14889 })
14890 .expect("signed native call marker must be replayed");
14891 assert_eq!(
14892 replay_batch.provider_state().unwrap().exchange_id(),
14893 "fixture-exchange-1"
14894 );
14895 assert!(calls[1].messages.iter().any(|message| {
14896 ai_agents_core::decode_native_tool_result_markers(&message.content)
14897 .is_ok_and(|results| results.is_some())
14898 }));
14899 }
14900
14901 #[tokio::test]
14902 async fn custom_memory_loss_stops_before_signed_tool_execution() {
14903 let mut mock = MockLLMProvider::new("native-custom-memory");
14904 mock.set_tool_choice(Some(ToolChoice::Required));
14905 let call = ToolCall {
14906 id: "provider-call-drop".to_string(),
14907 name: "calculator".to_string(),
14908 arguments: serde_json::json!({"expression": "3 + 4"}),
14909 };
14910 let state = ai_agents_core::NativeProviderState::new(
14911 "fixture-exchange-drop",
14912 "fixture",
14913 "native-tools",
14914 ai_agents_core::NativeProviderTarget::new("https://fixture.invalid/", "fixture-model")
14915 .unwrap(),
14916 serde_json::json!({
14917 "role": "model",
14918 "parts": [{
14919 "functionCall": {"name": "calculator", "args": {"expression": "3 + 4"}},
14920 "thoughtSignature": "fixture-signature-drop"
14921 }]
14922 }),
14923 vec![ai_agents_core::NativeCallBinding::new("provider-call-drop", 0).unwrap()],
14924 )
14925 .unwrap();
14926 mock.add_response(
14927 LLMResponse::new("", FinishReason::ToolCall)
14928 .with_provider_state(state)
14929 .unwrap()
14930 .with_tool_calls(vec![call])
14931 .unwrap(),
14932 );
14933 let agent = AgentBuilder::new()
14934 .system_prompt("Use the calculator.")
14935 .llm(Arc::new(mock))
14936 .memory(Arc::new(DroppingSignedAssistantMemory {
14937 messages: RwLock::new(Vec::new()),
14938 }))
14939 .tool(Arc::new(CalculatorTool::new()))
14940 .build()
14941 .unwrap();
14942
14943 let error = agent.chat("What is 3 + 4?").await.unwrap_err();
14944
14945 assert!(
14946 error
14947 .to_string()
14948 .contains("removed before provider continuation")
14949 );
14950 assert!(agent.tool_call_history.read().is_empty());
14951 }
14952
14953 #[tokio::test]
14954 async fn sequential_signed_history_validates_every_prior_exchange() {
14955 let mut mock = MockLLMProvider::new("native-sequential-memory");
14956 mock.set_tool_choice(Some(ToolChoice::Required));
14957 mock.add_response(signed_calculator_response(
14958 "seq-exchange-1",
14959 "seq-call-1",
14960 "1 + 1",
14961 ));
14962 mock.add_response(signed_calculator_response(
14963 "seq-exchange-2",
14964 "seq-call-2",
14965 "2 + 2",
14966 ));
14967 let agent = AgentBuilder::new()
14968 .system_prompt("Use the calculator sequentially.")
14969 .llm(Arc::new(mock))
14970 .memory(Arc::new(DroppingEarlierSequentialMemory {
14971 messages: RwLock::new(Vec::new()),
14972 signed_seen: std::sync::atomic::AtomicUsize::new(0),
14973 }))
14974 .tool(Arc::new(CalculatorTool::new()))
14975 .build()
14976 .unwrap();
14977
14978 let error = agent.chat("Calculate twice.").await.unwrap_err();
14979
14980 assert!(error.to_string().contains("seq-exchange-1"));
14981 assert_eq!(agent.tool_call_history.read().len(), 1);
14982 }
14983
14984 #[tokio::test]
14985 async fn post_transition_signed_hitl_rejection_stops_before_continuation() {
14986 let mut native = MockLLMProvider::new("post-transition-native");
14987 native.set_tool_choice(Some(ToolChoice::Auto));
14988 let call = ToolCall {
14989 id: "post-transition-call".to_string(),
14990 name: "echo".to_string(),
14991 arguments: serde_json::json!({"message": "hello"}),
14992 };
14993 let state = ai_agents_core::NativeProviderState::new(
14994 "post-transition-exchange",
14995 "fixture",
14996 "native-tools",
14997 ai_agents_core::NativeProviderTarget::new("https://fixture.invalid/", "fixture-model")
14998 .unwrap(),
14999 serde_json::json!({
15000 "role": "model",
15001 "parts": [{
15002 "functionCall": {"name": "echo", "args": {"message": "hello"}},
15003 "thoughtSignature": "post-transition-signature"
15004 }]
15005 }),
15006 vec![ai_agents_core::NativeCallBinding::new("post-transition-call", 0).unwrap()],
15007 )
15008 .unwrap();
15009 native.add_response(
15010 LLMResponse::new("", FinishReason::ToolCall)
15011 .with_provider_state(state)
15012 .unwrap()
15013 .with_tool_calls(vec![call])
15014 .unwrap(),
15015 );
15016 let observed_native = native.clone();
15017 let yaml = r#"
15018name: PostTransitionNativeReject
15019system_prompt: test
15020tools: [echo]
15021hitl:
15022 tools:
15023 echo:
15024 require_approval: true
15025states:
15026 initial: intake
15027 states:
15028 intake:
15029 prompt: intake
15030 transitions:
15031 - to: active
15032 guard:
15033 context:
15034 route:
15035 eq: active
15036 active:
15037 prompt: active
15038 llm: native
15039"#;
15040 let agent = AgentBuilder::from_yaml(yaml)
15041 .unwrap()
15042 .llm(Arc::new(mock_with_response("stale intake response")))
15043 .llm_alias("native", Arc::new(native))
15044 .auto_configure_features()
15045 .unwrap()
15046 .build()
15047 .unwrap();
15048 agent
15049 .set_context("route", serde_json::json!("active"))
15050 .unwrap();
15051
15052 let error = agent.chat("move to active").await.unwrap_err();
15053
15054 assert!(matches!(error, AgentError::HITLRejected(_)));
15055 assert_eq!(observed_native.call_count(), 1);
15056 }
15057
15058 #[test]
15059 fn runtime_overflow_removes_a_past_signed_user_turn_as_one_prefix() {
15060 let call = ToolCall {
15061 id: "overflow-call".to_string(),
15062 name: "calculator".to_string(),
15063 arguments: serde_json::json!({"expression": "1 + 1"}),
15064 };
15065 let state = ai_agents_core::NativeProviderState::new(
15066 "overflow-exchange",
15067 "google",
15068 "generateContent",
15069 ai_agents_core::NativeProviderTarget::new("https://example.invalid/", "gemini-3")
15070 .unwrap(),
15071 serde_json::json!({
15072 "role": "model",
15073 "parts": [{
15074 "functionCall": {"name": "calculator", "args": {"expression": "1 + 1"}},
15075 "thoughtSignature": "overflow-signature"
15076 }]
15077 }),
15078 vec![ai_agents_core::NativeCallBinding::new("overflow-call", 0).unwrap()],
15079 )
15080 .unwrap();
15081 let call_marker = ai_agents_core::encode_native_tool_call_markers(
15082 std::slice::from_ref(&call),
15083 Some(&state),
15084 )
15085 .unwrap();
15086 let result_marker = ai_agents_core::encode_native_tool_result_marker(
15087 &call,
15088 serde_json::json!({"result": 2}),
15089 )
15090 .unwrap();
15091 let history = vec![
15092 ChatMessage::user("old question"),
15093 ChatMessage::assistant(call_marker),
15094 ChatMessage::function("calculator", result_marker),
15095 ChatMessage::assistant("old answer"),
15096 ChatMessage::user("new question"),
15097 ];
15098
15099 let removable = RuntimeAgent::native_safe_prefix_at_least(&history, 1).unwrap();
15100
15101 assert_eq!(removable, 4);
15102 }
15103
15104 #[test]
15105 fn auxiliary_projection_does_not_interpret_user_marker_text() {
15106 let user_text = serde_json::json!({
15107 "_ai_agents_native_tool_call": true,
15108 "id": "",
15109 "tool": "user-data",
15110 "arguments": {}
15111 })
15112 .to_string();
15113
15114 let projected =
15115 RuntimeAgent::readable_native_messages(vec![ChatMessage::user(&user_text)]).unwrap();
15116
15117 assert_eq!(projected[0].content, user_text);
15118 }
15119
15120 #[tokio::test]
15121 async fn terminal_provider_history_error_skips_retry_and_static_fallback() {
15122 let calls = Arc::new(std::sync::atomic::AtomicU32::new(0));
15123 let recovery = RecoveryManager::new(ai_agents_recovery::ErrorRecoveryConfig {
15124 default: ai_agents_recovery::RetryConfig {
15125 max_retries: 3,
15126 ..Default::default()
15127 },
15128 llm: ai_agents_recovery::LLMRecoveryConfig {
15129 on_failure: LLMFailureAction::FallbackResponse {
15130 message: "must not be returned".to_string(),
15131 },
15132 ..Default::default()
15133 },
15134 ..Default::default()
15135 });
15136 let agent = AgentBuilder::new()
15137 .system_prompt("Reject corrupted native history.")
15138 .llm(Arc::new(TerminalHistoryProvider {
15139 calls: Arc::clone(&calls),
15140 }))
15141 .recovery_manager(recovery)
15142 .build()
15143 .unwrap();
15144
15145 let error = agent.chat("continue").await.unwrap_err();
15146
15147 assert!(
15148 error
15149 .to_string()
15150 .contains("native history integrity failure")
15151 );
15152 assert_eq!(calls.load(Ordering::SeqCst), 1);
15153 }
15154
15155 #[tokio::test]
15156 async fn prompt_fallback_uses_one_corrective_retry() {
15157 let mut mock = MockLLMProvider::new("prompt-required");
15158 mock.set_tool_choice(Some(ToolChoice::Required));
15159 mock.set_native_tool_support(false);
15160 mock.set_responses(
15161 vec![
15162 "I can calculate that.".to_string(),
15163 r#"{"tool":"calculator","arguments":{"expression":"2 + 2"}}"#.to_string(),
15164 "The answer is 4.".to_string(),
15165 ],
15166 false,
15167 );
15168 let observed = mock.clone();
15169 let agent = AgentBuilder::new()
15170 .system_prompt("Use tools.")
15171 .llm(Arc::new(mock))
15172 .tool(Arc::new(CalculatorTool::new()))
15173 .build()
15174 .unwrap();
15175
15176 let response = agent.chat("What is 2 + 2?").await.unwrap();
15177
15178 assert_eq!(response.content, "The answer is 4.");
15179 assert_eq!(observed.call_count(), 3);
15180 let corrective = &observed.call_history()[1].messages;
15181 assert!(
15182 corrective
15183 .last()
15184 .unwrap()
15185 .content
15186 .contains("previous response")
15187 );
15188 }
15189
15190 #[tokio::test]
15191 async fn prompt_fallback_fails_after_one_noncompliant_retry() {
15192 let mut mock = MockLLMProvider::new("prompt-required-failure");
15193 mock.set_tool_choice(Some(ToolChoice::Required));
15194 mock.set_native_tool_support(false);
15195 mock.set_responses(
15196 vec!["No tool.".to_string(), "Still no tool.".to_string()],
15197 false,
15198 );
15199 let observed = mock.clone();
15200 let agent = AgentBuilder::new()
15201 .system_prompt("Use tools.")
15202 .llm(Arc::new(mock))
15203 .tool(Arc::new(CalculatorTool::new()))
15204 .build()
15205 .unwrap();
15206
15207 let error = agent.chat("What is 2 + 2?").await.unwrap_err();
15208
15209 assert!(error.to_string().contains("one corrective retry"));
15210 assert_eq!(observed.call_count(), 2);
15211 }
15212
15213 #[tokio::test]
15214 async fn specific_choice_cannot_widen_the_effective_grant() {
15215 let mut mock = MockLLMProvider::new("specific-outside-grant");
15216 mock.set_tool_choice(Some(ToolChoice::Specific("random".to_string())));
15217 let observed = mock.clone();
15218 let agent = AgentBuilder::new()
15219 .system_prompt("Use tools.")
15220 .llm(Arc::new(mock))
15221 .tool(Arc::new(CalculatorTool::new()))
15222 .build()
15223 .unwrap();
15224
15225 let error = agent.chat("Generate a value.").await.unwrap_err();
15226
15227 assert!(error.to_string().contains("is not registered"));
15228 assert_eq!(observed.call_count(), 0);
15229 }
15230
15231 #[tokio::test]
15232 async fn none_choice_exposes_no_tool_protocol() {
15233 let mut mock = MockLLMProvider::new("no-tools");
15234 mock.set_tool_choice(Some(ToolChoice::None));
15235 mock.set_response(r#"{"tool":"calculator","arguments":{"expression":"2 + 2"}}"#);
15236 let observed = mock.clone();
15237 let agent = AgentBuilder::new()
15238 .system_prompt("Answer directly.")
15239 .llm(Arc::new(mock))
15240 .tool(Arc::new(CalculatorTool::new()))
15241 .build()
15242 .unwrap();
15243
15244 let response = agent.chat("Hello").await.unwrap();
15245
15246 assert!(response.tool_calls.is_none());
15247 assert_eq!(observed.call_count(), 1);
15248 let call = observed.last_call().unwrap();
15249 assert!(call.request.is_none());
15250 assert!(
15251 call.messages
15252 .iter()
15253 .all(|message| !message.content.contains("Available tools:"))
15254 );
15255 }
15256
15257 struct RuntimeStorage {
15258 capabilities: Box<[StorageCapability]>,
15259 snapshots: RwLock<HashMap<String, AgentSnapshot>>,
15260 metadata: RwLock<HashMap<String, ai_agents_core::SessionMetadata>>,
15261 metadata_save_calls: AtomicU64,
15262 metadata_load_calls: AtomicU64,
15263 fail_metadata_save: AtomicBool,
15264 fail_metadata_load: AtomicBool,
15265 }
15266
15267 impl RuntimeStorage {
15268 fn new(capabilities: impl IntoIterator<Item = StorageCapability>) -> Self {
15269 Self {
15270 capabilities: capabilities.into_iter().collect(),
15271 snapshots: RwLock::new(HashMap::new()),
15272 metadata: RwLock::new(HashMap::new()),
15273 metadata_save_calls: AtomicU64::new(0),
15274 metadata_load_calls: AtomicU64::new(0),
15275 fail_metadata_save: AtomicBool::new(false),
15276 fail_metadata_load: AtomicBool::new(false),
15277 }
15278 }
15279 }
15280
15281 #[async_trait]
15282 impl AgentStorage for RuntimeStorage {
15283 fn supports(&self, capability: StorageCapability) -> bool {
15284 self.capabilities.contains(&capability)
15285 }
15286
15287 async fn save(&self, session_id: &str, snapshot: &AgentSnapshot) -> Result<()> {
15288 self.snapshots
15289 .write()
15290 .insert(session_id.to_string(), snapshot.clone());
15291 Ok(())
15292 }
15293
15294 async fn load(&self, session_id: &str) -> Result<Option<AgentSnapshot>> {
15295 Ok(self.snapshots.read().get(session_id).cloned())
15296 }
15297
15298 async fn delete(&self, session_id: &str) -> Result<()> {
15299 self.snapshots.write().remove(session_id);
15300 Ok(())
15301 }
15302
15303 async fn list_sessions(&self) -> Result<Vec<String>> {
15304 Ok(self.snapshots.read().keys().cloned().collect())
15305 }
15306
15307 async fn save_snapshot_with_metadata(
15308 &self,
15309 session_id: &str,
15310 snapshot: &AgentSnapshot,
15311 metadata: &ai_agents_core::SessionMetadata,
15312 ) -> Result<()> {
15313 self.metadata_save_calls.fetch_add(1, Ordering::SeqCst);
15314 if self.fail_metadata_save.load(Ordering::SeqCst) {
15315 return Err(AgentError::Persistence("metadata save failed".into()));
15316 }
15317 self.snapshots
15318 .write()
15319 .insert(session_id.to_string(), snapshot.clone());
15320 self.metadata
15321 .write()
15322 .insert(session_id.to_string(), metadata.clone());
15323 Ok(())
15324 }
15325
15326 async fn save_metadata(
15327 &self,
15328 session_id: &str,
15329 metadata: &ai_agents_core::SessionMetadata,
15330 ) -> Result<()> {
15331 self.metadata_save_calls.fetch_add(1, Ordering::SeqCst);
15332 if self.fail_metadata_save.load(Ordering::SeqCst) {
15333 return Err(AgentError::Persistence("metadata save failed".into()));
15334 }
15335 self.metadata
15336 .write()
15337 .insert(session_id.to_string(), metadata.clone());
15338 Ok(())
15339 }
15340
15341 async fn load_metadata(
15342 &self,
15343 session_id: &str,
15344 ) -> Result<Option<ai_agents_core::SessionMetadata>> {
15345 self.metadata_load_calls.fetch_add(1, Ordering::SeqCst);
15346 if self.fail_metadata_load.load(Ordering::SeqCst) {
15347 return Err(AgentError::Persistence("metadata load failed".into()));
15348 }
15349 Ok(self.metadata.read().get(session_id).cloned())
15350 }
15351 }
15352
15353 fn runtime_storage_agent() -> RuntimeAgent {
15354 AgentBuilder::new()
15355 .system_prompt("Test runtime storage integration.")
15356 .llm(Arc::new(mock_with_response("done")))
15357 .build()
15358 .unwrap()
15359 }
15360
15361 fn restore_spec(id: &str) -> crate::spec::AgentSpec {
15362 crate::spec::AgentSpec {
15363 name: id.to_string(),
15364 system_prompt: format!("Restore child {id}."),
15365 ..crate::spec::AgentSpec::default()
15366 }
15367 }
15368
15369 fn restore_entry(id: &str) -> ai_agents_core::SpawnedAgentEntry {
15370 ai_agents_core::SpawnedAgentEntry {
15371 id: id.to_string(),
15372 name: id.to_string(),
15373 spec_yaml: serde_yaml::to_string(&restore_spec(id)).unwrap(),
15374 }
15375 }
15376
15377 fn restore_spawner(
15378 storage: Arc<RuntimeStorage>,
15379 max_agents: usize,
15380 ) -> (
15381 Arc<crate::spawner::AgentSpawner>,
15382 Arc<crate::spawner::AgentRegistry>,
15383 ) {
15384 let mut llms = LLMRegistry::new();
15385 llms.register("default", Arc::new(mock_with_response("done")));
15386 (
15387 Arc::new(
15388 crate::spawner::AgentSpawner::new()
15389 .with_shared_llms(llms)
15390 .with_shared_storage(storage)
15391 .with_max_agents(max_agents),
15392 ),
15393 Arc::new(crate::spawner::AgentRegistry::new()),
15394 )
15395 }
15396
15397 async fn save_restore_target(
15398 parent: &RuntimeAgent,
15399 storage: &RuntimeStorage,
15400 session_id: &str,
15401 entries: Vec<ai_agents_core::SpawnedAgentEntry>,
15402 ) {
15403 let mut snapshot = parent.save_state().await.unwrap();
15404 snapshot.spawned_agents = Some(entries);
15405 storage.save(session_id, &snapshot).await.unwrap();
15406 storage
15407 .save_metadata(session_id, &ai_agents_core::SessionMetadata::default())
15408 .await
15409 .unwrap();
15410 }
15411
15412 #[tokio::test]
15413 async fn storage_init_requires_storage_for_actor_facts() {
15414 let facts = ai_agents_facts::FactsConfig {
15415 enabled: true,
15416 ..Default::default()
15417 };
15418 let agent = runtime_storage_agent().with_facts_config(None, Some(facts));
15419
15420 let error = agent.init_storage().await.unwrap_err();
15421 assert!(matches!(
15422 error,
15423 AgentError::Config(message)
15424 if message.contains("actor facts or actor memory")
15425 && message.contains("none is configured or injected")
15426 ));
15427 }
15428
15429 #[tokio::test]
15430 async fn storage_init_validates_actor_facts_capability() {
15431 let storage = Arc::new(RuntimeStorage::new([StorageCapability::Snapshot]));
15432 let actor_memory = ai_agents_facts::ActorMemoryConfig {
15433 enabled: true,
15434 ..Default::default()
15435 };
15436 let agent = runtime_storage_agent()
15437 .with_storage(storage)
15438 .with_facts_config(Some(actor_memory), None);
15439
15440 assert!(matches!(
15441 agent.init_storage().await,
15442 Err(AgentError::UnsupportedStorageCapability(
15443 StorageCapability::ActorFacts
15444 ))
15445 ));
15446 }
15447
15448 #[tokio::test]
15449 async fn blocking_chat_rejects_unsupported_required_storage() {
15450 let storage = Arc::new(RuntimeStorage::new([StorageCapability::Snapshot]));
15451 let facts = ai_agents_facts::FactsConfig {
15452 enabled: true,
15453 ..Default::default()
15454 };
15455 let agent = runtime_storage_agent()
15456 .with_storage(storage)
15457 .with_facts_config(None, Some(facts));
15458
15459 assert!(matches!(
15460 agent.chat("hello").await,
15461 Err(AgentError::UnsupportedStorageCapability(
15462 StorageCapability::ActorFacts
15463 ))
15464 ));
15465 }
15466
15467 #[tokio::test]
15468 async fn streaming_chat_rejects_unsupported_required_storage_before_stream_creation() {
15469 let storage = Arc::new(RuntimeStorage::new([StorageCapability::Snapshot]));
15470 let config = ai_agents_relationships::RelationshipConfig {
15471 enabled: true,
15472 ..Default::default()
15473 };
15474 let manager = Arc::new(RelationshipManager::from_config(config).unwrap());
15475 let agent = runtime_storage_agent()
15476 .with_storage(storage)
15477 .with_relationships(manager);
15478
15479 assert!(matches!(
15480 agent.chat_stream("hello").await,
15481 Err(AgentError::UnsupportedStorageCapability(
15482 StorageCapability::ActorRelationships
15483 ))
15484 ));
15485 }
15486
15487 #[tokio::test]
15488 async fn storage_init_completes_facts_for_injected_storage() {
15489 let storage = Arc::new(RuntimeStorage::new([
15490 StorageCapability::Snapshot,
15491 StorageCapability::ActorFacts,
15492 ]));
15493 let facts = ai_agents_facts::FactsConfig {
15494 enabled: true,
15495 ..Default::default()
15496 };
15497 let agent = runtime_storage_agent()
15498 .with_storage(storage)
15499 .with_facts_config(None, Some(facts));
15500
15501 agent.init_storage().await.unwrap();
15502 assert!(agent.fact_store().is_some());
15503 }
15504
15505 #[tokio::test]
15506 async fn storage_init_requires_storage_for_persistent_relationships() {
15507 let config = ai_agents_relationships::RelationshipConfig {
15508 enabled: true,
15509 ..Default::default()
15510 };
15511 let manager = Arc::new(RelationshipManager::from_config(config).unwrap());
15512 let agent = runtime_storage_agent().with_relationships(manager);
15513
15514 let error = agent.init_storage().await.unwrap_err();
15515 assert!(matches!(
15516 error,
15517 AgentError::Config(message)
15518 if message.contains("persistent relationships")
15519 && message.contains("none is configured or injected")
15520 ));
15521 }
15522
15523 #[tokio::test]
15524 async fn storage_init_validates_persistent_relationships_capability() {
15525 let storage = Arc::new(RuntimeStorage::new([StorageCapability::Snapshot]));
15526 let config = ai_agents_relationships::RelationshipConfig {
15527 enabled: true,
15528 ..Default::default()
15529 };
15530 let manager = Arc::new(RelationshipManager::from_config(config).unwrap());
15531 let agent = runtime_storage_agent()
15532 .with_storage(storage)
15533 .with_relationships(manager);
15534
15535 assert!(matches!(
15536 agent.init_storage().await,
15537 Err(AgentError::UnsupportedStorageCapability(
15538 StorageCapability::ActorRelationships
15539 ))
15540 ));
15541 }
15542
15543 #[tokio::test]
15544 async fn session_restore_updates_identity_and_clears_stale_actor_binding() {
15545 let storage = Arc::new(RuntimeStorage::new([
15546 StorageCapability::Snapshot,
15547 StorageCapability::SessionMetadata,
15548 ]));
15549 let agent = runtime_storage_agent().with_storage(storage.clone());
15550 agent.set_actor_id("old-actor").unwrap();
15551 agent.save_session("old").await.unwrap();
15552 storage
15553 .save("target", &agent.save_state().await.unwrap())
15554 .await
15555 .unwrap();
15556 storage
15557 .save_metadata("target", &ai_agents_core::SessionMetadata::default())
15558 .await
15559 .unwrap();
15560
15561 assert!(agent.load_session("target").await.unwrap());
15562
15563 assert_eq!(agent.current_session_id.read().as_deref(), Some("target"));
15564 assert_eq!(agent.actor_id(), None);
15565 }
15566
15567 #[tokio::test]
15568 async fn complete_restore_reconciles_growth_shrink_and_empty_topologies() {
15569 let storage = Arc::new(RuntimeStorage::new([
15570 StorageCapability::Snapshot,
15571 StorageCapability::SessionMetadata,
15572 ]));
15573 let (spawner, registry) = restore_spawner(storage.clone(), 3);
15574 let parent = runtime_storage_agent()
15575 .with_storage(storage.clone())
15576 .with_spawner_handles(Arc::clone(&spawner), Arc::clone(®istry));
15577
15578 for id in ["a", "b"] {
15579 let spawned = spawner
15580 .spawn_with_id(id.to_string(), restore_spec(id))
15581 .await
15582 .unwrap();
15583 spawned.agent.save_session("grow").await.unwrap();
15584 registry.register(spawned).await.unwrap();
15585 }
15586 let staged_c = crate::spawner::storage::NamespacedStorage::new(storage.clone(), "c");
15587 staged_c
15588 .save("grow", &AgentSnapshot::new("c".into()))
15589 .await
15590 .unwrap();
15591 staged_c
15592 .save_metadata("grow", &ai_agents_core::SessionMetadata::default())
15593 .await
15594 .unwrap();
15595 save_restore_target(
15596 &parent,
15597 storage.as_ref(),
15598 "grow",
15599 vec![restore_entry("a"), restore_entry("b"), restore_entry("c")],
15600 )
15601 .await;
15602
15603 assert_eq!(parent.restore_session_full("grow").await.unwrap(), 3);
15604 assert_eq!(registry.count(), 3);
15605 assert!(registry.contains("c"));
15606 assert_eq!(spawner.spawned_count(), 3);
15607
15608 for id in ["a", "b"] {
15609 registry
15610 .get(id)
15611 .unwrap()
15612 .save_session("shrink")
15613 .await
15614 .unwrap();
15615 }
15616 save_restore_target(
15617 &parent,
15618 storage.as_ref(),
15619 "shrink",
15620 vec![restore_entry("a"), restore_entry("b")],
15621 )
15622 .await;
15623
15624 assert_eq!(parent.restore_session_full("shrink").await.unwrap(), 2);
15625 assert_eq!(registry.count(), 2);
15626 assert!(!registry.contains("c"));
15627 assert_eq!(spawner.spawned_count(), 2);
15628
15629 save_restore_target(&parent, storage.as_ref(), "empty", Vec::new()).await;
15630
15631 assert_eq!(parent.restore_session_full("empty").await.unwrap(), 0);
15632 assert_eq!(registry.count(), 0);
15633 assert_eq!(spawner.spawned_count(), 0);
15634 assert_eq!(parent.current_session_id.read().as_deref(), Some("empty"));
15635 }
15636
15637 #[tokio::test]
15638 async fn storage_session_metadata_is_called_only_when_advertised() {
15639 let storage = Arc::new(RuntimeStorage::new([StorageCapability::Snapshot]));
15640 storage.fail_metadata_save.store(true, Ordering::SeqCst);
15641 storage.fail_metadata_load.store(true, Ordering::SeqCst);
15642 let agent = runtime_storage_agent().with_storage(storage.clone());
15643
15644 agent.save_session("session").await.unwrap();
15645 assert!(agent.load_session("session").await.unwrap());
15646 assert_eq!(storage.metadata_save_calls.load(Ordering::SeqCst), 0);
15647 assert_eq!(storage.metadata_load_calls.load(Ordering::SeqCst), 0);
15648 }
15649
15650 #[cfg(feature = "sqlite")]
15651 #[tokio::test]
15652 async fn sqlite_runtime_save_filter_reopen_and_reload_stay_consistent() {
15653 let directory =
15654 std::env::temp_dir().join(format!("ai-agents-runtime-sqlite-{}", uuid::Uuid::new_v4()));
15655 let path = directory.join("sessions.sqlite");
15656 let path_string = path.to_string_lossy().into_owned();
15657 let storage = Arc::new(
15658 ai_agents_storage::SqliteStorage::new(&path_string)
15659 .await
15660 .unwrap(),
15661 );
15662 let agent = runtime_storage_agent().with_storage(storage.clone());
15663 agent.set_session_metadata(ai_agents_core::SessionMetadata {
15664 tags: vec!["initial".into()],
15665 ..Default::default()
15666 });
15667 agent.chat("persist this turn").await.unwrap();
15668 agent.save_session("session").await.unwrap();
15669
15670 agent.set_session_metadata(ai_agents_core::SessionMetadata {
15671 tags: vec!["updated".into()],
15672 ..Default::default()
15673 });
15674 agent.save_session("session").await.unwrap();
15675 assert!(
15676 agent
15677 .list_sessions_filtered(&ai_agents_core::SessionFilter {
15678 tags: Some(vec!["initial".into()]),
15679 ..Default::default()
15680 })
15681 .await
15682 .unwrap()
15683 .is_empty()
15684 );
15685 assert_eq!(
15686 agent
15687 .list_sessions_filtered(&ai_agents_core::SessionFilter {
15688 tags: Some(vec!["updated".into()]),
15689 ..Default::default()
15690 })
15691 .await
15692 .unwrap()
15693 .len(),
15694 1
15695 );
15696 drop(agent);
15697 storage.close().await;
15698 drop(storage);
15699
15700 let reopened_storage = Arc::new(
15701 ai_agents_storage::SqliteStorage::new(&path_string)
15702 .await
15703 .unwrap(),
15704 );
15705 let restored = runtime_storage_agent().with_storage(reopened_storage.clone());
15706 assert!(restored.load_session("session").await.unwrap());
15707 assert_eq!(restored.session_metadata().tags, vec!["updated"]);
15708 assert_eq!(
15709 restored.current_session_id.read().as_deref(),
15710 Some("session")
15711 );
15712 assert!(restored.save_state().await.unwrap().memory.messages.len() >= 2);
15713 assert_eq!(
15714 restored
15715 .list_sessions_filtered(&ai_agents_core::SessionFilter {
15716 tags: Some(vec!["updated".into()]),
15717 ..Default::default()
15718 })
15719 .await
15720 .unwrap()
15721 .len(),
15722 1
15723 );
15724
15725 drop(restored);
15726 reopened_storage.close().await;
15727 drop(reopened_storage);
15728 crate::remove_sqlite_test_directory(&directory)
15729 .await
15730 .unwrap();
15731 }
15732
15733 #[tokio::test]
15734 async fn storage_session_metadata_backend_failures_propagate() {
15735 let storage = Arc::new(RuntimeStorage::new([
15736 StorageCapability::Snapshot,
15737 StorageCapability::SessionMetadata,
15738 ]));
15739 let agent = runtime_storage_agent().with_storage(storage.clone());
15740
15741 agent.save_session("session").await.unwrap();
15742 storage
15743 .save("target", &agent.save_state().await.unwrap())
15744 .await
15745 .unwrap();
15746 storage.fail_metadata_load.store(true, Ordering::SeqCst);
15747 assert!(matches!(
15748 agent.load_session("target").await,
15749 Err(AgentError::Persistence(message)) if message == "metadata load failed"
15750 ));
15751 assert_eq!(agent.current_session_id.read().as_deref(), Some("session"));
15752
15753 storage.fail_metadata_save.store(true, Ordering::SeqCst);
15754 assert!(matches!(
15755 agent.save_session("session").await,
15756 Err(AgentError::Persistence(message)) if message == "metadata save failed"
15757 ));
15758 }
15759
15760 struct ProviderFutureDropSignal {
15761 dropped: Arc<AtomicBool>,
15762 }
15763
15764 impl Drop for ProviderFutureDropSignal {
15765 fn drop(&mut self) {
15766 self.dropped.store(true, Ordering::SeqCst);
15767 }
15768 }
15769
15770 struct BufferedLockingProvider {
15771 lock: Arc<tokio::sync::Mutex<()>>,
15772 stream_started: Arc<tokio::sync::Notify>,
15773 stream_dropped: Arc<AtomicBool>,
15774 committed_after_drop: Arc<AtomicBool>,
15775 }
15776
15777 #[async_trait]
15778 impl LLMProvider for BufferedLockingProvider {
15779 async fn complete(
15780 &self,
15781 _messages: &[ChatMessage],
15782 _config: Option<&LLMConfig>,
15783 ) -> std::result::Result<LLMResponse, LLMError> {
15784 let _guard = self.lock.lock().await;
15785 self.committed_after_drop
15786 .store(self.stream_dropped.load(Ordering::SeqCst), Ordering::SeqCst);
15787 Ok(LLMResponse::new(
15788 "Committed technical response.",
15789 FinishReason::Stop,
15790 ))
15791 }
15792
15793 async fn complete_stream(
15794 &self,
15795 _messages: &[ChatMessage],
15796 _config: Option<&LLMConfig>,
15797 ) -> std::result::Result<
15798 Box<dyn Stream<Item = std::result::Result<LLMChunk, LLMError>> + Unpin + Send>,
15799 LLMError,
15800 > {
15801 let _guard = self.lock.lock().await;
15802 let _drop_signal = ProviderFutureDropSignal {
15803 dropped: Arc::clone(&self.stream_dropped),
15804 };
15805 self.stream_started.notify_one();
15806 std::future::pending().await
15807 }
15808
15809 fn provider_name(&self) -> &str {
15810 "buffered-locking"
15811 }
15812
15813 fn supports(&self, _feature: LLMFeature) -> bool {
15814 false
15815 }
15816 }
15817
15818 struct PendingDropStream {
15819 dropped: Arc<AtomicBool>,
15820 dropped_notify: Arc<tokio::sync::Notify>,
15821 }
15822
15823 impl Stream for PendingDropStream {
15824 type Item = std::result::Result<LLMChunk, LLMError>;
15825
15826 fn poll_next(
15827 self: Pin<&mut Self>,
15828 _cx: &mut std::task::Context<'_>,
15829 ) -> std::task::Poll<Option<Self::Item>> {
15830 std::task::Poll::Pending
15831 }
15832 }
15833
15834 impl Drop for PendingDropStream {
15835 fn drop(&mut self) {
15836 self.dropped.store(true, Ordering::SeqCst);
15837 self.dropped_notify.notify_one();
15838 }
15839 }
15840
15841 struct EstablishedStreamProvider {
15842 stream_started: Arc<tokio::sync::Notify>,
15843 stream_dropped: Arc<AtomicBool>,
15844 stream_dropped_notify: Arc<tokio::sync::Notify>,
15845 committed_after_drop: Arc<AtomicBool>,
15846 }
15847
15848 #[async_trait]
15849 impl LLMProvider for EstablishedStreamProvider {
15850 async fn complete(
15851 &self,
15852 _messages: &[ChatMessage],
15853 _config: Option<&LLMConfig>,
15854 ) -> std::result::Result<LLMResponse, LLMError> {
15855 if !self.stream_dropped.load(Ordering::SeqCst) {
15856 self.stream_dropped_notify.notified().await;
15857 }
15858 self.committed_after_drop
15859 .store(self.stream_dropped.load(Ordering::SeqCst), Ordering::SeqCst);
15860 Ok(LLMResponse::new(
15861 "Committed technical response.",
15862 FinishReason::Stop,
15863 ))
15864 }
15865
15866 async fn complete_stream(
15867 &self,
15868 _messages: &[ChatMessage],
15869 _config: Option<&LLMConfig>,
15870 ) -> std::result::Result<
15871 Box<dyn Stream<Item = std::result::Result<LLMChunk, LLMError>> + Unpin + Send>,
15872 LLMError,
15873 > {
15874 self.stream_started.notify_one();
15875 Ok(Box::new(PendingDropStream {
15876 dropped: Arc::clone(&self.stream_dropped),
15877 dropped_notify: Arc::clone(&self.stream_dropped_notify),
15878 }))
15879 }
15880
15881 fn provider_name(&self) -> &str {
15882 "established-stream"
15883 }
15884
15885 fn supports(&self, _feature: LLMFeature) -> bool {
15886 false
15887 }
15888 }
15889
15890 struct FirstCallLockingProvider {
15891 lock: Arc<tokio::sync::Mutex<()>>,
15892 first_started: Arc<tokio::sync::Notify>,
15893 first_dropped: Arc<AtomicBool>,
15894 committed_after_drop: Arc<AtomicBool>,
15895 calls: AtomicU64,
15896 }
15897
15898 #[async_trait]
15899 impl LLMProvider for FirstCallLockingProvider {
15900 async fn complete(
15901 &self,
15902 _messages: &[ChatMessage],
15903 _config: Option<&LLMConfig>,
15904 ) -> std::result::Result<LLMResponse, LLMError> {
15905 let _guard = self.lock.lock().await;
15906 let call = self.calls.fetch_add(1, Ordering::SeqCst);
15907 if call == 0 {
15908 let _drop_signal = ProviderFutureDropSignal {
15909 dropped: Arc::clone(&self.first_dropped),
15910 };
15911 self.first_started.notify_one();
15912 return std::future::pending().await;
15913 }
15914 self.committed_after_drop
15915 .store(self.first_dropped.load(Ordering::SeqCst), Ordering::SeqCst);
15916 Ok(LLMResponse::new(
15917 "Committed technical response.",
15918 FinishReason::Stop,
15919 ))
15920 }
15921
15922 async fn complete_stream(
15923 &self,
15924 _messages: &[ChatMessage],
15925 _config: Option<&LLMConfig>,
15926 ) -> std::result::Result<
15927 Box<dyn Stream<Item = std::result::Result<LLMChunk, LLMError>> + Unpin + Send>,
15928 LLMError,
15929 > {
15930 Err(LLMError::Other(
15931 "streaming is not used in this test".to_string(),
15932 ))
15933 }
15934
15935 fn provider_name(&self) -> &str {
15936 "first-call-locking"
15937 }
15938
15939 fn supports(&self, _feature: LLMFeature) -> bool {
15940 false
15941 }
15942 }
15943
15944 struct RoutingAfterProviderStart {
15945 provider_started: Arc<tokio::sync::Notify>,
15946 }
15947
15948 #[async_trait]
15949 impl LLMProvider for RoutingAfterProviderStart {
15950 async fn complete(
15951 &self,
15952 _messages: &[ChatMessage],
15953 _config: Option<&LLMConfig>,
15954 ) -> std::result::Result<LLMResponse, LLMError> {
15955 self.provider_started.notified().await;
15956 Ok(LLMResponse::new("1", FinishReason::Stop))
15957 }
15958
15959 async fn complete_stream(
15960 &self,
15961 _messages: &[ChatMessage],
15962 _config: Option<&LLMConfig>,
15963 ) -> std::result::Result<
15964 Box<dyn Stream<Item = std::result::Result<LLMChunk, LLMError>> + Unpin + Send>,
15965 LLMError,
15966 > {
15967 Err(LLMError::Other(
15968 "streaming is not used in this test".to_string(),
15969 ))
15970 }
15971
15972 fn provider_name(&self) -> &str {
15973 "routing-after-start"
15974 }
15975
15976 fn supports(&self, _feature: LLMFeature) -> bool {
15977 false
15978 }
15979 }
15980
15981 struct ResponseCountingHooks {
15983 responses: Arc<std::sync::atomic::AtomicUsize>,
15984 }
15985
15986 struct RootTurnProbeProvider {
15988 complete_entered: tokio::sync::mpsc::UnboundedSender<()>,
15989 }
15990
15991 struct ResponseChatHooks {
15993 target: parking_lot::Mutex<Option<Weak<RuntimeAgent>>>,
15994 invoked: AtomicBool,
15995 nested_result: parking_lot::Mutex<Option<std::result::Result<String, String>>>,
15996 }
15997
15998 struct ConcurrentResponseHooks {
16000 registry: Weak<crate::spawner::AgentRegistry>,
16001 child_id: String,
16002 invoked: AtomicBool,
16003 nested_result: parking_lot::Mutex<Option<std::result::Result<String, String>>>,
16004 }
16005
16006 struct RetryDeadlineTool {
16008 calls: Arc<std::sync::atomic::AtomicUsize>,
16009 deadlines: Arc<parking_lot::Mutex<Vec<chrono::DateTime<chrono::Utc>>>>,
16010 remaining_ms: Arc<parking_lot::Mutex<Vec<i64>>>,
16011 }
16012
16013 struct ToolLifecycleRecordingHooks {
16015 events: parking_lot::Mutex<Vec<String>>,
16016 records: parking_lot::Mutex<Vec<ToolExecutionRecord>>,
16017 }
16018
16019 impl ToolLifecycleRecordingHooks {
16020 fn new() -> Self {
16022 Self {
16023 events: parking_lot::Mutex::new(Vec::new()),
16024 records: parking_lot::Mutex::new(Vec::new()),
16025 }
16026 }
16027
16028 fn events(&self) -> Vec<String> {
16030 self.events.lock().clone()
16031 }
16032
16033 fn records(&self) -> Vec<ToolExecutionRecord> {
16035 self.records.lock().clone()
16036 }
16037 }
16038
16039 struct ContextEchoTool;
16041
16042 #[async_trait]
16043 impl LLMProvider for RootTurnProbeProvider {
16044 async fn complete(
16045 &self,
16046 _messages: &[ChatMessage],
16047 _config: Option<&LLMConfig>,
16048 ) -> std::result::Result<LLMResponse, LLMError> {
16049 let _ = self.complete_entered.send(());
16050 Ok(LLMResponse::new("blocking complete", FinishReason::Stop))
16051 }
16052
16053 async fn complete_stream(
16054 &self,
16055 _messages: &[ChatMessage],
16056 _config: Option<&LLMConfig>,
16057 ) -> std::result::Result<
16058 Box<dyn Stream<Item = std::result::Result<LLMChunk, LLMError>> + Unpin + Send>,
16059 LLMError,
16060 > {
16061 Ok(Box::new(futures::stream::iter(vec![Ok(
16062 LLMChunk::final_chunk("stream complete", FinishReason::Stop, None),
16063 )])))
16064 }
16065
16066 fn provider_name(&self) -> &str {
16067 "root-turn-probe"
16068 }
16069
16070 fn supports(&self, feature: LLMFeature) -> bool {
16071 matches!(feature, LLMFeature::Streaming)
16072 }
16073 }
16074
16075 #[async_trait]
16076 impl ai_agents_core::Tool for ContextEchoTool {
16077 fn id(&self) -> &str {
16078 "context_echo"
16079 }
16080
16081 fn name(&self) -> &str {
16082 "Context Echo"
16083 }
16084
16085 fn description(&self) -> &str {
16086 "Returns selected execution context fields."
16087 }
16088
16089 fn input_schema(&self) -> Value {
16090 serde_json::json!({"type": "object"})
16091 }
16092
16093 fn policy_bindings(&self) -> ai_agents_core::ToolPolicyBindings {
16094 ai_agents_core::ToolPolicyBindings {
16095 path_fields: vec![ai_agents_core::PathPolicyBinding::read("path")],
16096 result_limit_fields: vec![ai_agents_core::ResultLimitBinding::new(
16097 "max_results",
16098 ai_agents_core::ResultLimitKind::MaxResults,
16099 )],
16100 ..Default::default()
16101 }
16102 }
16103
16104 async fn execute(
16105 &self,
16106 _args: Value,
16107 ctx: ai_agents_core::ToolExecutionContext,
16108 ) -> ToolResult {
16109 ToolResult::ok(
16110 serde_json::json!({
16111 "requested_name": ctx.requested_name,
16112 "canonical_id": ctx.canonical_id,
16113 "display_name": ctx.display_name,
16114 "max_results": ctx.limits.max_results,
16115 "custom_config": ctx.custom_config,
16116 })
16117 .to_string(),
16118 )
16119 }
16120 }
16121
16122 #[async_trait]
16123 impl ai_agents_core::Tool for RetryDeadlineTool {
16124 fn id(&self) -> &str {
16125 "retry_deadline"
16126 }
16127
16128 fn name(&self) -> &str {
16129 "Retry Deadline"
16130 }
16131
16132 fn description(&self) -> &str {
16133 "Records one deadline per retry invocation."
16134 }
16135
16136 fn input_schema(&self) -> Value {
16137 serde_json::json!({"type": "object"})
16138 }
16139
16140 fn safety_metadata(&self) -> ai_agents_core::ToolSafetyMetadata {
16141 ai_agents_core::ToolSafetyMetadata::compute()
16142 }
16143
16144 fn classify_call(&self, _args: &Value) -> ai_agents_core::ToolCallClassification {
16145 let mut classification =
16146 ai_agents_core::ToolCallClassification::from_metadata(&self.safety_metadata());
16147 classification.timeout_ms = Some(1_000);
16148 classification.safely_retryable = true;
16149 classification
16150 }
16151
16152 async fn execute(
16154 &self,
16155 _args: Value,
16156 ctx: ai_agents_core::ToolExecutionContext,
16157 ) -> ToolResult {
16158 let deadline = ctx
16159 .deadline
16160 .expect("each invocation must receive a deadline");
16161 self.remaining_ms.lock().push(
16162 deadline
16163 .signed_duration_since(chrono::Utc::now())
16164 .num_milliseconds(),
16165 );
16166 self.deadlines.lock().push(deadline);
16167 let call = self.calls.fetch_add(1, Ordering::SeqCst);
16168 if call == 0 {
16169 tokio::time::sleep(std::time::Duration::from_millis(20)).await;
16170 ToolResult::error("retry")
16171 } else {
16172 ToolResult::ok("done")
16173 }
16174 }
16175 }
16176
16177 struct ClassifiedTimeoutTool {
16179 id: &'static str,
16180 calls: Arc<std::sync::atomic::AtomicUsize>,
16181 timeout_ms: u64,
16182 sleep_ms: u64,
16183 requires_approval: bool,
16184 remaining_ms: Arc<parking_lot::Mutex<Vec<i64>>>,
16185 }
16186
16187 struct ApprovalModifiedTimeoutTool {
16189 calls: Arc<std::sync::atomic::AtomicUsize>,
16190 }
16191
16192 struct SlowTool;
16194
16195 struct FlakyWriteTool {
16197 calls: Arc<std::sync::atomic::AtomicUsize>,
16198 }
16199
16200 struct LockedWriteTool {
16202 active: Arc<std::sync::atomic::AtomicUsize>,
16203 max_active: Arc<std::sync::atomic::AtomicUsize>,
16204 }
16205
16206 struct MultiResourceWriteTool {
16207 active: Arc<std::sync::atomic::AtomicUsize>,
16208 max_active: Arc<std::sync::atomic::AtomicUsize>,
16209 }
16210
16211 #[derive(Clone)]
16212 struct PathMutationGate {
16213 entered: Arc<AtomicBool>,
16214 entered_notify: Arc<tokio::sync::Notify>,
16215 release: Arc<tokio::sync::Notify>,
16216 }
16217
16218 impl PathMutationGate {
16219 fn new() -> Self {
16220 Self {
16221 entered: Arc::new(AtomicBool::new(false)),
16222 entered_notify: Arc::new(tokio::sync::Notify::new()),
16223 release: Arc::new(tokio::sync::Notify::new()),
16224 }
16225 }
16226
16227 async fn wait_until_entered(&self) {
16228 if !self.entered.load(Ordering::SeqCst) {
16229 self.entered_notify.notified().await;
16230 }
16231 }
16232
16233 fn release(&self) {
16234 self.release.notify_one();
16235 }
16236 }
16237
16238 struct BlockingPathMutationTool {
16239 id: &'static str,
16240 path_fields: Vec<ai_agents_core::PathPolicyBinding>,
16241 gate: PathMutationGate,
16242 }
16243
16244 struct NoBindingWriteTool {
16245 active: Arc<std::sync::atomic::AtomicUsize>,
16246 max_active: Arc<std::sync::atomic::AtomicUsize>,
16247 }
16248
16249 struct RecoveryTestTool {
16250 id: String,
16251 succeeds: bool,
16252 calls: Arc<std::sync::atomic::AtomicUsize>,
16253 max_output_chars: Option<usize>,
16254 }
16255
16256 struct BlockingApprovalHandler {
16257 entered: Arc<tokio::sync::Barrier>,
16258 release: Arc<tokio::sync::Notify>,
16259 result: ApprovalResult,
16260 }
16261
16262 struct CountingApprovalHandler {
16263 calls: Arc<std::sync::atomic::AtomicUsize>,
16264 }
16265
16266 struct DriftingFallbackProvider {
16268 refreshed: AtomicBool,
16269 primary_calls: Arc<std::sync::atomic::AtomicUsize>,
16270 secondary_calls: Arc<std::sync::atomic::AtomicUsize>,
16271 }
16272
16273 struct RefreshFallbackProviderHooks {
16275 agent: parking_lot::Mutex<Option<Weak<RuntimeAgent>>>,
16276 lifecycle: Arc<ToolLifecycleRecordingHooks>,
16277 }
16278
16279 struct RuntimeWebFetchTransport {
16280 calls: Arc<std::sync::atomic::AtomicUsize>,
16281 }
16282
16283 struct RuntimeWebFetchResolver;
16284
16285 struct ReentrantToolHooks {
16286 agent: parking_lot::Mutex<Option<Weak<RuntimeAgent>>>,
16287 invoked: AtomicBool,
16288 nested_success: AtomicBool,
16289 }
16290
16291 #[async_trait]
16292 impl ai_agents_core::Tool for ClassifiedTimeoutTool {
16293 fn id(&self) -> &str {
16295 self.id
16296 }
16297
16298 fn name(&self) -> &str {
16300 "Classified Timeout"
16301 }
16302
16303 fn description(&self) -> &str {
16305 "Records and waits under one call-level timeout."
16306 }
16307
16308 fn input_schema(&self) -> Value {
16310 serde_json::json!({"type": "object"})
16311 }
16312
16313 fn classify_call(&self, _args: &Value) -> ai_agents_core::ToolCallClassification {
16315 let mut classification =
16316 ai_agents_core::ToolCallClassification::from_metadata(&self.safety_metadata());
16317 classification.timeout_ms = Some(self.timeout_ms);
16318 classification.requires_approval = self.requires_approval;
16319 classification
16320 }
16321
16322 async fn execute(
16324 &self,
16325 _args: Value,
16326 ctx: ai_agents_core::ToolExecutionContext,
16327 ) -> ToolResult {
16328 self.calls.fetch_add(1, Ordering::SeqCst);
16329 let deadline = ctx
16330 .deadline
16331 .expect("each invocation must receive a deadline");
16332 self.remaining_ms.lock().push(
16333 deadline
16334 .signed_duration_since(chrono::Utc::now())
16335 .num_milliseconds(),
16336 );
16337 tokio::time::sleep(Duration::from_millis(self.sleep_ms)).await;
16338 ToolResult::ok("done")
16339 }
16340 }
16341
16342 #[async_trait]
16343 impl ai_agents_core::Tool for ApprovalModifiedTimeoutTool {
16344 fn id(&self) -> &str {
16346 "approval_modified_timeout"
16347 }
16348
16349 fn name(&self) -> &str {
16351 "Approval Modified Timeout"
16352 }
16353
16354 fn description(&self) -> &str {
16356 "Becomes invalid only after approval modifies its arguments."
16357 }
16358
16359 fn input_schema(&self) -> Value {
16361 serde_json::json!({"type": "object"})
16362 }
16363
16364 fn policy_bindings(&self) -> ai_agents_core::ToolPolicyBindings {
16366 ai_agents_core::ToolPolicyBindings {
16367 path_fields: vec![ai_agents_core::PathPolicyBinding::write("path")],
16368 ..Default::default()
16369 }
16370 }
16371
16372 fn safety_metadata(&self) -> ai_agents_core::ToolSafetyMetadata {
16374 ai_agents_core::ToolSafetyMetadata {
16375 read_only: false,
16376 concurrency_safe: false,
16377 operation: ai_agents_core::ToolOperationKind::Write,
16378 side_effect_level: ai_agents_core::ToolSideEffectLevel::LocalWrite,
16379 requires_network: false,
16380 destructive: false,
16381 open_world: false,
16382 host_dependent: false,
16383 requires_user_interaction: false,
16384 supports_cancellation: true,
16385 default_requires_approval: true,
16386 should_defer_schema: false,
16387 max_output_chars: Some(1024),
16388 max_result_size_chars: Some(1024),
16389 }
16390 }
16391
16392 fn classify_call(&self, args: &Value) -> ai_agents_core::ToolCallClassification {
16394 let mut classification =
16395 ai_agents_core::ToolCallClassification::from_metadata(&self.safety_metadata());
16396 classification.timeout_ms = Some(if args["invalid_timeout"].as_bool() == Some(true) {
16397 u64::MAX
16398 } else {
16399 1_000
16400 });
16401 classification
16402 }
16403
16404 async fn execute(
16406 &self,
16407 _args: Value,
16408 _ctx: ai_agents_core::ToolExecutionContext,
16409 ) -> ToolResult {
16410 self.calls.fetch_add(1, Ordering::SeqCst);
16411 ToolResult::ok("unexpected")
16412 }
16413 }
16414
16415 #[async_trait]
16416 impl ai_agents_core::Tool for SlowTool {
16417 fn id(&self) -> &str {
16418 "slow"
16419 }
16420
16421 fn name(&self) -> &str {
16422 "Slow"
16423 }
16424
16425 fn description(&self) -> &str {
16426 "Waits until cancelled or timed out."
16427 }
16428
16429 fn input_schema(&self) -> Value {
16430 serde_json::json!({"type": "object"})
16431 }
16432
16433 async fn execute(
16434 &self,
16435 _args: Value,
16436 _ctx: ai_agents_core::ToolExecutionContext,
16437 ) -> ToolResult {
16438 tokio::time::sleep(std::time::Duration::from_secs(5)).await;
16439 ToolResult::ok("done")
16440 }
16441 }
16442
16443 #[async_trait]
16444 impl ai_agents_core::Tool for FlakyWriteTool {
16445 fn id(&self) -> &str {
16446 "flaky_write"
16447 }
16448
16449 fn name(&self) -> &str {
16450 "Flaky Write"
16451 }
16452
16453 fn description(&self) -> &str {
16454 "Fails on the first write attempt."
16455 }
16456
16457 fn input_schema(&self) -> Value {
16458 serde_json::json!({"type": "object", "properties": {"path": {"type": "string"}}})
16459 }
16460
16461 fn policy_bindings(&self) -> ai_agents_core::ToolPolicyBindings {
16462 ai_agents_core::ToolPolicyBindings {
16463 path_fields: vec![ai_agents_core::PathPolicyBinding::write("path")],
16464 ..Default::default()
16465 }
16466 }
16467
16468 fn safety_metadata(&self) -> ai_agents_core::ToolSafetyMetadata {
16469 ai_agents_core::ToolSafetyMetadata {
16470 read_only: false,
16471 concurrency_safe: false,
16472 operation: ai_agents_core::ToolOperationKind::Write,
16473 side_effect_level: ai_agents_core::ToolSideEffectLevel::LocalWrite,
16474 requires_network: false,
16475 destructive: false,
16476 open_world: false,
16477 host_dependent: false,
16478 requires_user_interaction: false,
16479 supports_cancellation: true,
16480 default_requires_approval: false,
16481 should_defer_schema: false,
16482 max_output_chars: Some(1024),
16483 max_result_size_chars: Some(1024),
16484 }
16485 }
16486
16487 fn classify_call(&self, _args: &Value) -> ai_agents_core::ToolCallClassification {
16488 let mut classification =
16489 ai_agents_core::ToolCallClassification::from_metadata(&self.safety_metadata());
16490 classification.safely_retryable = false;
16491 classification
16492 }
16493
16494 async fn execute(
16495 &self,
16496 _args: Value,
16497 _ctx: ai_agents_core::ToolExecutionContext,
16498 ) -> ToolResult {
16499 let call = self.calls.fetch_add(1, Ordering::SeqCst);
16500 if call == 0 {
16501 ToolResult::error("first failure")
16502 } else {
16503 ToolResult::ok("second success")
16504 }
16505 }
16506 }
16507
16508 #[async_trait]
16509 impl ai_agents_core::Tool for LockedWriteTool {
16510 fn id(&self) -> &str {
16511 "locked_write"
16512 }
16513
16514 fn name(&self) -> &str {
16515 "Locked Write"
16516 }
16517
16518 fn description(&self) -> &str {
16519 "Tracks concurrent execution on one resource."
16520 }
16521
16522 fn input_schema(&self) -> Value {
16523 serde_json::json!({"type": "object", "properties": {"path": {"type": "string"}}})
16524 }
16525
16526 fn policy_bindings(&self) -> ai_agents_core::ToolPolicyBindings {
16527 ai_agents_core::ToolPolicyBindings {
16528 path_fields: vec![ai_agents_core::PathPolicyBinding::write("path")],
16529 ..Default::default()
16530 }
16531 }
16532
16533 fn safety_metadata(&self) -> ai_agents_core::ToolSafetyMetadata {
16534 ai_agents_core::ToolSafetyMetadata {
16535 read_only: false,
16536 concurrency_safe: false,
16537 operation: ai_agents_core::ToolOperationKind::Write,
16538 side_effect_level: ai_agents_core::ToolSideEffectLevel::LocalWrite,
16539 requires_network: false,
16540 destructive: false,
16541 open_world: false,
16542 host_dependent: false,
16543 requires_user_interaction: false,
16544 supports_cancellation: true,
16545 default_requires_approval: false,
16546 should_defer_schema: false,
16547 max_output_chars: Some(1024),
16548 max_result_size_chars: Some(1024),
16549 }
16550 }
16551
16552 async fn execute(
16553 &self,
16554 _args: Value,
16555 _ctx: ai_agents_core::ToolExecutionContext,
16556 ) -> ToolResult {
16557 let active = self.active.fetch_add(1, Ordering::SeqCst) + 1;
16558 loop {
16559 let current_max = self.max_active.load(Ordering::SeqCst);
16560 if active <= current_max {
16561 break;
16562 }
16563 if self
16564 .max_active
16565 .compare_exchange(current_max, active, Ordering::SeqCst, Ordering::SeqCst)
16566 .is_ok()
16567 {
16568 break;
16569 }
16570 }
16571 tokio::time::sleep(std::time::Duration::from_millis(100)).await;
16572 self.active.fetch_sub(1, Ordering::SeqCst);
16573 ToolResult::ok("done")
16574 }
16575 }
16576
16577 #[async_trait]
16578 impl ai_agents_core::Tool for MultiResourceWriteTool {
16579 fn id(&self) -> &str {
16580 "multi_resource_write"
16581 }
16582
16583 fn name(&self) -> &str {
16584 "Multi Resource Write"
16585 }
16586
16587 fn description(&self) -> &str {
16588 "Tracks concurrent execution across source and destination resources."
16589 }
16590
16591 fn input_schema(&self) -> Value {
16592 serde_json::json!({"type": "object"})
16593 }
16594
16595 fn policy_bindings(&self) -> ai_agents_core::ToolPolicyBindings {
16596 ai_agents_core::ToolPolicyBindings {
16597 path_fields: vec![
16598 ai_agents_core::PathPolicyBinding::read_write("source_path"),
16599 ai_agents_core::PathPolicyBinding::write("destination_path"),
16600 ],
16601 ..Default::default()
16602 }
16603 }
16604
16605 fn safety_metadata(&self) -> ai_agents_core::ToolSafetyMetadata {
16606 LockedWriteTool {
16607 active: Arc::clone(&self.active),
16608 max_active: Arc::clone(&self.max_active),
16609 }
16610 .safety_metadata()
16611 }
16612
16613 async fn execute(
16614 &self,
16615 _args: Value,
16616 _ctx: ai_agents_core::ToolExecutionContext,
16617 ) -> ToolResult {
16618 let active = self.active.fetch_add(1, Ordering::SeqCst) + 1;
16619 self.max_active.fetch_max(active, Ordering::SeqCst);
16620 tokio::time::sleep(std::time::Duration::from_millis(75)).await;
16621 self.active.fetch_sub(1, Ordering::SeqCst);
16622 ToolResult::ok("done")
16623 }
16624 }
16625
16626 #[async_trait]
16627 impl ai_agents_core::Tool for BlockingPathMutationTool {
16628 fn id(&self) -> &str {
16629 self.id
16630 }
16631
16632 fn name(&self) -> &str {
16633 self.id
16634 }
16635
16636 fn description(&self) -> &str {
16637 "Blocks a path mutation until the test releases it."
16638 }
16639
16640 fn input_schema(&self) -> Value {
16641 serde_json::json!({"type": "object"})
16642 }
16643
16644 fn policy_bindings(&self) -> ai_agents_core::ToolPolicyBindings {
16645 ai_agents_core::ToolPolicyBindings {
16646 path_fields: self.path_fields.clone(),
16647 ..Default::default()
16648 }
16649 }
16650
16651 fn safety_metadata(&self) -> ai_agents_core::ToolSafetyMetadata {
16652 ai_agents_core::ToolSafetyMetadata {
16653 read_only: false,
16654 concurrency_safe: false,
16655 operation: ai_agents_core::ToolOperationKind::Write,
16656 side_effect_level: ai_agents_core::ToolSideEffectLevel::LocalWrite,
16657 requires_network: false,
16658 destructive: false,
16659 open_world: false,
16660 host_dependent: false,
16661 requires_user_interaction: false,
16662 supports_cancellation: true,
16663 default_requires_approval: false,
16664 should_defer_schema: false,
16665 max_output_chars: Some(1024),
16666 max_result_size_chars: Some(1024),
16667 }
16668 }
16669
16670 async fn execute(
16671 &self,
16672 _args: Value,
16673 _ctx: ai_agents_core::ToolExecutionContext,
16674 ) -> ToolResult {
16675 self.gate.entered.store(true, Ordering::SeqCst);
16676 self.gate.entered_notify.notify_one();
16677 self.gate.release.notified().await;
16678 ToolResult::ok("done")
16679 }
16680 }
16681
16682 #[async_trait]
16683 impl ai_agents_core::Tool for NoBindingWriteTool {
16684 fn id(&self) -> &str {
16685 "no_binding_write"
16686 }
16687
16688 fn name(&self) -> &str {
16689 "No Binding Write"
16690 }
16691
16692 fn description(&self) -> &str {
16693 "Tracks concurrent execution without resource bindings."
16694 }
16695
16696 fn input_schema(&self) -> Value {
16697 serde_json::json!({"type": "object"})
16698 }
16699
16700 fn safety_metadata(&self) -> ai_agents_core::ToolSafetyMetadata {
16701 LockedWriteTool {
16702 active: Arc::clone(&self.active),
16703 max_active: Arc::clone(&self.max_active),
16704 }
16705 .safety_metadata()
16706 }
16707
16708 async fn execute(
16709 &self,
16710 _args: Value,
16711 _ctx: ai_agents_core::ToolExecutionContext,
16712 ) -> ToolResult {
16713 let active = self.active.fetch_add(1, Ordering::SeqCst) + 1;
16714 self.max_active.fetch_max(active, Ordering::SeqCst);
16715 tokio::time::sleep(std::time::Duration::from_millis(75)).await;
16716 self.active.fetch_sub(1, Ordering::SeqCst);
16717 ToolResult::ok("done")
16718 }
16719 }
16720
16721 #[async_trait]
16722 impl ai_agents_core::Tool for RecoveryTestTool {
16723 fn id(&self) -> &str {
16724 &self.id
16725 }
16726
16727 fn name(&self) -> &str {
16728 &self.id
16729 }
16730
16731 fn description(&self) -> &str {
16732 "Records recovery execution and returns a configured result."
16733 }
16734
16735 fn input_schema(&self) -> Value {
16736 serde_json::json!({"type": "object"})
16737 }
16738
16739 fn policy_bindings(&self) -> ai_agents_core::ToolPolicyBindings {
16740 ai_agents_core::ToolPolicyBindings {
16741 path_fields: vec![ai_agents_core::PathPolicyBinding::write("path")],
16742 ..Default::default()
16743 }
16744 }
16745
16746 fn safety_metadata(&self) -> ai_agents_core::ToolSafetyMetadata {
16748 ai_agents_core::ToolSafetyMetadata {
16749 read_only: false,
16750 concurrency_safe: false,
16751 operation: ai_agents_core::ToolOperationKind::Write,
16752 side_effect_level: ai_agents_core::ToolSideEffectLevel::LocalWrite,
16753 requires_network: false,
16754 destructive: false,
16755 open_world: false,
16756 host_dependent: false,
16757 requires_user_interaction: false,
16758 supports_cancellation: true,
16759 default_requires_approval: false,
16760 should_defer_schema: false,
16761 max_output_chars: Some(self.max_output_chars.unwrap_or(1024)),
16762 max_result_size_chars: Some(1024),
16763 }
16764 }
16765
16766 async fn execute(
16768 &self,
16769 _args: Value,
16770 _ctx: ai_agents_core::ToolExecutionContext,
16771 ) -> ToolResult {
16772 self.calls.fetch_add(1, Ordering::SeqCst);
16773 let mut result = if self.succeeds {
16774 ToolResult::ok(format!("{} succeeded", self.id))
16775 } else {
16776 ToolResult::error(format!("{} failed", self.id))
16777 };
16778 result.metadata = Some(HashMap::from([(
16779 "recovery_test_tool".to_string(),
16780 Value::String(self.id.clone()),
16781 )]));
16782 result
16783 }
16784 }
16785
16786 #[async_trait]
16787 impl WebFetchTransport for RuntimeWebFetchTransport {
16788 async fn send(
16790 &self,
16791 _request: WebFetchTransportRequest,
16792 ) -> std::result::Result<WebFetchTransportResponse, String> {
16793 Err("validated addresses are required".to_string())
16794 }
16795
16796 async fn send_validated(
16798 &self,
16799 _request: WebFetchTransportRequest,
16800 _addresses: &[std::net::SocketAddr],
16801 ) -> std::result::Result<WebFetchTransportResponse, String> {
16802 self.calls.fetch_add(1, Ordering::SeqCst);
16803 Ok(WebFetchTransportResponse {
16804 status: 200,
16805 content_type: Some("text/plain".to_string()),
16806 location: None,
16807 body: b"approved".to_vec(),
16808 })
16809 }
16810 }
16811
16812 #[async_trait]
16813 impl WebFetchResolver for RuntimeWebFetchResolver {
16814 async fn resolve(
16816 &self,
16817 _host: &str,
16818 _port: u16,
16819 ) -> std::result::Result<Vec<std::net::IpAddr>, String> {
16820 Ok(vec![std::net::IpAddr::V4(std::net::Ipv4Addr::new(
16821 93, 184, 216, 34,
16822 ))])
16823 }
16824 }
16825
16826 #[async_trait]
16827 impl ToolProvider for DriftingFallbackProvider {
16828 fn id(&self) -> &str {
16830 "drifting_fallback"
16831 }
16832
16833 fn name(&self) -> &str {
16835 "Drifting Fallback"
16836 }
16837
16838 fn provider_type(&self) -> ToolProviderType {
16840 ToolProviderType::Custom
16841 }
16842
16843 async fn list_tools(&self) -> Vec<ToolDescriptor> {
16845 let alias = ToolAliases::new().with_name("en", "fallback alias");
16846 let mut primary = ToolDescriptor::new(
16847 "primary",
16848 "Primary",
16849 "Fails before fallback.",
16850 serde_json::json!({"type": "object"}),
16851 );
16852 let mut secondary = ToolDescriptor::new(
16853 "secondary",
16854 "Secondary",
16855 "Must not execute after final canonical drift.",
16856 serde_json::json!({"type": "object"}),
16857 );
16858 if self.refreshed.load(Ordering::SeqCst) {
16859 primary = primary.with_aliases(alias);
16860 } else {
16861 secondary = secondary.with_aliases(alias);
16862 }
16863 vec![primary, secondary]
16864 }
16865
16866 async fn get_tool(&self, tool_id: &str) -> Option<Arc<dyn Tool>> {
16868 let calls = match tool_id {
16869 "primary" => Arc::clone(&self.primary_calls),
16870 "secondary" => Arc::clone(&self.secondary_calls),
16871 _ => return None,
16872 };
16873 Some(Arc::new(RecoveryTestTool {
16874 id: tool_id.to_string(),
16875 succeeds: false,
16876 calls,
16877 max_output_chars: None,
16878 }))
16879 }
16880
16881 fn supports_refresh(&self) -> bool {
16883 true
16884 }
16885
16886 async fn refresh(&self) -> std::result::Result<(), ToolProviderError> {
16888 self.refreshed.store(true, Ordering::SeqCst);
16889 Ok(())
16890 }
16891 }
16892
16893 #[async_trait]
16894 impl AgentHooks for RefreshFallbackProviderHooks {
16895 async fn on_tool_start(&self, tool: &str, args: &Value) {
16897 self.lifecycle.on_tool_start(tool, args).await;
16898 if tool != "secondary" {
16899 return;
16900 }
16901 let agent = self.agent.lock().as_ref().and_then(Weak::upgrade);
16902 if let Some(agent) = agent {
16903 agent
16904 .tools
16905 .refresh_provider("drifting_fallback")
16906 .await
16907 .unwrap();
16908 }
16909 }
16910
16911 async fn on_tool_complete(&self, tool: &str, result: &ToolResult, duration_ms: u64) {
16912 self.lifecycle
16913 .on_tool_complete(tool, result, duration_ms)
16914 .await;
16915 }
16916
16917 async fn on_tool_execution_record(&self, record: &ToolExecutionRecord) {
16918 self.lifecycle.on_tool_execution_record(record).await;
16919 }
16920
16921 async fn on_error(&self, error: &AgentError) {
16922 self.lifecycle.on_error(error).await;
16923 }
16924 }
16925
16926 #[async_trait]
16927 impl ApprovalHandler for BlockingApprovalHandler {
16928 async fn request_approval(
16929 &self,
16930 _request: ai_agents_hitl::ApprovalRequest,
16931 ) -> ApprovalResult {
16932 self.entered.wait().await;
16933 self.release.notified().await;
16934 self.result.clone()
16935 }
16936 }
16937
16938 #[async_trait]
16939 impl ApprovalHandler for CountingApprovalHandler {
16940 async fn request_approval(
16941 &self,
16942 _request: ai_agents_hitl::ApprovalRequest,
16943 ) -> ApprovalResult {
16944 self.calls.fetch_add(1, Ordering::SeqCst);
16945 ApprovalResult::Approved
16946 }
16947 }
16948
16949 #[async_trait]
16950 impl AgentHooks for ReentrantToolHooks {
16951 async fn on_tool_complete(&self, tool: &str, _result: &ToolResult, _duration_ms: u64) {
16952 if tool != "reentrant_write" || self.invoked.swap(true, Ordering::SeqCst) {
16953 return;
16954 }
16955 let agent = self.agent.lock().as_ref().and_then(Weak::upgrade);
16956 if let Some(agent) = agent {
16957 let result = agent
16958 .invoke_tool(ToolExecutionRequest::new(
16959 "nested-hook-call",
16960 "reentrant_write",
16961 serde_json::json!({"path": "./hook.txt"}),
16962 ToolCallSource::Manual,
16963 ))
16964 .await;
16965 self.nested_success
16966 .store(result.is_ok_and(|record| record.success), Ordering::SeqCst);
16967 }
16968 }
16969 }
16970
16971 #[async_trait]
16972 impl AgentHooks for ResponseCountingHooks {
16973 async fn on_response(&self, _response: &AgentResponse) {
16974 self.responses.fetch_add(1, Ordering::SeqCst);
16975 }
16976 }
16977
16978 #[async_trait]
16979 impl AgentHooks for ResponseChatHooks {
16980 async fn on_response(&self, _response: &AgentResponse) {
16982 if self.invoked.swap(true, Ordering::SeqCst) {
16983 return;
16984 }
16985 let target = self.target.lock().as_ref().and_then(Weak::upgrade);
16986 let result = if let Some(target) = target {
16987 target
16988 .chat("nested response hook call")
16989 .await
16990 .map(|response| response.content)
16991 .map_err(|error| error.to_string())
16992 } else {
16993 Err("response hook target is unavailable".to_string())
16994 };
16995 *self.nested_result.lock() = Some(result);
16996 }
16997 }
16998
16999 #[async_trait]
17000 impl AgentHooks for ConcurrentResponseHooks {
17001 async fn on_response(&self, _response: &AgentResponse) {
17003 if self.invoked.swap(true, Ordering::SeqCst) {
17004 return;
17005 }
17006 let Some(registry) = self.registry.upgrade() else {
17007 *self.nested_result.lock() =
17008 Some(Err("concurrent registry is unavailable".to_string()));
17009 return;
17010 };
17011 let agents = [ai_agents_state::ConcurrentAgentRef::Id(
17012 self.child_id.clone(),
17013 )];
17014 let aggregation = ai_agents_state::AggregationConfig {
17015 strategy: ai_agents_state::AggregationStrategy::FirstWins,
17016 synthesizer_llm: None,
17017 synthesizer_prompt: None,
17018 vote: None,
17019 };
17020 let result = crate::orchestration::concurrent(
17021 ®istry,
17022 "nested concurrent response hook call",
17023 &agents,
17024 &aggregation,
17025 None,
17026 Some(1),
17027 None,
17028 ai_agents_state::PartialFailureAction::Abort,
17029 None,
17030 )
17031 .await
17032 .map(|result| result.response.content)
17033 .map_err(|error| error.to_string());
17034 *self.nested_result.lock() = Some(result);
17035 }
17036 }
17037
17038 #[async_trait]
17039 impl AgentHooks for ToolLifecycleRecordingHooks {
17040 async fn on_tool_start(&self, tool: &str, _args: &Value) {
17041 self.events.lock().push(format!("start:{tool}"));
17042 }
17043
17044 async fn on_tool_complete(&self, tool: &str, result: &ToolResult, _duration_ms: u64) {
17045 self.events
17046 .lock()
17047 .push(format!("complete:{tool}:{}", result.success));
17048 }
17049
17050 async fn on_tool_execution_record(&self, record: &ToolExecutionRecord) {
17051 self.events.lock().push(format!(
17052 "record:{}:{}",
17053 record.canonical_id, record.executed
17054 ));
17055 self.records.lock().push(record.clone());
17056 }
17057
17058 async fn on_error(&self, _error: &AgentError) {
17060 self.events.lock().push("error".to_string());
17061 }
17062 }
17063
17064 struct ApprovalRecordingHooks {
17065 events: parking_lot::Mutex<Vec<String>>,
17066 }
17067
17068 impl ApprovalRecordingHooks {
17069 fn new() -> Self {
17070 Self {
17071 events: parking_lot::Mutex::new(Vec::new()),
17072 }
17073 }
17074
17075 fn events(&self) -> Vec<String> {
17076 self.events.lock().clone()
17077 }
17078 }
17079
17080 #[async_trait]
17081 impl AgentHooks for ApprovalRecordingHooks {
17082 async fn on_approval_result(&self, request_id: &str, result: &ApprovalResult) {
17083 self.events.lock().push(format!(
17084 "raw:{}:{}",
17085 request_id,
17086 approval_result_name(result)
17087 ));
17088 }
17089
17090 async fn on_approval_resolved(
17091 &self,
17092 request: &ai_agents_hitl::ApprovalRequest,
17093 raw_result: &ApprovalResult,
17094 outcome: &ApprovalResolvedOutcome,
17095 ) {
17096 self.events.lock().push(format!(
17097 "resolved:{}:{}:{}",
17098 request.id,
17099 approval_result_name(raw_result),
17100 approval_outcome_name(outcome)
17101 ));
17102 }
17103 }
17104
17105 fn approval_result_name(result: &ApprovalResult) -> &'static str {
17106 match result {
17107 ApprovalResult::Approved => "approved",
17108 ApprovalResult::Rejected { .. } => "rejected",
17109 ApprovalResult::Modified { .. } => "modified",
17110 ApprovalResult::Timeout => "timeout",
17111 }
17112 }
17113
17114 fn approval_outcome_name(outcome: &ApprovalResolvedOutcome) -> &'static str {
17115 match outcome {
17116 ApprovalResolvedOutcome::Approved => "approved",
17117 ApprovalResolvedOutcome::Rejected { .. } => "rejected",
17118 ApprovalResolvedOutcome::Modified { .. } => "modified",
17119 ApprovalResolvedOutcome::Error { .. } => "error",
17120 }
17121 }
17122
17123 fn assert_correlated_approval_events(
17124 events: &[String],
17125 raw_status: &str,
17126 outcome_status: &str,
17127 ) {
17128 assert_eq!(events.len(), 2);
17129 let raw: Vec<_> = events[0].split(':').collect();
17130 let resolved: Vec<_> = events[1].split(':').collect();
17131 assert_eq!(raw[0], "raw");
17132 assert_eq!(resolved[0], "resolved");
17133 assert_eq!(raw[1], resolved[1]);
17134 assert_eq!(raw[2], raw_status);
17135 assert_eq!(resolved[2], raw_status);
17136 assert_eq!(resolved[3], outcome_status);
17137 }
17138
17139 fn approval_security_config(policy_enabled: bool) -> ToolSecurityConfig {
17140 let mut security = ToolSecurityConfig {
17141 enabled: true,
17142 fail_closed: true,
17143 ..Default::default()
17144 };
17145 let policy = ai_agents_tools::ToolPolicyConfig {
17146 enabled: policy_enabled,
17147 write_paths: vec![".".to_string()],
17148 require_confirmation: true,
17149 ..Default::default()
17150 };
17151 security.tools.insert("locked_write".to_string(), policy);
17152 security
17153 }
17154
17155 struct MutationTestWorkspace {
17156 root: std::path::PathBuf,
17157 }
17158
17159 impl MutationTestWorkspace {
17160 fn new() -> Self {
17161 let root = std::env::temp_dir().join(format!(
17162 "ai-agents-runtime-mutation-{}",
17163 uuid::Uuid::new_v4()
17164 ));
17165 std::fs::create_dir_all(&root).unwrap();
17166 Self { root }
17167 }
17168 }
17169
17170 impl Drop for MutationTestWorkspace {
17171 fn drop(&mut self) {
17172 let _ = std::fs::remove_dir_all(&self.root);
17173 }
17174 }
17175
17176 async fn wait_for_resource_lock_strong_count(locks: &ToolResourceLocks, minimum: usize) {
17177 tokio::time::timeout(std::time::Duration::from_secs(2), async {
17178 loop {
17179 let strong_count = locks
17180 .read()
17181 .get("path-mutation:global")
17182 .map_or(0, |lock| lock.strong_count());
17183 if strong_count >= minimum {
17184 break;
17185 }
17186 tokio::task::yield_now().await;
17187 }
17188 })
17189 .await
17190 .expect("path mutation call did not reach the shared lock");
17191 }
17192
17193 async fn assert_path_mutation_pair_serialized(
17194 first_id: &'static str,
17195 first_fields: Vec<ai_agents_core::PathPolicyBinding>,
17196 first_args: Value,
17197 second_id: &'static str,
17198 second_fields: Vec<ai_agents_core::PathPolicyBinding>,
17199 second_args: Value,
17200 ) {
17201 let locks = new_tool_resource_locks();
17202 let first_gate = PathMutationGate::new();
17203 let second_gate = PathMutationGate::new();
17204 second_gate.release();
17205 let agent = Arc::new(
17206 AgentBuilder::new()
17207 .system_prompt("Test global path mutation locking.")
17208 .llm(Arc::new(mock_with_response("done")))
17209 .tool(Arc::new(BlockingPathMutationTool {
17210 id: first_id,
17211 path_fields: first_fields,
17212 gate: first_gate.clone(),
17213 }))
17214 .tool(Arc::new(BlockingPathMutationTool {
17215 id: second_id,
17216 path_fields: second_fields,
17217 gate: second_gate.clone(),
17218 }))
17219 .build()
17220 .unwrap()
17221 .with_shared_resource_locks(Arc::clone(&locks)),
17222 );
17223
17224 let first = {
17225 let agent = Arc::clone(&agent);
17226 tokio::spawn(async move {
17227 agent
17228 .invoke_tool(ToolExecutionRequest::new(
17229 format!("{}-first", first_id),
17230 first_id,
17231 first_args,
17232 ToolCallSource::Manual,
17233 ))
17234 .await
17235 .unwrap()
17236 })
17237 };
17238 first_gate.wait_until_entered().await;
17239
17240 let second = {
17241 let agent = Arc::clone(&agent);
17242 tokio::spawn(async move {
17243 agent
17244 .invoke_tool(ToolExecutionRequest::new(
17245 format!("{}-second", second_id),
17246 second_id,
17247 second_args,
17248 ToolCallSource::Manual,
17249 ))
17250 .await
17251 .unwrap()
17252 })
17253 };
17254 wait_for_resource_lock_strong_count(&locks, 2).await;
17255 assert!(!second_gate.entered.load(Ordering::SeqCst));
17256 assert!(!second.is_finished());
17257
17258 first_gate.release();
17259 let (first, second) = tokio::time::timeout(std::time::Duration::from_secs(2), async {
17260 tokio::join!(first, second)
17261 })
17262 .await
17263 .expect("serialized path mutation calls did not finish");
17264 assert!(first.unwrap().success);
17265 assert!(second.unwrap().success);
17266 assert!(second_gate.entered.load(Ordering::SeqCst));
17267 assert!(locks.read().is_empty());
17268 }
17269
17270 #[derive(Clone, Copy)]
17271 enum MutationDenial {
17272 Policy,
17273 Approval,
17274 }
17275
17276 fn mutation_denial_security_config(
17277 tool_id: &str,
17278 workspace: &std::path::Path,
17279 denial: MutationDenial,
17280 ) -> ToolSecurityConfig {
17281 let workspace = workspace.to_string_lossy().into_owned();
17282 let mut policy = ai_agents_tools::ToolPolicyConfig {
17283 read_paths: vec![workspace.clone()],
17284 write_paths: vec![workspace.clone()],
17285 ..Default::default()
17286 };
17287 match denial {
17288 MutationDenial::Policy => policy.blocked_paths = vec![workspace],
17289 MutationDenial::Approval => policy.require_confirmation = true,
17290 }
17291
17292 let mut security = ToolSecurityConfig {
17293 enabled: true,
17294 fail_closed: true,
17295 ..Default::default()
17296 };
17297 security.tools.insert(tool_id.to_string(), policy);
17298 security
17299 }
17300
17301 async fn assert_path_mutation_denied(tool: Arc<dyn Tool>, denial: MutationDenial) {
17302 let workspace = MutationTestWorkspace::new();
17303 let tool_id = tool.id().to_string();
17304 let preserved = workspace.root.join(format!("{}-preserved.txt", tool_id));
17305 let destination = workspace.root.join(format!("{}-destination.txt", tool_id));
17306 std::fs::write(&preserved, "preserved").unwrap();
17307 let arguments = match tool_id.as_str() {
17308 "copy_path" | "move_path" => serde_json::json!({
17309 "source_path": preserved.to_string_lossy(),
17310 "destination_path": destination.to_string_lossy(),
17311 "dry_run": false
17312 }),
17313 "delete_path" => serde_json::json!({
17314 "path": preserved.to_string_lossy(),
17315 "recursive": false,
17316 "dry_run": false
17317 }),
17318 _ => panic!("unsupported mutation tool: {}", tool_id),
17319 };
17320 let security = mutation_denial_security_config(&tool_id, &workspace.root, denial);
17321 let builder = AgentBuilder::new()
17322 .system_prompt("Test mutation denial.")
17323 .llm(Arc::new(mock_with_response("done")))
17324 .tool(tool)
17325 .tool_security(ToolSecurityEngine::new(security));
17326 let builder = match denial {
17327 MutationDenial::Policy => builder,
17328 MutationDenial::Approval => builder
17329 .hitl_engine(HITLEngine::new(ai_agents_hitl::HITLConfig::default()))
17330 .approval_handler(Arc::new(RejectAllHandler::new())),
17331 };
17332 let agent = builder.build().unwrap();
17333
17334 let record = agent
17335 .invoke_tool(ToolExecutionRequest::new(
17336 format!("{}-denied", tool_id),
17337 tool_id.clone(),
17338 arguments,
17339 ToolCallSource::Manual,
17340 ))
17341 .await
17342 .unwrap();
17343
17344 assert!(!record.executed, "{} must not be invoked", tool_id);
17345 assert!(!record.success);
17346 match denial {
17347 MutationDenial::Policy => {
17348 assert_eq!(record.policy.outcome, PermissionOutcome::Deny);
17349 assert!(record.approval.as_ref().is_some_and(|approval| matches!(
17350 &approval.status,
17351 ToolApprovalStatus::NotRequired
17352 )));
17353 }
17354 MutationDenial::Approval => {
17355 assert_eq!(record.policy.outcome, PermissionOutcome::RequiresApproval);
17356 assert!(record.approval.as_ref().is_some_and(|approval| matches!(
17357 &approval.status,
17358 ToolApprovalStatus::Rejected
17359 )));
17360 }
17361 }
17362 assert_eq!(std::fs::read_to_string(&preserved).unwrap(), "preserved");
17363 assert!(!destination.exists());
17364 }
17365
17366 fn recovery_manager_with_fallbacks(
17367 fallbacks: impl IntoIterator<Item = (String, String)>,
17368 ) -> RecoveryManager {
17369 use ai_agents_recovery::{ErrorRecoveryConfig, ToolRecoveryConfig, ToolRetryConfig};
17370
17371 let per_tool = fallbacks
17372 .into_iter()
17373 .map(|(tool, fallback_tool)| {
17374 (
17375 tool,
17376 ToolRetryConfig {
17377 max_retries: 0,
17378 timeout_ms: Some(1_000),
17379 on_failure: ToolFailureAction::Fallback { fallback_tool },
17380 },
17381 )
17382 })
17383 .collect();
17384 RecoveryManager::new(ErrorRecoveryConfig {
17385 tools: ToolRecoveryConfig {
17386 per_tool,
17387 ..Default::default()
17388 },
17389 ..Default::default()
17390 })
17391 }
17392
17393 fn approval_check() -> HITLCheckResult {
17394 HITLCheckResult::required(
17395 ApprovalTrigger::tool("test", serde_json::json!({})),
17396 HashMap::new(),
17397 "Approve?",
17398 None,
17399 )
17400 }
17401
17402 fn agent_with_approval_result(
17403 raw_result: ApprovalResult,
17404 timeout_action: TimeoutAction,
17405 hooks: Arc<ApprovalRecordingHooks>,
17406 ) -> RuntimeAgent {
17407 use ai_agents_hitl::{CallbackHandler, HITLConfig};
17408
17409 let config = HITLConfig {
17410 on_timeout: timeout_action,
17411 ..Default::default()
17412 };
17413 let handler = CallbackHandler::new(move |_| raw_result.clone());
17414 AgentBuilder::new()
17415 .system_prompt("Test HITL hooks.")
17416 .llm(Arc::new(mock_with_response("done")))
17417 .build()
17418 .unwrap()
17419 .with_hooks(hooks)
17420 .with_hitl(HITLEngine::new(config), Arc::new(handler))
17421 }
17422
17423 #[tokio::test]
17424 async fn approval_hooks_expose_direct_effective_decisions_after_raw_results() {
17425 let cases = vec![
17426 (ApprovalResult::Approved, "approved"),
17427 (
17428 ApprovalResult::Rejected {
17429 reason: Some("denied".to_string()),
17430 },
17431 "rejected",
17432 ),
17433 (
17434 ApprovalResult::Modified {
17435 changes: HashMap::from([("value".to_string(), serde_json::json!(2))]),
17436 },
17437 "modified",
17438 ),
17439 ];
17440
17441 for (raw_result, expected) in cases {
17442 let hooks = Arc::new(ApprovalRecordingHooks::new());
17443 let agent =
17444 agent_with_approval_result(raw_result, TimeoutAction::Reject, hooks.clone());
17445
17446 let result = agent.request_hitl_approval(approval_check()).await.unwrap();
17447
17448 assert_eq!(approval_result_name(&result), expected);
17449 assert_correlated_approval_events(&hooks.events(), expected, expected);
17450 }
17451 }
17452
17453 #[tokio::test]
17454 async fn approval_hooks_expose_timeout_policy_decisions() {
17455 for (timeout_action, expected) in [
17456 (TimeoutAction::Approve, "approved"),
17457 (TimeoutAction::Reject, "rejected"),
17458 ] {
17459 let hooks = Arc::new(ApprovalRecordingHooks::new());
17460 let agent =
17461 agent_with_approval_result(ApprovalResult::Timeout, timeout_action, hooks.clone());
17462
17463 let result = agent.request_hitl_approval(approval_check()).await.unwrap();
17464
17465 assert_eq!(approval_result_name(&result), expected);
17466 assert_correlated_approval_events(&hooks.events(), "timeout", expected);
17467 }
17468 }
17469
17470 #[tokio::test]
17471 async fn timeout_error_fires_correlated_resolved_error_before_returning() {
17472 let hooks = Arc::new(ApprovalRecordingHooks::new());
17473 let agent = agent_with_approval_result(
17474 ApprovalResult::Timeout,
17475 TimeoutAction::Error,
17476 hooks.clone(),
17477 );
17478
17479 let error = agent
17480 .request_hitl_approval(approval_check())
17481 .await
17482 .unwrap_err();
17483
17484 assert!(error.to_string().contains("HITL approval timeout"));
17485 assert_correlated_approval_events(&hooks.events(), "timeout", "error");
17486 }
17487
17488 #[tokio::test]
17490 async fn test_integration_yaml_to_chat_basic() {
17491 let mock = mock_with_response("Hello! How can I help you?");
17492 let agent = AgentBuilder::new()
17493 .system_prompt("You are a test assistant.")
17494 .llm(Arc::new(mock))
17495 .build()
17496 .unwrap();
17497
17498 let response = agent.chat("Hi").await.unwrap();
17499 assert!(!response.content.is_empty());
17500 assert_eq!(response.content, "Hello! How can I help you?");
17501 }
17502
17503 #[tokio::test]
17504 async fn stream_events_emit_one_authoritative_final_without_legacy_done() {
17505 let agent = AgentBuilder::new()
17506 .system_prompt("You are a test assistant.")
17507 .llm(Arc::new(mock_with_response(
17508 "Hello from the final response.",
17509 )))
17510 .build()
17511 .unwrap();
17512
17513 let mut stream = agent.chat_stream_events("Hi").await.unwrap();
17514 let mut final_responses = Vec::new();
17515 let mut legacy_done = 0;
17516 while let Some(event) = stream.next().await {
17517 match event {
17518 AgentStreamEvent::Chunk(StreamChunk::Done {}) => legacy_done += 1,
17519 AgentStreamEvent::Chunk(StreamChunk::Error { message }) => {
17520 panic!("unexpected stream error: {message}")
17521 }
17522 AgentStreamEvent::Final(response) => final_responses.push(response),
17523 AgentStreamEvent::Chunk(_) => {}
17524 }
17525 }
17526
17527 assert_eq!(legacy_done, 0);
17528 assert_eq!(final_responses.len(), 1);
17529 let response = final_responses.pop().unwrap();
17530 assert_eq!(response.content, "Hello from the final response.");
17531 assert!(
17532 response
17533 .metadata
17534 .as_ref()
17535 .is_some_and(|metadata| { metadata.contains_key("reasoning") })
17536 );
17537 }
17538
17539 #[tokio::test]
17540 async fn stream_final_content_includes_output_processing_after_provisional_chunks() {
17541 let yaml = r#"
17542name: ProcessedStreamAgent
17543system_prompt: "Answer directly."
17544process:
17545 output:
17546 - type: format
17547 config:
17548 template: "{{ response }} [finalized]"
17549streaming:
17550 enabled: true
17551"#;
17552 let agent = AgentBuilder::from_yaml(yaml)
17553 .unwrap()
17554 .llm(Arc::new(mock_with_response("provisional answer")))
17555 .auto_configure_features()
17556 .unwrap()
17557 .build()
17558 .unwrap();
17559
17560 let mut stream = agent.chat_stream_events("Hi").await.unwrap();
17561 let mut provisional = String::new();
17562 let mut final_content = None;
17563 while let Some(event) = stream.next().await {
17564 match event {
17565 AgentStreamEvent::Chunk(StreamChunk::Content { text }) => {
17566 provisional.push_str(&text)
17567 }
17568 AgentStreamEvent::Chunk(StreamChunk::Error { message }) => {
17569 panic!("unexpected stream error: {message}")
17570 }
17571 AgentStreamEvent::Final(response) => final_content = Some(response.content),
17572 AgentStreamEvent::Chunk(_) => {}
17573 }
17574 }
17575
17576 assert_eq!(provisional, "provisional answer");
17577 assert_eq!(
17578 final_content.as_deref(),
17579 Some("provisional answer [finalized]")
17580 );
17581 }
17582
17583 #[tokio::test]
17584 async fn stream_events_preserve_tool_progress_and_final_tool_calls() {
17585 let agent = AgentBuilder::new()
17586 .system_prompt("Use the echo tool once, then answer.")
17587 .llm(Arc::new(mock_with_responses(vec![
17588 r#"{"tool":"echo","arguments":{"message":"hello"}}"#,
17589 "Echo completed.",
17590 ])))
17591 .tool(Arc::new(ai_agents_tools::EchoTool::new()))
17592 .build()
17593 .unwrap();
17594
17595 let mut stream = agent.chat_stream_events("echo hello").await.unwrap();
17596 let mut starts = 0;
17597 let mut results = 0;
17598 let mut ends = 0;
17599 let mut final_response = None;
17600 while let Some(event) = stream.next().await {
17601 match event {
17602 AgentStreamEvent::Chunk(StreamChunk::ToolCallStart { name, .. }) => {
17603 assert_eq!(name, "echo");
17604 starts += 1;
17605 }
17606 AgentStreamEvent::Chunk(StreamChunk::ToolResult { name, success, .. }) => {
17607 assert_eq!(name, "echo");
17608 assert!(success);
17609 results += 1;
17610 }
17611 AgentStreamEvent::Chunk(StreamChunk::ToolCallEnd { .. }) => ends += 1,
17612 AgentStreamEvent::Chunk(StreamChunk::Error { message }) => {
17613 panic!("unexpected stream error: {message}")
17614 }
17615 AgentStreamEvent::Final(response) => final_response = Some(response),
17616 AgentStreamEvent::Chunk(_) => {}
17617 }
17618 }
17619
17620 assert_eq!((starts, results, ends), (1, 1, 1));
17621 let response = final_response.expect("tool stream must finalize");
17622 assert_eq!(response.content, "Echo completed.");
17623 assert_eq!(
17624 response.tool_calls.as_ref().map(|calls| calls
17625 .iter()
17626 .map(|call| call.name.as_str())
17627 .collect::<Vec<_>>()),
17628 Some(vec!["echo"])
17629 );
17630 }
17631
17632 #[tokio::test]
17633 async fn legacy_stream_still_emits_one_done_chunk() {
17634 let agent = AgentBuilder::new()
17635 .system_prompt("You are a test assistant.")
17636 .llm(Arc::new(mock_with_response(
17637 "Hello from the legacy stream.",
17638 )))
17639 .build()
17640 .unwrap();
17641
17642 let mut stream = agent.chat_stream("Hi").await.unwrap();
17643 let mut done = 0;
17644 while let Some(chunk) = stream.next().await {
17645 match chunk {
17646 StreamChunk::Done {} => done += 1,
17647 StreamChunk::Error { message } => panic!("unexpected stream error: {message}"),
17648 _ => {}
17649 }
17650 }
17651
17652 assert_eq!(done, 1);
17653 }
17654
17655 #[tokio::test]
17657 async fn test_integration_multi_turn_conversation() {
17658 let mock = mock_with_responses(vec![
17659 "Hello! I'm your assistant.",
17660 "The weather is sunny today.",
17661 "Goodbye!",
17662 ]);
17663 let agent = AgentBuilder::new()
17664 .system_prompt("You are helpful.")
17665 .llm(Arc::new(mock))
17666 .build()
17667 .unwrap();
17668
17669 let r1 = agent.chat("Hi").await.unwrap();
17670 assert_eq!(r1.content, "Hello! I'm your assistant.");
17671
17672 let r2 = agent.chat("What's the weather?").await.unwrap();
17673 assert_eq!(r2.content, "The weather is sunny today.");
17674
17675 let r3 = agent.chat("Bye").await.unwrap();
17676 assert_eq!(r3.content, "Goodbye!");
17677
17678 let messages = agent.memory.get_messages(None).await.unwrap();
17680 assert_eq!(messages.len(), 6);
17682 }
17683
17684 #[test]
17685 fn later_approval_preserves_modified_evidence() {
17686 let arguments = serde_json::json!({"dry_run": true});
17687 let mut record = Some(ToolApprovalRecord {
17688 status: ToolApprovalStatus::Modified,
17689 reason: None,
17690 modified_arguments: Some(arguments.clone()),
17691 });
17692
17693 merge_approved_record(&mut record);
17694
17695 let record = record.unwrap();
17696 assert!(matches!(record.status, ToolApprovalStatus::Modified));
17697 assert_eq!(record.modified_arguments, Some(arguments));
17698 }
17699
17700 #[test]
17701 fn approval_binding_rejects_replaced_tool_implementation() {
17702 let reviewed_tool: Arc<dyn ai_agents_core::Tool> = Arc::new(ContextEchoTool);
17703 let same_tool = Arc::clone(&reviewed_tool);
17704 let replacement_tool: Arc<dyn ai_agents_core::Tool> = Arc::new(ContextEchoTool);
17705 let arguments = serde_json::json!({"path": "."});
17706 let versions = ToolDecisionVersions {
17707 policy: 2,
17708 registry: 3,
17709 runtime_control: 4,
17710 state: Some(5),
17711 };
17712 let binding = ToolApprovalBinding {
17713 canonical_id: "context_echo".to_string(),
17714 arguments: arguments.clone(),
17715 confirmation_required: true,
17716 policy_version: versions.policy,
17717 runtime_control_version: versions.runtime_control,
17718 state_generation: versions.state,
17719 reviewed_tool,
17720 };
17721
17722 assert!(!binding.is_stale("context_echo", &arguments, true, versions, &same_tool,));
17723 assert!(binding.is_stale(
17724 "context_echo",
17725 &arguments,
17726 true,
17727 versions,
17728 &replacement_tool,
17729 ));
17730 }
17731
17732 #[tokio::test]
17733 async fn approved_mutation_to_dry_run_remains_executable() {
17734 use ai_agents_hitl::CallbackHandler;
17735
17736 let handler = CallbackHandler::new(|_| ApprovalResult::Modified {
17737 changes: HashMap::from([("dry_run".to_string(), serde_json::json!(true))]),
17738 });
17739 let agent = AgentBuilder::new()
17740 .system_prompt("Test safer approval modifications.")
17741 .llm(Arc::new(mock_with_response("done")))
17742 .tool(Arc::new(ai_agents_tools::FileWriteTool::new()))
17743 .hitl_engine(HITLEngine::new(ai_agents_hitl::HITLConfig::default()))
17744 .approval_handler(Arc::new(handler))
17745 .build()
17746 .unwrap();
17747
17748 let record = agent
17749 .invoke_tool(ToolExecutionRequest::new(
17750 "approved-dry-run",
17751 "file_write",
17752 serde_json::json!({
17753 "path": "./approval-dry-run.txt",
17754 "content": "not written"
17755 }),
17756 ToolCallSource::Manual,
17757 ))
17758 .await
17759 .unwrap();
17760
17761 assert!(record.executed);
17762 assert!(record.success);
17763 assert_eq!(record.executed_arguments["dry_run"], true);
17764 assert!(matches!(
17765 record.approval.as_ref().map(|approval| &approval.status),
17766 Some(ToolApprovalStatus::Modified)
17767 ));
17768 let output: Value = serde_json::from_str(&record.output).unwrap();
17769 assert_eq!(output["mutation_performed"], false);
17770 }
17771
17772 #[tokio::test]
17774 async fn shared_executor_approval_reaches_web_fetch_transport() {
17775 use ai_agents_hitl::{CallbackHandler, HITLConfig};
17776 use ai_agents_tools::{DomainPolicyConfig, ToolPolicyConfig};
17777
17778 let calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
17779 let tool = WebFetchTool::with_transport_and_resolver(
17780 Arc::new(RuntimeWebFetchTransport {
17781 calls: Arc::clone(&calls),
17782 }),
17783 Arc::new(RuntimeWebFetchResolver),
17784 );
17785 let mut security = ToolSecurityConfig {
17786 enabled: true,
17787 fail_closed: true,
17788 ..Default::default()
17789 };
17790 security.tools.insert(
17791 "web_fetch".to_string(),
17792 ToolPolicyConfig {
17793 domains: DomainPolicyConfig {
17794 requires_approval: vec!["approval.test".to_string()],
17795 ..Default::default()
17796 },
17797 allowed_schemes: vec!["https".to_string()],
17798 allowed_ports: vec![443],
17799 ..Default::default()
17800 },
17801 );
17802 let handler = CallbackHandler::new(|_| ApprovalResult::Approved);
17803 let agent = AgentBuilder::new()
17804 .system_prompt("Test approved web fetch execution.")
17805 .llm(Arc::new(mock_with_response("done")))
17806 .tool(Arc::new(tool))
17807 .tool_security(ToolSecurityEngine::new(security))
17808 .build()
17809 .unwrap()
17810 .with_hitl(HITLEngine::new(HITLConfig::default()), Arc::new(handler));
17811
17812 let record = agent
17813 .invoke_tool(ToolExecutionRequest::new(
17814 "approved-web-fetch",
17815 "web_fetch",
17816 serde_json::json!({
17817 "url": "https://approval.test/page",
17818 "cache_ttl_seconds": 0
17819 }),
17820 ToolCallSource::Manual,
17821 ))
17822 .await
17823 .unwrap();
17824
17825 assert!(record.success);
17826 assert!(
17827 record
17828 .approval
17829 .as_ref()
17830 .is_some_and(|approval| matches!(approval.status, ToolApprovalStatus::Approved))
17831 );
17832 assert_eq!(calls.load(Ordering::SeqCst), 1);
17833 }
17834
17835 #[tokio::test]
17836 async fn context_preserves_requested_and_canonical_identity() {
17837 let mock = mock_with_response("hello");
17838 let mut tools = ai_agents_tools::ToolRegistry::new();
17839 tools.register(Arc::new(ContextEchoTool)).unwrap();
17840
17841 let mut security = ToolSecurityConfig {
17842 enabled: true,
17843 fail_closed: true,
17844 ..Default::default()
17845 };
17846 let mut policy = ai_agents_tools::ToolPolicyConfig {
17847 read_paths: vec![".".to_string()],
17848 max_results: Some(7),
17849 ..Default::default()
17850 };
17851 policy
17852 .config
17853 .insert("backend".to_string(), serde_json::json!("memory"));
17854 security.tools.insert("context_echo".to_string(), policy);
17855
17856 let agent = AgentBuilder::new()
17857 .system_prompt("You are helpful.")
17858 .llm(Arc::new(mock))
17859 .tools(tools)
17860 .tool_security(ToolSecurityEngine::new(security))
17861 .build()
17862 .unwrap();
17863
17864 let record = agent
17865 .invoke_tool(ToolExecutionRequest::new(
17866 "ctx-call",
17867 "Context Echo",
17868 serde_json::json!({"path": ".", "max_results": 99}),
17869 ToolCallSource::Manual,
17870 ))
17871 .await
17872 .unwrap();
17873
17874 assert!(record.success);
17875 assert!(matches!(&record.source, ToolCallSource::Manual));
17876 assert_eq!(record.requested_name, "Context Echo");
17877 assert_eq!(record.canonical_id, "context_echo");
17878 assert_eq!(record.policy.outcome, PermissionOutcome::Allow);
17879 assert_eq!(record.executed_arguments["max_results"], 7);
17880 let output: Value = serde_json::from_str(&record.output).unwrap();
17881 assert_eq!(output["requested_name"], "Context Echo");
17882 assert_eq!(output["canonical_id"], "context_echo");
17883 assert_eq!(output["max_results"], 7);
17884 assert_eq!(output["custom_config"]["backend"], "memory");
17885 assert!(record.metadata.contains_key("effective_limits"));
17886 assert!(record.metadata.contains_key("policy_snapshot"));
17887 }
17888
17889 #[tokio::test]
17890 async fn test_runtime_control_cancels_active_tool_call() {
17891 let mock = mock_with_response("hello");
17892 let agent = Arc::new(
17893 AgentBuilder::new()
17894 .system_prompt("You are helpful.")
17895 .llm(Arc::new(mock))
17896 .tool(Arc::new(SlowTool))
17897 .build()
17898 .unwrap(),
17899 );
17900 let control = agent.runtime_control();
17901 let running_agent = Arc::clone(&agent);
17902 let handle = tokio::spawn(async move {
17903 running_agent
17904 .invoke_tool(ToolExecutionRequest::new(
17905 "slow-call",
17906 "slow",
17907 serde_json::json!({}),
17908 ToolCallSource::Manual,
17909 ))
17910 .await
17911 .unwrap()
17912 });
17913
17914 tokio::time::sleep(std::time::Duration::from_millis(100)).await;
17915 control.cancel_all();
17916 let record = handle.await.unwrap();
17917
17918 assert!(record.executed);
17919 assert!(record.cancelled);
17920 assert!(!record.success);
17921 assert!(record.cancellation_reason.is_some());
17922 }
17923
17924 #[tokio::test]
17926 async fn cancelled_tool_does_not_enter_fallback() {
17927 let fallback_calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
17928 let agent = Arc::new(
17929 AgentBuilder::new()
17930 .system_prompt("Test cancellation before fallback.")
17931 .llm(Arc::new(mock_with_response("done")))
17932 .tool(Arc::new(SlowTool))
17933 .tool(Arc::new(RecoveryTestTool {
17934 id: "fallback".to_string(),
17935 succeeds: true,
17936 calls: Arc::clone(&fallback_calls),
17937 max_output_chars: None,
17938 }))
17939 .recovery_manager(recovery_manager_with_fallbacks([(
17940 "slow".to_string(),
17941 "fallback".to_string(),
17942 )]))
17943 .build()
17944 .unwrap(),
17945 );
17946 let control = agent.runtime_control();
17947 let running_agent = Arc::clone(&agent);
17948 let handle = tokio::spawn(async move {
17949 running_agent
17950 .invoke_tool(ToolExecutionRequest::new(
17951 "cancelled-fallback-call",
17952 "slow",
17953 serde_json::json!({}),
17954 ToolCallSource::Manual,
17955 ))
17956 .await
17957 .unwrap()
17958 });
17959
17960 tokio::time::sleep(Duration::from_millis(100)).await;
17961 control.cancel_all();
17962 let record = handle.await.unwrap();
17963
17964 assert!(record.executed);
17965 assert!(record.cancelled);
17966 assert!(!record.success);
17967 assert_eq!(record.canonical_id, "slow");
17968 assert_eq!(fallback_calls.load(Ordering::SeqCst), 0);
17969 assert_eq!(agent.tool_call_history().len(), 1);
17970 }
17971
17972 #[tokio::test]
17973 async fn non_idempotent_tool_calls_are_not_retried() {
17974 use ai_agents_recovery::{ErrorRecoveryConfig, ToolRecoveryConfig, ToolRetryConfig};
17975
17976 let mock = mock_with_response("hello");
17977 let calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
17978 let agent = AgentBuilder::new()
17979 .system_prompt("You are helpful.")
17980 .llm(Arc::new(mock))
17981 .tool(Arc::new(FlakyWriteTool {
17982 calls: Arc::clone(&calls),
17983 }))
17984 .recovery_manager(RecoveryManager::new(ErrorRecoveryConfig {
17985 tools: ToolRecoveryConfig {
17986 default: ToolRetryConfig {
17987 max_retries: 2,
17988 ..Default::default()
17989 },
17990 ..Default::default()
17991 },
17992 ..Default::default()
17993 }))
17994 .build()
17995 .unwrap();
17996
17997 let record = agent
17998 .invoke_tool(ToolExecutionRequest::new(
17999 "flaky-call",
18000 "flaky_write",
18001 serde_json::json!({"path": "./tmp.txt"}),
18002 ToolCallSource::Manual,
18003 ))
18004 .await
18005 .unwrap();
18006
18007 assert!(!record.success);
18008 assert_eq!(calls.load(Ordering::SeqCst), 1);
18009 }
18010
18011 #[tokio::test]
18012 async fn safely_retryable_tool_receives_a_fresh_deadline_per_attempt() {
18013 use ai_agents_recovery::{ErrorRecoveryConfig, ToolRecoveryConfig, ToolRetryConfig};
18014
18015 let calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
18016 let deadlines = Arc::new(parking_lot::Mutex::new(Vec::new()));
18017 let remaining_ms = Arc::new(parking_lot::Mutex::new(Vec::new()));
18018 let agent = AgentBuilder::new()
18019 .system_prompt("Test retry deadlines.")
18020 .llm(Arc::new(mock_with_response("done")))
18021 .tool(Arc::new(RetryDeadlineTool {
18022 calls: Arc::clone(&calls),
18023 deadlines: Arc::clone(&deadlines),
18024 remaining_ms: Arc::clone(&remaining_ms),
18025 }))
18026 .recovery_manager(RecoveryManager::new(ErrorRecoveryConfig {
18027 tools: ToolRecoveryConfig {
18028 per_tool: HashMap::from([(
18029 "retry_deadline".to_string(),
18030 ToolRetryConfig {
18031 max_retries: 1,
18032 ..Default::default()
18033 },
18034 )]),
18035 ..Default::default()
18036 },
18037 ..Default::default()
18038 }))
18039 .build()
18040 .unwrap();
18041
18042 let record = agent
18043 .invoke_tool(ToolExecutionRequest::new(
18044 "retry-deadline-call",
18045 "retry_deadline",
18046 serde_json::json!({}),
18047 ToolCallSource::Manual,
18048 ))
18049 .await
18050 .unwrap();
18051
18052 assert!(record.executed);
18053 assert!(record.success);
18054 assert_eq!(calls.load(Ordering::SeqCst), 2);
18055 let deadlines = deadlines.lock();
18056 assert_eq!(deadlines.len(), 2);
18057 assert!(
18058 deadlines[1] > deadlines[0],
18059 "retry inherited the first invocation deadline"
18060 );
18061 let remaining_ms = remaining_ms.lock();
18062 assert_eq!(remaining_ms.len(), 2);
18063 assert!(
18064 remaining_ms
18065 .iter()
18066 .all(|remaining| (800..=1_000).contains(remaining))
18067 );
18068 }
18069
18070 #[tokio::test]
18072 async fn call_classification_timeout_controls_deadline_and_timer() {
18073 let calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
18074 let remaining_ms = Arc::new(parking_lot::Mutex::new(Vec::new()));
18075 let agent = AgentBuilder::new()
18076 .system_prompt("Test call-level timeout.")
18077 .llm(Arc::new(mock_with_response("done")))
18078 .tool(Arc::new(ClassifiedTimeoutTool {
18079 id: "classified_timeout",
18080 calls: Arc::clone(&calls),
18081 timeout_ms: 100,
18082 sleep_ms: 150,
18083 requires_approval: false,
18084 remaining_ms: Arc::clone(&remaining_ms),
18085 }))
18086 .build()
18087 .unwrap();
18088
18089 let started = Instant::now();
18090 let record = agent
18091 .invoke_tool(ToolExecutionRequest::new(
18092 "classified-timeout-call",
18093 "classified_timeout",
18094 serde_json::json!({}),
18095 ToolCallSource::Manual,
18096 ))
18097 .await
18098 .unwrap();
18099
18100 assert!(record.executed);
18101 assert!(record.timed_out);
18102 assert!(!record.success);
18103 assert_eq!(calls.load(Ordering::SeqCst), 1);
18104 assert!(started.elapsed() < Duration::from_secs(1));
18105 let remaining_ms = remaining_ms.lock();
18106 assert_eq!(remaining_ms.len(), 1);
18107 assert!((1..=100).contains(&remaining_ms[0]));
18108 }
18109
18110 #[tokio::test]
18112 async fn recovery_timeout_only_lowers_call_and_policy_timeouts() {
18113 use ai_agents_recovery::{ErrorRecoveryConfig, ToolRecoveryConfig, ToolRetryConfig};
18114
18115 let calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
18116 let remaining_ms = Arc::new(parking_lot::Mutex::new(Vec::new()));
18117 let agent = AgentBuilder::new()
18118 .system_prompt("Test recovery timeout.")
18119 .llm(Arc::new(mock_with_response("done")))
18120 .tool(Arc::new(ClassifiedTimeoutTool {
18121 id: "recovery_timeout",
18122 calls: Arc::clone(&calls),
18123 timeout_ms: 1_000,
18124 sleep_ms: 150,
18125 requires_approval: false,
18126 remaining_ms: Arc::clone(&remaining_ms),
18127 }))
18128 .recovery_manager(RecoveryManager::new(ErrorRecoveryConfig {
18129 tools: ToolRecoveryConfig {
18130 per_tool: HashMap::from([(
18131 "recovery_timeout".to_string(),
18132 ToolRetryConfig {
18133 timeout_ms: Some(100),
18134 ..Default::default()
18135 },
18136 )]),
18137 ..Default::default()
18138 },
18139 ..Default::default()
18140 }))
18141 .build()
18142 .unwrap();
18143
18144 let started = Instant::now();
18145 let record = agent
18146 .invoke_tool(ToolExecutionRequest::new(
18147 "recovery-timeout-call",
18148 "recovery_timeout",
18149 serde_json::json!({}),
18150 ToolCallSource::Manual,
18151 ))
18152 .await
18153 .unwrap();
18154
18155 assert!(record.executed);
18156 assert!(record.timed_out);
18157 assert!(!record.success);
18158 assert_eq!(calls.load(Ordering::SeqCst), 1);
18159 assert!(started.elapsed() < Duration::from_secs(1));
18160 assert_eq!(record.metadata["effective_limits"]["timeout_ms"], 100);
18161 let remaining_ms = remaining_ms.lock();
18162 assert_eq!(remaining_ms.len(), 1);
18163 assert!((1..=100).contains(&remaining_ms[0]));
18164 }
18165
18166 #[tokio::test]
18168 async fn recovery_default_timeout_controls_deadline_and_timer() {
18169 use ai_agents_recovery::{ErrorRecoveryConfig, ToolRecoveryConfig, ToolRetryConfig};
18170
18171 let calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
18172 let remaining_ms = Arc::new(parking_lot::Mutex::new(Vec::new()));
18173 let agent = AgentBuilder::new()
18174 .system_prompt("Test default recovery timeout.")
18175 .llm(Arc::new(mock_with_response("done")))
18176 .tool(Arc::new(ClassifiedTimeoutTool {
18177 id: "default_recovery_timeout",
18178 calls: Arc::clone(&calls),
18179 timeout_ms: 1_000,
18180 sleep_ms: 150,
18181 requires_approval: false,
18182 remaining_ms: Arc::clone(&remaining_ms),
18183 }))
18184 .recovery_manager(RecoveryManager::new(ErrorRecoveryConfig {
18185 tools: ToolRecoveryConfig {
18186 default: ToolRetryConfig {
18187 timeout_ms: Some(100),
18188 ..Default::default()
18189 },
18190 ..Default::default()
18191 },
18192 ..Default::default()
18193 }))
18194 .build()
18195 .unwrap();
18196
18197 let started = Instant::now();
18198 let record = agent
18199 .invoke_tool(ToolExecutionRequest::new(
18200 "default-recovery-timeout-call",
18201 "default_recovery_timeout",
18202 serde_json::json!({}),
18203 ToolCallSource::Manual,
18204 ))
18205 .await
18206 .unwrap();
18207
18208 assert!(record.executed);
18209 assert!(record.timed_out);
18210 assert!(!record.success);
18211 assert_eq!(calls.load(Ordering::SeqCst), 1);
18212 assert!(started.elapsed() < Duration::from_secs(1));
18213 assert_eq!(record.metadata["effective_limits"]["timeout_ms"], 100);
18214 let remaining_ms = remaining_ms.lock();
18215 assert_eq!(remaining_ms.len(), 1);
18216 assert!((1..=100).contains(&remaining_ms[0]));
18217 }
18218
18219 #[test]
18221 fn recovery_timeout_cannot_widen_security_baseline() {
18222 let security_engine = ToolSecurityEngine::new(ToolSecurityConfig {
18223 default_timeout_ms: 100,
18224 ..Default::default()
18225 });
18226 let safety = ToolSafetyMetadata::compute();
18227 let mut classification = ToolCallClassification::from_metadata(&safety);
18228 classification.timeout_ms = Some(500);
18229
18230 let (limits, timeout) = RuntimeAgent::effective_tool_limits(
18231 &security_engine,
18232 "recovery_cannot_widen",
18233 &safety,
18234 &classification,
18235 Some(1_000),
18236 )
18237 .unwrap();
18238
18239 assert_eq!(limits.timeout_ms, Some(100));
18240 assert_eq!(timeout.timer, Duration::from_millis(100));
18241 }
18242
18243 #[tokio::test]
18245 async fn invalid_call_timeout_stops_before_approval_or_tool_invocation() {
18246 let tool_calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
18247 let approval_calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
18248 let remaining_ms = Arc::new(parking_lot::Mutex::new(Vec::new()));
18249 let mut security = ToolSecurityConfig {
18250 enabled: true,
18251 ..Default::default()
18252 };
18253 security.tools.insert(
18254 "invalid_call_timeout".to_string(),
18255 ai_agents_tools::ToolPolicyConfig {
18256 require_confirmation: true,
18257 ..Default::default()
18258 },
18259 );
18260 let agent = AgentBuilder::new()
18261 .system_prompt("Test invalid call timeout.")
18262 .llm(Arc::new(mock_with_response("done")))
18263 .tool(Arc::new(ClassifiedTimeoutTool {
18264 id: "invalid_call_timeout",
18265 calls: Arc::clone(&tool_calls),
18266 timeout_ms: u64::MAX,
18267 sleep_ms: 0,
18268 requires_approval: false,
18269 remaining_ms,
18270 }))
18271 .tool_security(ToolSecurityEngine::new(security))
18272 .hitl_engine(HITLEngine::new(ai_agents_hitl::HITLConfig::default()))
18273 .approval_handler(Arc::new(CountingApprovalHandler {
18274 calls: Arc::clone(&approval_calls),
18275 }))
18276 .build()
18277 .unwrap();
18278
18279 let error = agent
18280 .invoke_tool(ToolExecutionRequest::new(
18281 "invalid-call-timeout",
18282 "invalid_call_timeout",
18283 serde_json::json!({}),
18284 ToolCallSource::Manual,
18285 ))
18286 .await
18287 .unwrap_err();
18288
18289 assert!(error.to_string().contains(
18290 "effective tool timeout_ms must be no greater than 3153600000000000 milliseconds"
18291 ));
18292 assert_eq!(approval_calls.load(Ordering::SeqCst), 0);
18293 assert_eq!(tool_calls.load(Ordering::SeqCst), 0);
18294 }
18295
18296 #[tokio::test]
18298 async fn invalid_modified_call_timeout_stops_before_lock_or_invocation() {
18299 use ai_agents_hitl::CallbackHandler;
18300
18301 let blocker_gate = PathMutationGate::new();
18302 let tool_calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
18303 let hooks = Arc::new(ToolLifecycleRecordingHooks::new());
18304 let handler = CallbackHandler::new(|_| ApprovalResult::Modified {
18305 changes: HashMap::from([("invalid_timeout".to_string(), Value::Bool(true))]),
18306 });
18307 let agent = Arc::new(
18308 AgentBuilder::new()
18309 .system_prompt("Test final call timeout validation.")
18310 .llm(Arc::new(mock_with_response("done")))
18311 .tool(Arc::new(BlockingPathMutationTool {
18312 id: "timeout_lock_blocker",
18313 path_fields: vec![ai_agents_core::PathPolicyBinding::write("path")],
18314 gate: blocker_gate.clone(),
18315 }))
18316 .tool(Arc::new(ApprovalModifiedTimeoutTool {
18317 calls: Arc::clone(&tool_calls),
18318 }))
18319 .hitl_engine(HITLEngine::new(ai_agents_hitl::HITLConfig::default()))
18320 .approval_handler(Arc::new(handler))
18321 .hooks(hooks.clone())
18322 .build()
18323 .unwrap(),
18324 );
18325 let blocking_agent = Arc::clone(&agent);
18326 let blocker = tokio::spawn(async move {
18327 blocking_agent
18328 .invoke_tool(ToolExecutionRequest::new(
18329 "timeout-lock-blocker",
18330 "timeout_lock_blocker",
18331 serde_json::json!({"path": "./shared-timeout.txt"}),
18332 ToolCallSource::Manual,
18333 ))
18334 .await
18335 .unwrap()
18336 });
18337 blocker_gate.wait_until_entered().await;
18338
18339 let record = tokio::time::timeout(
18340 Duration::from_millis(500),
18341 agent.invoke_tool(ToolExecutionRequest::new(
18342 "invalid-modified-timeout",
18343 "approval_modified_timeout",
18344 serde_json::json!({
18345 "path": "./shared-timeout.txt",
18346 "invalid_timeout": false
18347 }),
18348 ToolCallSource::Manual,
18349 )),
18350 )
18351 .await
18352 .expect("final timeout validation must not wait for the held path lock")
18353 .unwrap();
18354
18355 blocker_gate.release();
18356 assert!(blocker.await.unwrap().success);
18357 assert!(!record.executed);
18358 assert!(!record.success);
18359 assert_eq!(record.policy.outcome, PermissionOutcome::Deny);
18360 assert!(record.output.contains(
18361 "effective tool timeout_ms must be no greater than 3153600000000000 milliseconds"
18362 ));
18363 assert_eq!(tool_calls.load(Ordering::SeqCst), 0);
18364 let invalid_request_events = hooks
18365 .events()
18366 .into_iter()
18367 .filter(|event| event.contains("approval_modified_timeout") || event == "error")
18368 .collect::<Vec<_>>();
18369 assert_eq!(
18370 invalid_request_events,
18371 vec![
18372 "start:approval_modified_timeout",
18373 "complete:approval_modified_timeout:false",
18374 "record:approval_modified_timeout:false",
18375 "error"
18376 ]
18377 );
18378 }
18379
18380 #[tokio::test]
18381 async fn side_effecting_tools_are_serialized_per_resource() {
18382 let mock = mock_with_response("hello");
18383 let active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
18384 let max_active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
18385 let agent = Arc::new(
18386 AgentBuilder::new()
18387 .system_prompt("You are helpful.")
18388 .llm(Arc::new(mock))
18389 .tool(Arc::new(LockedWriteTool {
18390 active: Arc::clone(&active),
18391 max_active: Arc::clone(&max_active),
18392 }))
18393 .build()
18394 .unwrap(),
18395 );
18396
18397 let left = {
18398 let agent = Arc::clone(&agent);
18399 tokio::spawn(async move {
18400 agent
18401 .invoke_tool(ToolExecutionRequest::new(
18402 "lock-1",
18403 "locked_write",
18404 serde_json::json!({"path": "./same.txt"}),
18405 ToolCallSource::Manual,
18406 ))
18407 .await
18408 .unwrap()
18409 })
18410 };
18411 let right = {
18412 let agent = Arc::clone(&agent);
18413 tokio::spawn(async move {
18414 agent
18415 .invoke_tool(ToolExecutionRequest::new(
18416 "lock-2",
18417 "locked_write",
18418 serde_json::json!({"path": "./same.txt"}),
18419 ToolCallSource::Manual,
18420 ))
18421 .await
18422 .unwrap()
18423 })
18424 };
18425
18426 let left = left.await.unwrap();
18427 let right = right.await.unwrap();
18428 assert!(left.success);
18429 assert!(right.success);
18430 assert_eq!(max_active.load(Ordering::SeqCst), 1);
18431 }
18432
18433 #[tokio::test]
18434 async fn path_resources_use_shared_global_lock_and_cleanup() {
18435 let active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
18436 let max_active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
18437 let bindings = ai_agents_core::ToolPolicyBindings {
18438 path_fields: vec![
18439 ai_agents_core::PathPolicyBinding::read_write("source_path"),
18440 ai_agents_core::PathPolicyBinding::write("destination_path"),
18441 ],
18442 ..Default::default()
18443 };
18444 let classification = ai_agents_core::ToolCallClassification::from_metadata(
18445 &MultiResourceWriteTool {
18446 active: Arc::clone(&active),
18447 max_active: Arc::clone(&max_active),
18448 }
18449 .safety_metadata(),
18450 );
18451 let left_args = serde_json::json!({
18452 "source_path": "./a/../first.txt",
18453 "destination_path": "./second.txt"
18454 });
18455 let right_args = serde_json::json!({
18456 "source_path": "./second.txt",
18457 "destination_path": "./first.txt"
18458 });
18459 let left_keys = tool_resource_lock_keys(
18460 "multi_resource_write",
18461 &left_args,
18462 &bindings,
18463 &classification,
18464 );
18465 let right_keys = tool_resource_lock_keys(
18466 "multi_resource_write",
18467 &right_args,
18468 &bindings,
18469 &classification,
18470 );
18471 assert_eq!(left_keys, right_keys);
18472 assert_eq!(left_keys, vec!["path-mutation:global".to_string()]);
18473
18474 let locks = new_tool_resource_locks();
18475 let build_agent = || {
18476 AgentBuilder::new()
18477 .system_prompt("Test shared resource locks.")
18478 .llm(Arc::new(mock_with_response("done")))
18479 .tool(Arc::new(MultiResourceWriteTool {
18480 active: Arc::clone(&active),
18481 max_active: Arc::clone(&max_active),
18482 }))
18483 .build()
18484 .unwrap()
18485 .with_shared_resource_locks(Arc::clone(&locks))
18486 };
18487 let left_agent = Arc::new(build_agent());
18488 let right_agent = Arc::new(build_agent());
18489 let left = tokio::spawn(async move {
18490 left_agent
18491 .invoke_tool(ToolExecutionRequest::new(
18492 "multi-left",
18493 "multi_resource_write",
18494 left_args,
18495 ToolCallSource::Manual,
18496 ))
18497 .await
18498 .unwrap()
18499 });
18500 let right = tokio::spawn(async move {
18501 right_agent
18502 .invoke_tool(ToolExecutionRequest::new(
18503 "multi-right",
18504 "multi_resource_write",
18505 right_args,
18506 ToolCallSource::Manual,
18507 ))
18508 .await
18509 .unwrap()
18510 });
18511 let (left, right) = tokio::time::timeout(std::time::Duration::from_secs(2), async {
18512 tokio::join!(left, right)
18513 })
18514 .await
18515 .expect("reversed resource acquisition must not deadlock");
18516
18517 assert!(left.unwrap().success);
18518 assert!(right.unwrap().success);
18519 assert_eq!(max_active.load(Ordering::SeqCst), 1);
18520 assert!(locks.read().is_empty());
18521 }
18522
18523 #[tokio::test]
18524 async fn global_path_lock_serializes_copy_destination_with_file_write() {
18525 assert_path_mutation_pair_serialized(
18526 "copy_path",
18527 CopyPathTool::new().policy_bindings().path_fields,
18528 serde_json::json!({
18529 "source_path": "./source.txt",
18530 "destination_path": "./shared.txt"
18531 }),
18532 "file_write",
18533 FileWriteTool::new().policy_bindings().path_fields,
18534 serde_json::json!({"path": "./shared.txt"}),
18535 )
18536 .await;
18537 }
18538
18539 #[tokio::test]
18540 async fn parent_and_spawned_runtime_share_global_path_lock() {
18541 let workspace = MutationTestWorkspace::new();
18542 let destination = workspace.root.join("spawned.txt");
18543 let parent_gate = PathMutationGate::new();
18544 let parent = Arc::new(
18545 AgentBuilder::from_yaml(
18546 r#"
18547name: LockParent
18548system_prompt: parent
18549llm:
18550 default: default
18551tools:
18552 - parent_path_write
18553spawner:
18554 shared_llms: true
18555"#,
18556 )
18557 .unwrap()
18558 .llm(Arc::new(mock_with_response("done")))
18559 .auto_configure_spawner()
18560 .await
18561 .unwrap()
18562 .tool(Arc::new(BlockingPathMutationTool {
18563 id: "parent_path_write",
18564 path_fields: vec![ai_agents_core::PathPolicyBinding::write("path")],
18565 gate: parent_gate.clone(),
18566 }))
18567 .build()
18568 .unwrap(),
18569 );
18570
18571 let mut child_spec = crate::spec::AgentSpec {
18572 name: "LockChild".to_string(),
18573 system_prompt: "child".to_string(),
18574 tools: Some(vec![crate::spec::ToolEntry::Simple(
18575 "file_write".to_string(),
18576 )]),
18577 ..Default::default()
18578 };
18579 child_spec.tool_security.enabled = true;
18580 child_spec.tool_security.fail_closed = true;
18581 let file_write_policy = ai_agents_tools::ToolPolicyConfig {
18582 write_paths: vec![workspace.root.to_string_lossy().into_owned()],
18583 allow_without_confirmation: true,
18584 ..Default::default()
18585 };
18586 child_spec
18587 .tool_security
18588 .tools
18589 .insert("file_write".to_string(), file_write_policy);
18590 let spawned = parent
18591 .spawner()
18592 .unwrap()
18593 .spawn_from_spec(child_spec)
18594 .await
18595 .unwrap();
18596 assert!(Arc::ptr_eq(
18597 &parent.resource_locks,
18598 &spawned.agent.resource_locks
18599 ));
18600 assert!(!Arc::ptr_eq(
18601 &parent.runtime_control,
18602 &spawned.agent.runtime_control
18603 ));
18604
18605 let parent_call = {
18606 let parent = Arc::clone(&parent);
18607 let destination = destination.clone();
18608 tokio::spawn(async move {
18609 parent
18610 .invoke_tool(ToolExecutionRequest::new(
18611 "parent-lock-holder",
18612 "parent_path_write",
18613 serde_json::json!({"path": destination}),
18614 ToolCallSource::Manual,
18615 ))
18616 .await
18617 .unwrap()
18618 })
18619 };
18620 parent_gate.wait_until_entered().await;
18621
18622 let child_call = {
18623 let child = Arc::clone(&spawned.agent);
18624 let destination = destination.clone();
18625 tokio::spawn(async move {
18626 child
18627 .invoke_tool(ToolExecutionRequest::new(
18628 "spawned-file-write",
18629 "file_write",
18630 serde_json::json!({
18631 "path": destination,
18632 "content": "spawned",
18633 "dry_run": false
18634 }),
18635 ToolCallSource::Manual,
18636 ))
18637 .await
18638 .unwrap()
18639 })
18640 };
18641 wait_for_resource_lock_strong_count(&parent.resource_locks, 2).await;
18642 assert!(!child_call.is_finished());
18643
18644 parent_gate.release();
18645 let (parent_record, child_record) =
18646 tokio::time::timeout(std::time::Duration::from_secs(2), async {
18647 tokio::join!(parent_call, child_call)
18648 })
18649 .await
18650 .expect("parent and spawned path mutations did not finish");
18651 assert!(parent_record.unwrap().success);
18652 assert!(child_record.unwrap().success);
18653 assert_eq!(std::fs::read_to_string(destination).unwrap(), "spawned");
18654 assert!(parent.resource_locks.read().is_empty());
18655 }
18656
18657 #[tokio::test]
18658 async fn cancelled_global_path_lock_waiter_does_not_retain_weak_entry() {
18659 let locks = new_tool_resource_locks();
18660 let holder_gate = PathMutationGate::new();
18661 let waiter_gate = PathMutationGate::new();
18662 waiter_gate.release();
18663 let holder = Arc::new(
18664 AgentBuilder::new()
18665 .system_prompt("Hold the global path lock.")
18666 .llm(Arc::new(mock_with_response("done")))
18667 .tool(Arc::new(BlockingPathMutationTool {
18668 id: "holder_write",
18669 path_fields: vec![ai_agents_core::PathPolicyBinding::write("path")],
18670 gate: holder_gate.clone(),
18671 }))
18672 .build()
18673 .unwrap()
18674 .with_shared_resource_locks(Arc::clone(&locks)),
18675 );
18676 let waiter = Arc::new(
18677 AgentBuilder::new()
18678 .system_prompt("Wait for the global path lock.")
18679 .llm(Arc::new(mock_with_response("done")))
18680 .tool(Arc::new(BlockingPathMutationTool {
18681 id: "waiter_write",
18682 path_fields: vec![ai_agents_core::PathPolicyBinding::write("path")],
18683 gate: waiter_gate.clone(),
18684 }))
18685 .build()
18686 .unwrap()
18687 .with_shared_resource_locks(Arc::clone(&locks)),
18688 );
18689
18690 let holder_call = {
18691 let holder = Arc::clone(&holder);
18692 tokio::spawn(async move {
18693 holder
18694 .invoke_tool(ToolExecutionRequest::new(
18695 "holder-call",
18696 "holder_write",
18697 serde_json::json!({"path": "./shared.txt"}),
18698 ToolCallSource::Manual,
18699 ))
18700 .await
18701 .unwrap()
18702 })
18703 };
18704 holder_gate.wait_until_entered().await;
18705
18706 let waiter_call = {
18707 let waiter = Arc::clone(&waiter);
18708 tokio::spawn(async move {
18709 waiter
18710 .invoke_tool(ToolExecutionRequest::new(
18711 "waiter-call",
18712 "waiter_write",
18713 serde_json::json!({"path": "./shared.txt"}),
18714 ToolCallSource::Manual,
18715 ))
18716 .await
18717 .unwrap()
18718 })
18719 };
18720 wait_for_resource_lock_strong_count(&locks, 2).await;
18721 waiter.runtime_control().cancel_all();
18722
18723 let waiter_record = tokio::time::timeout(std::time::Duration::from_secs(2), waiter_call)
18724 .await
18725 .expect("cancelled lock waiter did not finish")
18726 .unwrap();
18727 assert!(!waiter_record.success);
18728 assert!(!waiter_record.executed);
18729 assert!(waiter_record.cancelled);
18730 assert_eq!(
18731 waiter_record.cancellation_reason.as_deref(),
18732 Some("runtime control cancellation")
18733 );
18734 assert!(!waiter_gate.entered.load(Ordering::SeqCst));
18735 assert_eq!(
18736 locks
18737 .read()
18738 .get("path-mutation:global")
18739 .map_or(0, |lock| lock.strong_count()),
18740 1
18741 );
18742
18743 holder_gate.release();
18744 let holder_record = tokio::time::timeout(std::time::Duration::from_secs(2), holder_call)
18745 .await
18746 .expect("lock holder did not finish")
18747 .unwrap();
18748 assert!(holder_record.success);
18749 assert!(locks.read().is_empty());
18750 }
18751
18752 #[tokio::test]
18753 async fn path_mutation_policy_and_approval_denials_do_not_invoke_tools() {
18754 for denial in [MutationDenial::Policy, MutationDenial::Approval] {
18755 let tools: [Arc<dyn Tool>; 3] = [
18756 Arc::new(CopyPathTool::new()),
18757 Arc::new(MovePathTool::new()),
18758 Arc::new(DeletePathTool::new()),
18759 ];
18760 for tool in tools {
18761 assert_path_mutation_denied(tool, denial).await;
18762 }
18763 }
18764 }
18765
18766 #[tokio::test]
18767 async fn policy_denial_keeps_executor_hook_lifecycle_and_record_authority() {
18768 let workspace = MutationTestWorkspace::new();
18769 let target = workspace.root.join("denied.txt");
18770 let hooks = Arc::new(ToolLifecycleRecordingHooks::new());
18771 let agent = AgentBuilder::new()
18772 .system_prompt("Test denied tool hooks.")
18773 .llm(Arc::new(mock_with_response("done")))
18774 .tool(Arc::new(FileWriteTool::new()))
18775 .tool_security(ToolSecurityEngine::new(mutation_denial_security_config(
18776 "file_write",
18777 &workspace.root,
18778 MutationDenial::Policy,
18779 )))
18780 .hooks(hooks.clone())
18781 .build()
18782 .unwrap();
18783
18784 let record = agent
18785 .invoke_tool(ToolExecutionRequest::new(
18786 "denied-hook-call",
18787 "file_write",
18788 serde_json::json!({
18789 "path": target.to_string_lossy(),
18790 "content": "blocked"
18791 }),
18792 ToolCallSource::Manual,
18793 ))
18794 .await
18795 .unwrap();
18796
18797 assert!(!record.executed);
18798 assert!(!record.success);
18799 assert_eq!(record.policy.outcome, PermissionOutcome::Deny);
18800 assert_eq!(
18801 hooks.events(),
18802 vec![
18803 "start:file_write",
18804 "complete:file_write:false",
18805 "record:file_write:false",
18806 "error"
18807 ]
18808 );
18809 assert!(!target.exists());
18810 }
18811
18812 #[tokio::test]
18813 async fn approval_argument_changes_are_rechecked_against_final_scope() {
18814 let active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
18815 let max_active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
18816 let entered = Arc::new(tokio::sync::Barrier::new(2));
18817 let release = Arc::new(tokio::sync::Notify::new());
18818 let handler = Arc::new(BlockingApprovalHandler {
18819 entered: Arc::clone(&entered),
18820 release: Arc::clone(&release),
18821 result: ApprovalResult::Modified {
18822 changes: HashMap::from([(
18823 "path".to_string(),
18824 Value::String("./after-approval.txt".to_string()),
18825 )]),
18826 },
18827 });
18828 let agent = Arc::new(
18829 AgentBuilder::new()
18830 .system_prompt("Test final scope validation.")
18831 .llm(Arc::new(mock_with_response("done")))
18832 .tool(Arc::new(LockedWriteTool {
18833 active: Arc::clone(&active),
18834 max_active: Arc::clone(&max_active),
18835 }))
18836 .tool_security(ToolSecurityEngine::new(approval_security_config(true)))
18837 .hitl_engine(HITLEngine::new(ai_agents_hitl::HITLConfig::default()))
18838 .approval_handler(handler)
18839 .build()
18840 .unwrap(),
18841 );
18842 let control = agent.runtime_control();
18843 let running = Arc::clone(&agent);
18844 let call = tokio::spawn(async move {
18845 running
18846 .invoke_tool(ToolExecutionRequest::new(
18847 "approval-scope",
18848 "locked_write",
18849 serde_json::json!({"path": "./before-approval.txt"}),
18850 ToolCallSource::Manual,
18851 ))
18852 .await
18853 .unwrap()
18854 });
18855 entered.wait().await;
18856 let expected_version = control.set_tool_scope(Vec::new());
18857 release.notify_one();
18858 let record = call.await.unwrap();
18859
18860 assert!(!record.executed);
18861 assert!(!record.success);
18862 assert_eq!(record.runtime_config_version, expected_version);
18863 assert_eq!(record.executed_arguments["path"], "./after-approval.txt");
18864 assert_eq!(max_active.load(Ordering::SeqCst), 0);
18865 assert_eq!(
18866 record.metadata["runtime_scope_snapshot"],
18867 serde_json::json!([])
18868 );
18869 }
18870
18871 #[tokio::test]
18872 async fn approval_is_rechecked_against_final_policy_snapshot() {
18873 let active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
18874 let max_active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
18875 let entered = Arc::new(tokio::sync::Barrier::new(2));
18876 let release = Arc::new(tokio::sync::Notify::new());
18877 let handler = Arc::new(BlockingApprovalHandler {
18878 entered: Arc::clone(&entered),
18879 release: Arc::clone(&release),
18880 result: ApprovalResult::Approved,
18881 });
18882 let agent = Arc::new(
18883 AgentBuilder::new()
18884 .system_prompt("Test final policy validation.")
18885 .llm(Arc::new(mock_with_response("done")))
18886 .tool(Arc::new(LockedWriteTool {
18887 active: Arc::clone(&active),
18888 max_active: Arc::clone(&max_active),
18889 }))
18890 .tool_security(ToolSecurityEngine::new(approval_security_config(true)))
18891 .hitl_engine(HITLEngine::new(ai_agents_hitl::HITLConfig::default()))
18892 .approval_handler(handler)
18893 .build()
18894 .unwrap(),
18895 );
18896 let control = agent.runtime_control();
18897 let running = Arc::clone(&agent);
18898 let call = tokio::spawn(async move {
18899 running
18900 .invoke_tool(ToolExecutionRequest::new(
18901 "approval-policy",
18902 "locked_write",
18903 serde_json::json!({"path": "./policy.txt"}),
18904 ToolCallSource::Manual,
18905 ))
18906 .await
18907 .unwrap()
18908 });
18909 entered.wait().await;
18910 let expected_version = control.set_tool_security(approval_security_config(false));
18911 release.notify_one();
18912 let record = call.await.unwrap();
18913
18914 assert!(!record.executed);
18915 assert!(!record.success);
18916 assert_eq!(record.runtime_config_version, expected_version);
18917 assert_eq!(record.policy.outcome, PermissionOutcome::Unavailable);
18918 assert_eq!(max_active.load(Ordering::SeqCst), 0);
18919 assert!(record.metadata.contains_key("policy_snapshot"));
18920 }
18921
18922 #[test]
18923 fn invalid_live_policy_does_not_replace_snapshot_or_generation() {
18924 let agent = AgentBuilder::new()
18925 .system_prompt("Test runtime policy validation.")
18926 .llm(Arc::new(mock_with_response("done")))
18927 .build()
18928 .unwrap();
18929 let control = agent.runtime_control();
18930 let mut valid = ToolSecurityConfig::default();
18931 valid.tools.insert(
18932 "web_search".to_string(),
18933 ai_agents_tools::ToolPolicyConfig {
18934 max_results: Some(5),
18935 ..Default::default()
18936 },
18937 );
18938 let generation = control.try_set_tool_security(valid).unwrap();
18939
18940 let mut invalid = ToolSecurityConfig::default();
18941 invalid.tools.insert(
18942 "web_search".to_string(),
18943 ai_agents_tools::ToolPolicyConfig {
18944 max_results: Some(0),
18945 ..Default::default()
18946 },
18947 );
18948 let error = control.try_set_tool_security(invalid).unwrap_err();
18949
18950 assert!(
18951 error
18952 .to_string()
18953 .contains("max_results must be greater than 0")
18954 );
18955 assert_eq!(control.version(), generation);
18956 assert_eq!(
18957 control
18958 .state
18959 .tool_security_override
18960 .read()
18961 .as_ref()
18962 .unwrap()
18963 .config()
18964 .tools["web_search"]
18965 .max_results,
18966 Some(5)
18967 );
18968 }
18969
18970 #[test]
18972 fn invalid_timeout_config_stops_before_approval_or_tool_invocation() {
18973 let tool_calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
18974 let approval_calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
18975 let spec = crate::spec::AgentSpec {
18976 tool_security: ToolSecurityConfig {
18977 enabled: true,
18978 default_timeout_ms: u64::MAX,
18979 ..Default::default()
18980 },
18981 ..Default::default()
18982 };
18983
18984 let result = AgentBuilder::from_spec(spec)
18985 .llm(Arc::new(mock_with_response("done")))
18986 .tool(Arc::new(FlakyWriteTool {
18987 calls: Arc::clone(&tool_calls),
18988 }))
18989 .hitl_engine(HITLEngine::new(ai_agents_hitl::HITLConfig::default()))
18990 .approval_handler(Arc::new(CountingApprovalHandler {
18991 calls: Arc::clone(&approval_calls),
18992 }))
18993 .build();
18994
18995 assert!(result.is_err());
18996 assert_eq!(approval_calls.load(Ordering::SeqCst), 0);
18997 assert_eq!(tool_calls.load(Ordering::SeqCst), 0);
18998 }
18999
19000 #[test]
19002 fn invalid_recovery_timeout_config_stops_before_approval_or_tool_invocation() {
19003 use ai_agents_recovery::{ErrorRecoveryConfig, ToolRecoveryConfig, ToolRetryConfig};
19004
19005 let tool_calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19006 let approval_calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19007 let spec = crate::spec::AgentSpec {
19008 error_recovery: ErrorRecoveryConfig {
19009 tools: ToolRecoveryConfig {
19010 default: ToolRetryConfig {
19011 timeout_ms: Some(u64::MAX),
19012 ..Default::default()
19013 },
19014 ..Default::default()
19015 },
19016 ..Default::default()
19017 },
19018 ..Default::default()
19019 };
19020
19021 let result = AgentBuilder::from_spec(spec)
19022 .llm(Arc::new(mock_with_response("done")))
19023 .tool(Arc::new(FlakyWriteTool {
19024 calls: Arc::clone(&tool_calls),
19025 }))
19026 .hitl_engine(HITLEngine::new(ai_agents_hitl::HITLConfig::default()))
19027 .approval_handler(Arc::new(CountingApprovalHandler {
19028 calls: Arc::clone(&approval_calls),
19029 }))
19030 .build();
19031
19032 assert!(result.is_err());
19033 assert_eq!(approval_calls.load(Ordering::SeqCst), 0);
19034 assert_eq!(tool_calls.load(Ordering::SeqCst), 0);
19035 }
19036
19037 #[test]
19039 fn invalid_timeout_policy_does_not_replace_snapshot_or_generation() {
19040 let agent = AgentBuilder::new()
19041 .system_prompt("Test runtime timeout policy validation.")
19042 .llm(Arc::new(mock_with_response("done")))
19043 .build()
19044 .unwrap();
19045 let control = agent.runtime_control();
19046 let valid = ToolSecurityConfig {
19047 default_timeout_ms: 5_000,
19048 ..Default::default()
19049 };
19050 let generation = control.try_set_tool_security(valid).unwrap();
19051
19052 let invalid = ToolSecurityConfig {
19053 default_timeout_ms: MAX_TOOL_TIMEOUT_MS + 1,
19054 ..Default::default()
19055 };
19056 let error = control.try_set_tool_security(invalid).unwrap_err();
19057
19058 assert!(error.to_string().contains(&format!(
19059 "tool_security.default_timeout_ms must be no greater than {MAX_TOOL_TIMEOUT_MS} milliseconds"
19060 )));
19061 assert_eq!(control.version(), generation);
19062 assert_eq!(
19063 control
19064 .state
19065 .tool_security_override
19066 .read()
19067 .as_ref()
19068 .unwrap()
19069 .config()
19070 .default_timeout_ms,
19071 5_000
19072 );
19073 }
19074
19075 #[test]
19077 fn runtime_tool_timeout_conversion_enforces_the_stable_boundary() {
19078 let timeout = RuntimeAgent::validated_tool_timeout(MAX_TOOL_TIMEOUT_MS).unwrap();
19079 assert_eq!(timeout.timer, Duration::from_millis(MAX_TOOL_TIMEOUT_MS));
19080 assert_eq!(
19081 timeout.deadline_delta,
19082 chrono::Duration::milliseconds(MAX_TOOL_TIMEOUT_MS as i64)
19083 );
19084
19085 for timeout_ms in [MAX_TOOL_TIMEOUT_MS + 1, u64::MAX] {
19086 let error = RuntimeAgent::validated_tool_timeout(timeout_ms).unwrap_err();
19087 assert!(error.to_string().contains(&format!(
19088 "effective tool timeout_ms must be no greater than {MAX_TOOL_TIMEOUT_MS} milliseconds"
19089 )));
19090 }
19091 }
19092
19093 #[tokio::test]
19094 async fn persistent_override_preserves_rate_history_within_generation() {
19095 let calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19096 let agent = AgentBuilder::new()
19097 .system_prompt("Test persistent policy overrides.")
19098 .llm(Arc::new(mock_with_response("done")))
19099 .tool(Arc::new(RecoveryTestTool {
19100 id: "limited_override".to_string(),
19101 succeeds: true,
19102 calls: Arc::clone(&calls),
19103 max_output_chars: None,
19104 }))
19105 .build()
19106 .unwrap();
19107 let mut security = ToolSecurityConfig {
19108 enabled: true,
19109 fail_closed: true,
19110 ..Default::default()
19111 };
19112 let policy = ai_agents_tools::ToolPolicyConfig {
19113 write_paths: vec![".".to_string()],
19114 rate_limit: Some(1),
19115 ..Default::default()
19116 };
19117 security
19118 .tools
19119 .insert("limited_override".to_string(), policy);
19120 let generation = agent.runtime_control().set_tool_security(security);
19121
19122 let first = agent
19123 .invoke_tool(ToolExecutionRequest::new(
19124 "limited-first",
19125 "limited_override",
19126 serde_json::json!({"path": "./limited.txt"}),
19127 ToolCallSource::Manual,
19128 ))
19129 .await
19130 .unwrap();
19131 let second = agent
19132 .invoke_tool(ToolExecutionRequest::new(
19133 "limited-second",
19134 "limited_override",
19135 serde_json::json!({"path": "./limited.txt"}),
19136 ToolCallSource::Manual,
19137 ))
19138 .await
19139 .unwrap();
19140
19141 assert!(first.success);
19142 assert_eq!(first.policy_version, generation);
19143 assert!(!second.executed);
19144 assert!(second.output.contains("Rate limit exceeded"));
19145 assert_eq!(second.policy_version, generation);
19146 assert_eq!(calls.load(Ordering::SeqCst), 1);
19147 }
19148
19149 #[tokio::test]
19150 async fn concurrent_rate_admission_consumes_capacity_atomically() {
19151 let calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19152 let tool = Arc::new(RecoveryTestTool {
19153 id: "atomic_rate".to_string(),
19154 succeeds: true,
19155 calls: Arc::clone(&calls),
19156 max_output_chars: None,
19157 });
19158 let arguments = serde_json::json!({"path": "./atomic-rate.txt"});
19159 let bindings = tool.policy_bindings();
19160 let classification = tool.classify_call(&arguments);
19161 let resource_keys =
19162 tool_resource_lock_keys(tool.id(), &arguments, &bindings, &classification);
19163 let mut security = ToolSecurityConfig {
19164 enabled: true,
19165 fail_closed: true,
19166 ..Default::default()
19167 };
19168 let policy = ai_agents_tools::ToolPolicyConfig {
19169 write_paths: vec![".".to_string()],
19170 rate_limit: Some(1),
19171 ..Default::default()
19172 };
19173 security.tools.insert(tool.id().to_string(), policy);
19174 let agent = Arc::new(
19175 AgentBuilder::new()
19176 .system_prompt("Test atomic rate admission.")
19177 .llm(Arc::new(mock_with_response("done")))
19178 .tool(tool)
19179 .tool_security(ToolSecurityEngine::new(security))
19180 .build()
19181 .unwrap(),
19182 );
19183 let held = agent
19184 .acquire_tool_resource_locks(&resource_keys)
19185 .await
19186 .unwrap();
19187 let left = {
19188 let agent = Arc::clone(&agent);
19189 let arguments = arguments.clone();
19190 tokio::spawn(async move {
19191 agent
19192 .invoke_tool(ToolExecutionRequest::new(
19193 "atomic-rate-left",
19194 "atomic_rate",
19195 arguments,
19196 ToolCallSource::Manual,
19197 ))
19198 .await
19199 .unwrap()
19200 })
19201 };
19202 let right = {
19203 let agent = Arc::clone(&agent);
19204 tokio::spawn(async move {
19205 agent
19206 .invoke_tool(ToolExecutionRequest::new(
19207 "atomic-rate-right",
19208 "atomic_rate",
19209 arguments,
19210 ToolCallSource::Manual,
19211 ))
19212 .await
19213 .unwrap()
19214 })
19215 };
19216 tokio::time::sleep(std::time::Duration::from_millis(25)).await;
19217 drop(held);
19218 let (left, right) = tokio::join!(left, right);
19219 let records = [left.unwrap(), right.unwrap()];
19220
19221 assert_eq!(records.iter().filter(|record| record.success).count(), 1);
19222 assert_eq!(records.iter().filter(|record| record.executed).count(), 1);
19223 assert!(
19224 records.iter().any(|record| {
19225 !record.executed && record.output.contains("Rate limit exceeded")
19226 })
19227 );
19228 assert_eq!(calls.load(Ordering::SeqCst), 1);
19229 }
19230
19231 #[tokio::test]
19232 async fn changed_policy_generation_invalidates_pending_approval() {
19233 let active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19234 let max_active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19235 let entered = Arc::new(tokio::sync::Barrier::new(2));
19236 let release = Arc::new(tokio::sync::Notify::new());
19237 let handler = Arc::new(BlockingApprovalHandler {
19238 entered: Arc::clone(&entered),
19239 release: Arc::clone(&release),
19240 result: ApprovalResult::Approved,
19241 });
19242 let agent = Arc::new(
19243 AgentBuilder::new()
19244 .system_prompt("Test stale approval denial.")
19245 .llm(Arc::new(mock_with_response("done")))
19246 .tool(Arc::new(LockedWriteTool {
19247 active: Arc::clone(&active),
19248 max_active: Arc::clone(&max_active),
19249 }))
19250 .tool_security(ToolSecurityEngine::new(approval_security_config(true)))
19251 .hitl_engine(HITLEngine::new(ai_agents_hitl::HITLConfig::default()))
19252 .approval_handler(handler)
19253 .build()
19254 .unwrap(),
19255 );
19256 let running = Arc::clone(&agent);
19257 let call = tokio::spawn(async move {
19258 running
19259 .invoke_tool(ToolExecutionRequest::new(
19260 "stale-approval",
19261 "locked_write",
19262 serde_json::json!({"path": "./stale.txt"}),
19263 ToolCallSource::Manual,
19264 ))
19265 .await
19266 .unwrap()
19267 });
19268 entered.wait().await;
19269 let generation = agent
19270 .runtime_control()
19271 .set_tool_security(approval_security_config(true));
19272 release.notify_one();
19273 let record = call.await.unwrap();
19274
19275 assert!(!record.executed);
19276 assert!(record.output.contains("Approval became stale"));
19277 assert_eq!(record.policy_version, generation);
19278 assert_eq!(max_active.load(Ordering::SeqCst), 0);
19279 }
19280
19281 #[tokio::test]
19282 async fn final_policy_reapplies_argument_caps_after_approval_changes() {
19283 use ai_agents_hitl::CallbackHandler;
19284
19285 let mut security = ToolSecurityConfig {
19286 enabled: true,
19287 fail_closed: true,
19288 ..Default::default()
19289 };
19290 let policy = ai_agents_tools::ToolPolicyConfig {
19291 read_paths: vec![".".to_string()],
19292 max_results: Some(5),
19293 require_confirmation: true,
19294 ..Default::default()
19295 };
19296 security.tools.insert("context_echo".to_string(), policy);
19297 let handler = CallbackHandler::new(|_| ApprovalResult::Modified {
19298 changes: HashMap::from([("max_results".to_string(), serde_json::json!(99))]),
19299 });
19300 let agent = AgentBuilder::new()
19301 .system_prompt("Test final argument caps.")
19302 .llm(Arc::new(mock_with_response("done")))
19303 .tool(Arc::new(ContextEchoTool))
19304 .tool_security(ToolSecurityEngine::new(security))
19305 .hitl_engine(HITLEngine::new(ai_agents_hitl::HITLConfig::default()))
19306 .approval_handler(Arc::new(handler))
19307 .build()
19308 .unwrap();
19309
19310 let record = agent
19311 .invoke_tool(ToolExecutionRequest::new(
19312 "final-cap",
19313 "context_echo",
19314 serde_json::json!({"path": ".", "max_results": 1}),
19315 ToolCallSource::Manual,
19316 ))
19317 .await
19318 .unwrap();
19319
19320 assert!(record.success);
19321 assert_eq!(record.executed_arguments["max_results"], 5);
19322 assert_eq!(
19323 record.approval.unwrap().modified_arguments.unwrap()["max_results"],
19324 5
19325 );
19326 }
19327
19328 #[tokio::test]
19329 async fn no_binding_writes_use_canonical_fallback_lock() {
19330 let active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19331 let max_active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19332 let agent = Arc::new(
19333 AgentBuilder::new()
19334 .system_prompt("Test fallback resource locks.")
19335 .llm(Arc::new(mock_with_response("done")))
19336 .tool(Arc::new(NoBindingWriteTool {
19337 active: Arc::clone(&active),
19338 max_active: Arc::clone(&max_active),
19339 }))
19340 .build()
19341 .unwrap(),
19342 );
19343 let left = {
19344 let agent = Arc::clone(&agent);
19345 tokio::spawn(async move {
19346 agent
19347 .invoke_tool(ToolExecutionRequest::new(
19348 "no-binding-left",
19349 "no_binding_write",
19350 serde_json::json!({}),
19351 ToolCallSource::Manual,
19352 ))
19353 .await
19354 .unwrap()
19355 })
19356 };
19357 let right = {
19358 let agent = Arc::clone(&agent);
19359 tokio::spawn(async move {
19360 agent
19361 .invoke_tool(ToolExecutionRequest::new(
19362 "no-binding-right",
19363 "no_binding_write",
19364 serde_json::json!({}),
19365 ToolCallSource::Manual,
19366 ))
19367 .await
19368 .unwrap()
19369 })
19370 };
19371 let (left, right) = tokio::join!(left, right);
19372
19373 assert!(left.unwrap().success);
19374 assert!(right.unwrap().success);
19375 assert_eq!(max_active.load(Ordering::SeqCst), 1);
19376 }
19377
19378 #[tokio::test]
19379 async fn parent_and_child_paths_share_a_resource_lock() {
19380 let active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19381 let max_active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19382 let agent = Arc::new(
19383 AgentBuilder::new()
19384 .system_prompt("Test parent child resource locks.")
19385 .llm(Arc::new(mock_with_response("done")))
19386 .tool(Arc::new(LockedWriteTool {
19387 active: Arc::clone(&active),
19388 max_active: Arc::clone(&max_active),
19389 }))
19390 .build()
19391 .unwrap(),
19392 );
19393 let parent = format!("./lock-parent-{}", uuid::Uuid::new_v4());
19394 let child = format!("{}/child.txt", parent);
19395 let left = {
19396 let agent = Arc::clone(&agent);
19397 tokio::spawn(async move {
19398 agent
19399 .invoke_tool(ToolExecutionRequest::new(
19400 "parent-lock",
19401 "locked_write",
19402 serde_json::json!({"path": parent}),
19403 ToolCallSource::Manual,
19404 ))
19405 .await
19406 .unwrap()
19407 })
19408 };
19409 let right = {
19410 let agent = Arc::clone(&agent);
19411 tokio::spawn(async move {
19412 agent
19413 .invoke_tool(ToolExecutionRequest::new(
19414 "child-lock",
19415 "locked_write",
19416 serde_json::json!({"path": child}),
19417 ToolCallSource::Manual,
19418 ))
19419 .await
19420 .unwrap()
19421 })
19422 };
19423 let (left, right) = tokio::join!(left, right);
19424
19425 assert!(left.unwrap().success);
19426 assert!(right.unwrap().success);
19427 assert_eq!(max_active.load(Ordering::SeqCst), 1);
19428 }
19429
19430 #[tokio::test]
19431 async fn tool_hooks_can_reenter_after_resource_guards_are_dropped() {
19432 let calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19433 let hooks = Arc::new(ReentrantToolHooks {
19434 agent: parking_lot::Mutex::new(None),
19435 invoked: AtomicBool::new(false),
19436 nested_success: AtomicBool::new(false),
19437 });
19438 let agent = Arc::new(
19439 AgentBuilder::new()
19440 .system_prompt("Test hook reentrancy.")
19441 .llm(Arc::new(mock_with_response("done")))
19442 .tool(Arc::new(RecoveryTestTool {
19443 id: "reentrant_write".to_string(),
19444 succeeds: true,
19445 calls: Arc::clone(&calls),
19446 max_output_chars: None,
19447 }))
19448 .hooks(hooks.clone())
19449 .build()
19450 .unwrap(),
19451 );
19452 *hooks.agent.lock() = Some(Arc::downgrade(&agent));
19453 let record = tokio::time::timeout(
19454 std::time::Duration::from_secs(2),
19455 agent.invoke_tool(ToolExecutionRequest::new(
19456 "outer-hook-call",
19457 "reentrant_write",
19458 serde_json::json!({"path": "./hook.txt"}),
19459 ToolCallSource::Manual,
19460 )),
19461 )
19462 .await
19463 .expect("tool completion hook must not retain resource guards")
19464 .unwrap();
19465
19466 assert!(record.success);
19467 assert!(hooks.nested_success.load(Ordering::SeqCst));
19468 assert_eq!(calls.load(Ordering::SeqCst), 2);
19469 }
19470
19471 #[tokio::test]
19473 async fn fallback_finalizes_original_record_before_shared_execution() {
19474 let primary_calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19475 let fallback_calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19476 let hooks = Arc::new(ToolLifecycleRecordingHooks::new());
19477 let agent = AgentBuilder::new()
19478 .system_prompt("Test fallback execution.")
19479 .llm(Arc::new(mock_with_response("done")))
19480 .tool(Arc::new(RecoveryTestTool {
19481 id: "primary".to_string(),
19482 succeeds: false,
19483 calls: Arc::clone(&primary_calls),
19484 max_output_chars: None,
19485 }))
19486 .tool(Arc::new(RecoveryTestTool {
19487 id: "fallback".to_string(),
19488 succeeds: true,
19489 calls: Arc::clone(&fallback_calls),
19490 max_output_chars: None,
19491 }))
19492 .recovery_manager(recovery_manager_with_fallbacks([(
19493 "primary".to_string(),
19494 "fallback".to_string(),
19495 )]))
19496 .hooks(hooks.clone())
19497 .build()
19498 .unwrap();
19499 let record = tokio::time::timeout(
19500 std::time::Duration::from_secs(2),
19501 agent.invoke_tool(ToolExecutionRequest::new(
19502 "fallback-call",
19503 "primary",
19504 serde_json::json!({"path": "./shared.txt"}),
19505 ToolCallSource::Manual,
19506 )),
19507 )
19508 .await
19509 .expect("fallback must not retain the primary resource guard")
19510 .unwrap();
19511
19512 assert_eq!(
19513 hooks.events(),
19514 vec![
19515 "start:primary",
19516 "complete:primary:false",
19517 "record:primary:true",
19518 "error",
19519 "start:fallback",
19520 "complete:fallback:true",
19521 "record:fallback:true",
19522 ]
19523 );
19524 let records = hooks.records();
19525 assert_eq!(records.len(), 2);
19526 let original = &records[0];
19527 assert_eq!(original.canonical_id, "primary");
19528 assert!(matches!(original.source, ToolCallSource::Manual));
19529 assert!(original.executed);
19530 assert!(!original.success);
19531
19532 let fallback = &records[1];
19533 assert_eq!(fallback.canonical_id, "fallback");
19534 assert_eq!(fallback.call_id, "fallback-call");
19535 assert!(matches!(
19536 &fallback.source,
19537 ToolCallSource::Fallback { original_tool } if original_tool == "primary"
19538 ));
19539 assert!(fallback.executed);
19540 assert!(fallback.success);
19541 assert_eq!(record.canonical_id, fallback.canonical_id);
19542 assert_eq!(record.output, fallback.output);
19543
19544 let history = agent.tool_call_history();
19545 assert_eq!(
19546 history
19547 .iter()
19548 .map(|entry| entry.tool_id.as_str())
19549 .collect::<Vec<_>>(),
19550 vec!["primary", "fallback"]
19551 );
19552 assert_eq!(history[0].result.get("success"), Some(&Value::Bool(false)));
19553 assert_eq!(primary_calls.load(Ordering::SeqCst), 1);
19554 assert_eq!(fallback_calls.load(Ordering::SeqCst), 1);
19555 }
19556
19557 #[tokio::test]
19559 async fn self_fallback_cycle_is_denied_before_reinvocation() {
19560 let calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19561 let hooks = Arc::new(ToolLifecycleRecordingHooks::new());
19562 let agent = AgentBuilder::new()
19563 .system_prompt("Test self-fallback cycle admission.")
19564 .llm(Arc::new(mock_with_response("done")))
19565 .tool(Arc::new(RecoveryTestTool {
19566 id: "primary".to_string(),
19567 succeeds: false,
19568 calls: Arc::clone(&calls),
19569 max_output_chars: None,
19570 }))
19571 .recovery_manager(recovery_manager_with_fallbacks([(
19572 "primary".to_string(),
19573 "primary".to_string(),
19574 )]))
19575 .hooks(hooks.clone())
19576 .build()
19577 .unwrap();
19578
19579 let record = tokio::time::timeout(
19580 std::time::Duration::from_secs(2),
19581 agent.invoke_tool(ToolExecutionRequest::new(
19582 "self-fallback-call",
19583 "primary",
19584 serde_json::json!({"path": "./shared.txt"}),
19585 ToolCallSource::Manual,
19586 )),
19587 )
19588 .await
19589 .expect("self fallback must terminate without recursive execution")
19590 .unwrap();
19591
19592 assert_eq!(calls.load(Ordering::SeqCst), 1);
19593 assert_eq!(record.canonical_id, "primary");
19594 assert!(!record.executed);
19595 assert!(!record.success);
19596 assert_eq!(record.policy.outcome, PermissionOutcome::Deny);
19597 assert!(record.output.contains("fallback cycle"));
19598 assert!(matches!(
19599 record.source,
19600 ToolCallSource::Fallback { ref original_tool } if original_tool == "primary"
19601 ));
19602 assert_eq!(
19603 record.metadata.get("fallback_chain"),
19604 Some(&serde_json::json!(["primary"]))
19605 );
19606 assert_eq!(
19607 hooks.events(),
19608 vec![
19609 "start:primary",
19610 "complete:primary:false",
19611 "record:primary:true",
19612 "error",
19613 "complete:primary:false",
19614 "record:primary:false",
19615 "error",
19616 ]
19617 );
19618 assert_eq!(agent.tool_call_history().len(), 2);
19619 }
19620
19621 #[tokio::test]
19623 async fn alias_mediated_fallback_cycle_is_denied_canonically() {
19624 let primary_calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19625 let secondary_calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19626 let hooks = Arc::new(ToolLifecycleRecordingHooks::new());
19627 let agent = AgentBuilder::new()
19628 .system_prompt("Test canonical fallback cycle admission.")
19629 .llm(Arc::new(mock_with_response("done")))
19630 .tool(Arc::new(RecoveryTestTool {
19631 id: "primary".to_string(),
19632 succeeds: false,
19633 calls: Arc::clone(&primary_calls),
19634 max_output_chars: None,
19635 }))
19636 .tool(Arc::new(RecoveryTestTool {
19637 id: "secondary".to_string(),
19638 succeeds: false,
19639 calls: Arc::clone(&secondary_calls),
19640 max_output_chars: None,
19641 }))
19642 .recovery_manager(recovery_manager_with_fallbacks([
19643 ("primary".to_string(), "secondary".to_string()),
19644 ("secondary".to_string(), "primary alias".to_string()),
19645 ]))
19646 .hooks(hooks.clone())
19647 .build()
19648 .unwrap();
19649 agent.tools.set_tool_aliases(
19650 "primary",
19651 ToolAliases::new().with_name("en", "primary alias"),
19652 );
19653
19654 let record = tokio::time::timeout(
19655 std::time::Duration::from_secs(2),
19656 agent.invoke_tool(ToolExecutionRequest::new(
19657 "alias-fallback-call",
19658 "primary",
19659 serde_json::json!({"path": "./shared.txt"}),
19660 ToolCallSource::Manual,
19661 )),
19662 )
19663 .await
19664 .expect("alias-mediated fallback cycle must terminate")
19665 .unwrap();
19666
19667 assert_eq!(primary_calls.load(Ordering::SeqCst), 1);
19668 assert_eq!(secondary_calls.load(Ordering::SeqCst), 1);
19669 assert_eq!(record.requested_name, "primary alias");
19670 assert_eq!(record.canonical_id, "primary");
19671 assert!(!record.executed);
19672 assert!(record.output.contains("fallback cycle"));
19673 assert_eq!(
19674 record.metadata.get("fallback_chain"),
19675 Some(&serde_json::json!(["primary", "secondary"]))
19676 );
19677 assert_eq!(hooks.records().len(), 3);
19678 assert_eq!(agent.tool_call_history().len(), 3);
19679 }
19680
19681 #[tokio::test]
19683 async fn final_canonical_drift_cannot_bypass_fallback_ancestry() {
19684 let primary_calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19685 let secondary_calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19686 let provider = Arc::new(DriftingFallbackProvider {
19687 refreshed: AtomicBool::new(false),
19688 primary_calls: Arc::clone(&primary_calls),
19689 secondary_calls: Arc::clone(&secondary_calls),
19690 });
19691 let registry = ToolRegistry::new();
19692 registry.register_provider(provider).await.unwrap();
19693 let lifecycle = Arc::new(ToolLifecycleRecordingHooks::new());
19694 let hooks = Arc::new(RefreshFallbackProviderHooks {
19695 agent: parking_lot::Mutex::new(None),
19696 lifecycle: Arc::clone(&lifecycle),
19697 });
19698 let agent = Arc::new(
19699 AgentBuilder::new()
19700 .system_prompt("Test final canonical fallback admission.")
19701 .llm(Arc::new(mock_with_response("done")))
19702 .tools(registry)
19703 .recovery_manager(recovery_manager_with_fallbacks([(
19704 "primary".to_string(),
19705 "fallback alias".to_string(),
19706 )]))
19707 .hooks(hooks.clone())
19708 .build()
19709 .unwrap(),
19710 );
19711 *hooks.agent.lock() = Some(Arc::downgrade(&agent));
19712
19713 let record = agent
19714 .invoke_tool(ToolExecutionRequest::new(
19715 "drifting-fallback-call",
19716 "primary",
19717 serde_json::json!({"path": "./shared.txt"}),
19718 ToolCallSource::Manual,
19719 ))
19720 .await
19721 .unwrap();
19722
19723 assert_eq!(primary_calls.load(Ordering::SeqCst), 1);
19724 assert_eq!(secondary_calls.load(Ordering::SeqCst), 0);
19725 assert_eq!(record.canonical_id, "secondary");
19726 assert!(!record.executed);
19727 assert_eq!(record.policy.outcome, PermissionOutcome::Deny);
19728 assert!(record.output.contains("fallback cycle"));
19729 assert_eq!(
19730 record.metadata.get("fallback_chain"),
19731 Some(&serde_json::json!(["primary", "secondary"]))
19732 );
19733 assert_eq!(
19734 record.metadata.get("final_resolved_canonical_id"),
19735 Some(&serde_json::json!("primary"))
19736 );
19737 assert_eq!(
19738 lifecycle.events(),
19739 vec![
19740 "start:primary",
19741 "complete:primary:false",
19742 "record:primary:true",
19743 "error",
19744 "start:secondary",
19745 "complete:secondary:false",
19746 "record:secondary:false",
19747 "error",
19748 ]
19749 );
19750 let records = lifecycle.records();
19751 assert_eq!(records.len(), 2);
19752 assert_eq!(records[1].canonical_id, "secondary");
19753 assert_eq!(
19754 records[1].metadata.get("final_resolved_canonical_id"),
19755 Some(&serde_json::json!("primary"))
19756 );
19757 let history = agent.tool_call_history();
19758 assert_eq!(
19759 history
19760 .iter()
19761 .map(|entry| entry.tool_id.as_str())
19762 .collect::<Vec<_>>(),
19763 vec!["primary", "secondary"]
19764 );
19765 }
19766
19767 #[tokio::test]
19769 async fn acyclic_fallback_chain_is_denied_after_the_hop_limit() {
19770 let tool_count = MAX_TOOL_FALLBACK_HOPS + 2;
19771 let calls = (0..tool_count)
19772 .map(|_| Arc::new(std::sync::atomic::AtomicUsize::new(0)))
19773 .collect::<Vec<_>>();
19774 let mut builder = AgentBuilder::new()
19775 .system_prompt("Test bounded acyclic fallback admission.")
19776 .llm(Arc::new(mock_with_response("done")));
19777 for (index, counter) in calls.iter().enumerate() {
19778 builder = builder.tool(Arc::new(RecoveryTestTool {
19779 id: format!("fallback_{index}"),
19780 succeeds: false,
19781 calls: Arc::clone(counter),
19782 max_output_chars: None,
19783 }));
19784 }
19785 let fallbacks = (0..tool_count - 1).map(|index| {
19786 (
19787 format!("fallback_{index}"),
19788 format!("fallback_{}", index + 1),
19789 )
19790 });
19791 let agent = builder
19792 .recovery_manager(recovery_manager_with_fallbacks(fallbacks))
19793 .build()
19794 .unwrap();
19795
19796 let record = tokio::time::timeout(
19797 std::time::Duration::from_secs(2),
19798 agent.invoke_tool(ToolExecutionRequest::new(
19799 "bounded-fallback-call",
19800 "fallback_0",
19801 serde_json::json!({"path": "./shared.txt"}),
19802 ToolCallSource::Manual,
19803 )),
19804 )
19805 .await
19806 .expect("bounded fallback chain must terminate")
19807 .unwrap();
19808
19809 for counter in calls.iter().take(MAX_TOOL_FALLBACK_HOPS + 1) {
19810 assert_eq!(counter.load(Ordering::SeqCst), 1);
19811 }
19812 assert_eq!(calls[MAX_TOOL_FALLBACK_HOPS + 1].load(Ordering::SeqCst), 0);
19813 assert_eq!(
19814 record.canonical_id,
19815 format!("fallback_{}", MAX_TOOL_FALLBACK_HOPS + 1)
19816 );
19817 assert!(!record.executed);
19818 assert!(record.output.contains("maximum of 16 hops"));
19819 assert_eq!(agent.tool_call_history().len(), tool_count);
19820 }
19821
19822 #[tokio::test]
19823 async fn diagnostics_without_provider_records_unavailable_without_execution() {
19824 let mock = mock_with_response("hello");
19825 let yaml = r#"
19826name: DiagnosticsNoProviderAgent
19827system_prompt: "Review diagnostics."
19828tools: [diagnostics]
19829"#;
19830 let agent = AgentBuilder::from_yaml(yaml)
19831 .unwrap()
19832 .llm(Arc::new(mock))
19833 .auto_configure_features()
19834 .unwrap()
19835 .build()
19836 .unwrap();
19837
19838 let record = agent
19839 .invoke_tool(ToolExecutionRequest::new(
19840 "diagnostics-call",
19841 "diagnostics",
19842 serde_json::json!({}),
19843 ToolCallSource::Manual,
19844 ))
19845 .await
19846 .unwrap();
19847
19848 assert!(!record.executed);
19849 assert!(!record.success);
19850 assert_eq!(record.policy.outcome, PermissionOutcome::Unavailable);
19851 }
19852
19853 #[tokio::test]
19854 async fn web_search_without_provider_records_unavailable_without_execution() {
19855 let mock = mock_with_response("hello");
19856 let yaml = r#"
19857name: WebSearchNoProviderAgent
19858system_prompt: "You search the web."
19859tools: [web_search]
19860"#;
19861 let agent = AgentBuilder::from_yaml(yaml)
19862 .unwrap()
19863 .llm(Arc::new(mock))
19864 .auto_configure_features()
19865 .unwrap()
19866 .build()
19867 .unwrap();
19868
19869 let record = agent
19870 .invoke_tool(ToolExecutionRequest::new(
19871 "web-search-call",
19872 "web_search",
19873 serde_json::json!({"query": "rust async"}),
19874 ToolCallSource::Manual,
19875 ))
19876 .await
19877 .unwrap();
19878
19879 assert!(!record.executed);
19880 assert!(!record.success);
19881 assert_eq!(record.policy.outcome, PermissionOutcome::Unavailable);
19882 }
19883
19884 #[tokio::test]
19885 async fn unavailable_host_tool_does_not_request_approval() {
19886 let approvals = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19887 let handler = Arc::new(CountingApprovalHandler {
19888 calls: Arc::clone(&approvals),
19889 });
19890 let mut security = ToolSecurityConfig {
19891 enabled: true,
19892 fail_closed: true,
19893 ..Default::default()
19894 };
19895 security.tools.insert(
19896 "web_search".to_string(),
19897 ai_agents_tools::ToolPolicyConfig {
19898 enabled: true,
19899 require_confirmation: true,
19900 ..Default::default()
19901 },
19902 );
19903 let yaml = r#"
19904name: UnavailableApprovalAgent
19905system_prompt: "Search only with approval."
19906tools: [web_search]
19907"#;
19908 let agent = AgentBuilder::from_yaml(yaml)
19909 .unwrap()
19910 .llm(Arc::new(mock_with_response("done")))
19911 .auto_configure_features()
19912 .unwrap()
19913 .tool_security(ToolSecurityEngine::new(security))
19914 .hitl_engine(HITLEngine::new(ai_agents_hitl::HITLConfig::default()))
19915 .approval_handler(handler)
19916 .build()
19917 .unwrap();
19918
19919 let record = agent
19920 .invoke_tool(ToolExecutionRequest::new(
19921 "unavailable-before-approval",
19922 "web_search",
19923 serde_json::json!({"query": "rust async"}),
19924 ToolCallSource::Manual,
19925 ))
19926 .await
19927 .unwrap();
19928
19929 assert_eq!(approvals.load(Ordering::SeqCst), 0);
19930 assert!(!record.executed);
19931 assert!(!record.success);
19932 assert_eq!(record.policy.outcome, PermissionOutcome::Unavailable);
19933 assert!(
19934 record
19935 .approval
19936 .as_ref()
19937 .is_some_and(|approval| matches!(approval.status, ToolApprovalStatus::Unavailable))
19938 );
19939 }
19940
19941 #[tokio::test]
19942 async fn test_spawner_section_does_not_grant_core_tools_when_top_level_tools_omitted() {
19943 let mock = mock_with_response("hello");
19944 let yaml = r#"
19945name: SpawnerNoGrantAgent
19946system_prompt: "You manage agents."
19947spawner:
19948 max_agents: 2
19949"#;
19950 let agent = AgentBuilder::from_yaml(yaml)
19951 .unwrap()
19952 .llm(Arc::new(mock))
19953 .auto_configure_features()
19954 .unwrap()
19955 .auto_configure_spawner()
19956 .await
19957 .unwrap()
19958 .build()
19959 .unwrap();
19960
19961 let available = agent.get_available_tool_ids().await.unwrap();
19962 assert!(available.is_empty());
19963 }
19964
19965 #[tokio::test]
19966 async fn test_spawner_section_does_not_grant_core_tools_when_top_level_tools_empty() {
19967 let mock = mock_with_response("hello");
19968 let yaml = r#"
19969name: EmptySpawnerNoGrantAgent
19970system_prompt: "You manage agents."
19971tools: []
19972spawner:
19973 max_agents: 2
19974"#;
19975 let agent = AgentBuilder::from_yaml(yaml)
19976 .unwrap()
19977 .llm(Arc::new(mock))
19978 .auto_configure_features()
19979 .unwrap()
19980 .auto_configure_spawner()
19981 .await
19982 .unwrap()
19983 .build()
19984 .unwrap();
19985
19986 let available = agent.get_available_tool_ids().await.unwrap();
19987 assert!(available.is_empty());
19988 }
19989
19990 #[tokio::test]
19991 async fn test_management_tools_flag_grants_core_tools_when_top_level_tools_empty() {
19992 let mock = mock_with_response("hello");
19993 let yaml = r#"
19994name: ManagementGrantAgent
19995system_prompt: "You manage agents."
19996tools: []
19997spawner:
19998 management_tools: true
19999"#;
20000 let agent = AgentBuilder::from_yaml(yaml)
20001 .unwrap()
20002 .llm(Arc::new(mock))
20003 .auto_configure_features()
20004 .unwrap()
20005 .auto_configure_spawner()
20006 .await
20007 .unwrap()
20008 .build()
20009 .unwrap();
20010
20011 let available = agent.get_available_tool_ids().await.unwrap();
20012 assert_eq!(available.len(), 4);
20013 assert!(available.contains(&"spawn_agent".to_string()));
20014 assert!(available.contains(&"send_agent_message".to_string()));
20015 assert!(available.contains(&"list_agents".to_string()));
20016 assert!(available.contains(&"remove_agent".to_string()));
20017 }
20018
20019 #[tokio::test]
20020 async fn test_management_tools_flag_grants_core_tools_when_top_level_tools_omitted() {
20021 let mock = mock_with_response("hello");
20022 let yaml = r#"
20023name: ManagementOmittedToolsGrantAgent
20024system_prompt: "You manage agents."
20025spawner:
20026 management_tools: true
20027"#;
20028 let agent = AgentBuilder::from_yaml(yaml)
20029 .unwrap()
20030 .llm(Arc::new(mock))
20031 .auto_configure_features()
20032 .unwrap()
20033 .auto_configure_spawner()
20034 .await
20035 .unwrap()
20036 .build()
20037 .unwrap();
20038
20039 let available = agent.get_available_tool_ids().await.unwrap();
20040 assert_eq!(available.len(), 4);
20041 assert!(available.contains(&"spawn_agent".to_string()));
20042 assert!(available.contains(&"send_agent_message".to_string()));
20043 assert!(available.contains(&"list_agents".to_string()));
20044 assert!(available.contains(&"remove_agent".to_string()));
20045 }
20046
20047 #[tokio::test]
20048 async fn test_management_tools_selected_grants_only_selected_tools() {
20049 let mock = mock_with_response("hello");
20050 let yaml = r#"
20051name: ManagementSelectedGrantAgent
20052system_prompt: "You manage agents."
20053tools: []
20054spawner:
20055 management_tools:
20056 - spawn_agent
20057 - send_agent_message
20058 - list_agents
20059"#;
20060 let agent = AgentBuilder::from_yaml(yaml)
20061 .unwrap()
20062 .llm(Arc::new(mock))
20063 .auto_configure_features()
20064 .unwrap()
20065 .auto_configure_spawner()
20066 .await
20067 .unwrap()
20068 .build()
20069 .unwrap();
20070
20071 let available = agent.get_available_tool_ids().await.unwrap();
20072 assert_eq!(available.len(), 3);
20073 assert!(available.contains(&"spawn_agent".to_string()));
20074 assert!(available.contains(&"send_agent_message".to_string()));
20075 assert!(available.contains(&"list_agents".to_string()));
20076 assert!(!available.contains(&"remove_agent".to_string()));
20077 }
20078
20079 #[tokio::test]
20080 async fn test_orchestration_tools_flag_grants_tools_when_top_level_tools_empty() {
20081 let mock = mock_with_response("hello");
20082 let yaml = r#"
20083name: OrchestrationGrantAgent
20084system_prompt: "You coordinate agents."
20085llms:
20086 default:
20087 provider: openai
20088 model: gpt-4
20089 router:
20090 provider: openai
20091 model: gpt-4
20092llm:
20093 default: default
20094 router: router
20095tools: []
20096spawner:
20097 orchestration_tools: true
20098"#;
20099 let agent = AgentBuilder::from_yaml(yaml)
20100 .unwrap()
20101 .llm(Arc::new(mock))
20102 .auto_configure_features()
20103 .unwrap()
20104 .auto_configure_spawner()
20105 .await
20106 .unwrap()
20107 .build()
20108 .unwrap();
20109
20110 let available = agent.get_available_tool_ids().await.unwrap();
20111 assert_eq!(available.len(), 5);
20112 assert!(available.contains(&"route_to_agent".to_string()));
20113 assert!(available.contains(&"pipeline_process".to_string()));
20114 assert!(available.contains(&"concurrent_ask".to_string()));
20115 assert!(available.contains(&"group_discussion".to_string()));
20116 assert!(available.contains(&"handoff_conversation".to_string()));
20117 }
20118
20119 #[tokio::test]
20120 async fn test_persona_evolve_flag_grants_tool_when_top_level_tools_empty() {
20121 let mock = mock_with_response("hello");
20122 let yaml = r#"
20123name: PersonaGrantAgent
20124system_prompt: "You can evolve persona."
20125llm:
20126 provider: openai
20127 model: gpt-4
20128tools: []
20129persona:
20130 identity:
20131 name: "Guide"
20132 role: "Helper"
20133 evolution:
20134 enabled: true
20135 allow_llm_evolve: true
20136 mutable_fields:
20137 - traits.personality
20138"#;
20139 let agent = AgentBuilder::from_yaml(yaml)
20140 .unwrap()
20141 .llm(Arc::new(mock))
20142 .build()
20143 .unwrap();
20144
20145 let available = agent.get_available_tool_ids().await.unwrap();
20146 assert_eq!(available, vec!["persona_evolve".to_string()]);
20147 }
20148
20149 #[tokio::test]
20150 async fn test_persona_evolve_flag_grants_tool_when_top_level_tools_omitted() {
20151 let mock = mock_with_response("hello");
20152 let yaml = r#"
20153name: PersonaOmittedToolsGrantAgent
20154system_prompt: "You can evolve persona."
20155llm:
20156 provider: openai
20157 model: gpt-4
20158persona:
20159 identity:
20160 name: "Guide"
20161 role: "Helper"
20162 evolution:
20163 enabled: true
20164 allow_llm_evolve: true
20165 mutable_fields:
20166 - traits.personality
20167"#;
20168 let agent = AgentBuilder::from_yaml(yaml)
20169 .unwrap()
20170 .llm(Arc::new(mock))
20171 .build()
20172 .unwrap();
20173
20174 let available = agent.get_available_tool_ids().await.unwrap();
20175 assert_eq!(available, vec!["persona_evolve".to_string()]);
20176 }
20177
20178 #[tokio::test]
20179 async fn test_omitted_yaml_tools_exposes_no_tools() {
20180 let mock = mock_with_response("hello");
20181 let yaml = r#"
20182name: NoToolsAgent
20183system_prompt: "You are helpful."
20184"#;
20185 let agent = AgentBuilder::from_yaml(yaml)
20186 .unwrap()
20187 .llm(Arc::new(mock))
20188 .auto_configure_features()
20189 .unwrap()
20190 .build()
20191 .unwrap();
20192
20193 let available = agent.get_available_tool_ids().await.unwrap();
20194 assert!(available.is_empty());
20195 }
20196
20197 #[tokio::test]
20198 async fn runtime_scope_cannot_widen_omitted_or_empty_yaml_grants() {
20199 for tools in ["", "tools: []"] {
20200 let yaml = format!(
20201 r#"
20202name: RuntimeScopeNoGrantAgent
20203system_prompt: "No ordinary tools are granted."
20204{tools}
20205"#
20206 );
20207 let agent = AgentBuilder::from_yaml(&yaml)
20208 .unwrap()
20209 .llm(Arc::new(mock_with_response("done")))
20210 .auto_configure_features()
20211 .unwrap()
20212 .build()
20213 .unwrap();
20214
20215 agent
20216 .runtime_control()
20217 .set_tool_scope(vec!["calculator".to_string()]);
20218
20219 assert!(agent.get_available_tool_ids().await.unwrap().is_empty());
20220 }
20221 }
20222
20223 #[tokio::test]
20224 async fn runtime_scope_widening_attempt_keeps_only_declared_tools() {
20225 let yaml = r#"
20226name: RuntimeScopeWideningAgent
20227system_prompt: "Runtime scope cannot add authority."
20228tools: [calculator]
20229"#;
20230 let agent = AgentBuilder::from_yaml(yaml)
20231 .unwrap()
20232 .llm(Arc::new(mock_with_response("done")))
20233 .auto_configure_features()
20234 .unwrap()
20235 .build()
20236 .unwrap();
20237
20238 agent
20239 .runtime_control()
20240 .set_tool_scope(vec!["calculator".to_string(), "datetime".to_string()]);
20241
20242 assert_eq!(
20243 agent.get_available_tool_ids().await.unwrap(),
20244 vec!["calculator".to_string()]
20245 );
20246 }
20247
20248 #[tokio::test]
20249 async fn runtime_scope_is_canonical_unique_ordered_and_clear_restores_declared_grant() {
20250 let yaml = r#"
20251name: RuntimeScopeIntersectionAgent
20252system_prompt: "Use only declared tools."
20253tools: [calculator, datetime]
20254"#;
20255 let agent = AgentBuilder::from_yaml(yaml)
20256 .unwrap()
20257 .llm(Arc::new(mock_with_response("done")))
20258 .auto_configure_features()
20259 .unwrap()
20260 .build()
20261 .unwrap();
20262 let mut aliases = ai_agents_tools::ToolAliases::default();
20263 aliases
20264 .names
20265 .insert("en".to_string(), "calculate_alias".to_string());
20266 agent.tools.set_tool_aliases("calculator", aliases);
20267 let control = agent.runtime_control();
20268
20269 control.set_tool_scope(vec![
20270 "datetime".to_string(),
20271 "calculate_alias".to_string(),
20272 "calculator".to_string(),
20273 "unknown".to_string(),
20274 "datetime".to_string(),
20275 ]);
20276 assert_eq!(
20277 agent.get_available_tool_ids().await.unwrap(),
20278 vec!["calculator".to_string(), "datetime".to_string()]
20279 );
20280
20281 control.set_tool_scope(vec!["datetime".to_string()]);
20282 assert_eq!(
20283 agent.get_available_tool_ids().await.unwrap(),
20284 vec!["datetime".to_string()]
20285 );
20286
20287 control.clear_tool_scope_override();
20288 assert_eq!(
20289 agent.get_available_tool_ids().await.unwrap(),
20290 vec!["calculator".to_string(), "datetime".to_string()]
20291 );
20292 }
20293
20294 #[tokio::test]
20295 async fn runtime_scope_preserves_programmatic_registration_as_declared_grant() {
20296 let agent = AgentBuilder::new()
20297 .system_prompt("Use registered tools.")
20298 .llm(Arc::new(mock_with_response("done")))
20299 .tool(Arc::new(ContextEchoTool))
20300 .tool(Arc::new(SlowTool))
20301 .build()
20302 .unwrap();
20303
20304 agent.runtime_control().set_tool_scope(vec![
20305 "Context Echo".to_string(),
20306 "context_echo".to_string(),
20307 "unknown".to_string(),
20308 ]);
20309
20310 assert_eq!(
20311 agent.get_available_tool_ids().await.unwrap(),
20312 vec!["context_echo".to_string()]
20313 );
20314 }
20315
20316 #[tokio::test]
20317 async fn nested_state_scopes_intersect_every_ancestor_with_aliases() {
20318 let yaml = r#"
20319name: NestedStateScopeAgent
20320system_prompt: "Honor every state scope."
20321tools: [calculator, datetime, echo]
20322states:
20323 initial: root
20324 states:
20325 root:
20326 tools: [calculate_alias, datetime]
20327 initial: middle
20328 states:
20329 middle:
20330 initial: leaf
20331 states:
20332 leaf:
20333 tools: [datetime_alias, echo]
20334"#;
20335 let agent = AgentBuilder::from_yaml(yaml)
20336 .unwrap()
20337 .llm(Arc::new(mock_with_response("done")))
20338 .auto_configure_features()
20339 .unwrap()
20340 .build()
20341 .unwrap();
20342 let mut calculator_aliases = ai_agents_tools::ToolAliases::default();
20343 calculator_aliases
20344 .names
20345 .insert("en".to_string(), "calculate_alias".to_string());
20346 agent
20347 .tools
20348 .set_tool_aliases("calculator", calculator_aliases);
20349 let mut datetime_aliases = ai_agents_tools::ToolAliases::default();
20350 datetime_aliases
20351 .names
20352 .insert("en".to_string(), "datetime_alias".to_string());
20353 agent.tools.set_tool_aliases("datetime", datetime_aliases);
20354 agent.runtime_control().set_tool_scope(vec![
20355 "unknown".to_string(),
20356 "datetime_alias".to_string(),
20357 "calculate_alias".to_string(),
20358 "datetime".to_string(),
20359 ]);
20360
20361 assert_eq!(agent.current_state().as_deref(), Some("root.middle.leaf"));
20362 assert_eq!(
20363 agent.get_available_tool_ids().await.unwrap(),
20364 vec!["datetime".to_string()]
20365 );
20366 }
20367
20368 #[tokio::test]
20369 async fn ancestor_empty_state_scope_denies_omitted_descendants() {
20370 let yaml = r#"
20371name: NestedEmptyStateScopeAgent
20372system_prompt: "An empty ancestor scope denies all tools."
20373tools: [calculator]
20374states:
20375 initial: root
20376 states:
20377 root:
20378 tools: []
20379 initial: middle
20380 states:
20381 middle:
20382 initial: leaf
20383 states:
20384 leaf: {}
20385"#;
20386 let agent = AgentBuilder::from_yaml(yaml)
20387 .unwrap()
20388 .llm(Arc::new(mock_with_response("done")))
20389 .auto_configure_features()
20390 .unwrap()
20391 .build()
20392 .unwrap();
20393
20394 assert!(agent.get_available_tool_ids().await.unwrap().is_empty());
20395 }
20396
20397 #[tokio::test]
20398 async fn state_change_during_approval_invalidates_the_reviewed_authority() {
20399 let active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
20400 let max_active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
20401 let entered = Arc::new(tokio::sync::Barrier::new(2));
20402 let release = Arc::new(tokio::sync::Notify::new());
20403 let handler = Arc::new(BlockingApprovalHandler {
20404 entered: Arc::clone(&entered),
20405 release: Arc::clone(&release),
20406 result: ApprovalResult::Approved,
20407 });
20408 let yaml = r#"
20409name: ApprovalStateGenerationAgent
20410system_prompt: "State authority may change during approval."
20411tools: [locked_write]
20412states:
20413 initial: first
20414 states:
20415 first:
20416 tools: [locked_write]
20417 second:
20418 tools: [locked_write]
20419"#;
20420 let agent = Arc::new(
20421 AgentBuilder::from_yaml(yaml)
20422 .unwrap()
20423 .llm(Arc::new(mock_with_response("done")))
20424 .tool(Arc::new(LockedWriteTool {
20425 active: Arc::clone(&active),
20426 max_active: Arc::clone(&max_active),
20427 }))
20428 .tool_security(ToolSecurityEngine::new(approval_security_config(true)))
20429 .hitl_engine(HITLEngine::new(ai_agents_hitl::HITLConfig::default()))
20430 .approval_handler(handler)
20431 .build()
20432 .unwrap(),
20433 );
20434 let running = Arc::clone(&agent);
20435 let call = tokio::spawn(async move {
20436 running
20437 .invoke_tool(ToolExecutionRequest::new(
20438 "approval-state-generation",
20439 "locked_write",
20440 serde_json::json!({"path": "./state-generation.txt"}),
20441 ToolCallSource::Manual,
20442 ))
20443 .await
20444 .unwrap()
20445 });
20446
20447 entered.wait().await;
20448 agent.transition_to("second").await.unwrap();
20449 release.notify_one();
20450 let record = call.await.unwrap();
20451
20452 assert!(!record.executed);
20453 assert!(record.output.contains("Approval became stale"));
20454 assert_eq!(max_active.load(Ordering::SeqCst), 0);
20455 }
20456
20457 #[tokio::test]
20458 async fn state_change_while_waiting_for_resource_lock_fails_final_admission() {
20459 let holder_gate = PathMutationGate::new();
20460 let waiter_gate = PathMutationGate::new();
20461 let yaml = r#"
20462name: LockedStateGenerationAgent
20463system_prompt: "State authority must remain stable through admission."
20464tools: [state_lock_holder, state_lock_waiter]
20465states:
20466 initial: first
20467 states:
20468 first:
20469 tools: [state_lock_holder, state_lock_waiter]
20470 second:
20471 tools: [state_lock_holder, state_lock_waiter]
20472"#;
20473 let agent = Arc::new(
20474 AgentBuilder::from_yaml(yaml)
20475 .unwrap()
20476 .llm(Arc::new(mock_with_response("done")))
20477 .tool(Arc::new(BlockingPathMutationTool {
20478 id: "state_lock_holder",
20479 path_fields: vec![ai_agents_core::PathPolicyBinding::write("path")],
20480 gate: holder_gate.clone(),
20481 }))
20482 .tool(Arc::new(BlockingPathMutationTool {
20483 id: "state_lock_waiter",
20484 path_fields: vec![ai_agents_core::PathPolicyBinding::write("path")],
20485 gate: waiter_gate.clone(),
20486 }))
20487 .build()
20488 .unwrap(),
20489 );
20490 let holder_call = {
20491 let agent = Arc::clone(&agent);
20492 tokio::spawn(async move {
20493 agent
20494 .invoke_tool(ToolExecutionRequest::new(
20495 "state-lock-holder",
20496 "state_lock_holder",
20497 serde_json::json!({"path": "./shared-state-path.txt"}),
20498 ToolCallSource::Manual,
20499 ))
20500 .await
20501 .unwrap()
20502 })
20503 };
20504 holder_gate.wait_until_entered().await;
20505 let waiter_call = {
20506 let agent = Arc::clone(&agent);
20507 tokio::spawn(async move {
20508 agent
20509 .invoke_tool(ToolExecutionRequest::new(
20510 "state-lock-waiter",
20511 "state_lock_waiter",
20512 serde_json::json!({"path": "./shared-state-path.txt"}),
20513 ToolCallSource::Manual,
20514 ))
20515 .await
20516 .unwrap()
20517 })
20518 };
20519
20520 wait_for_resource_lock_strong_count(&agent.resource_locks, 2).await;
20521 agent.transition_to("second").await.unwrap();
20522 holder_gate.release();
20523 let holder_record = holder_call.await.unwrap();
20524 let waiter_record = waiter_call.await.unwrap();
20525
20526 assert!(holder_record.success);
20527 assert!(!waiter_record.executed);
20528 assert!(
20529 waiter_record
20530 .output
20531 .contains("state scope changed before admission")
20532 );
20533 assert!(!waiter_gate.entered.load(Ordering::SeqCst));
20534 }
20535
20536 #[tokio::test]
20537 async fn test_state_tools_cannot_widen_top_level_grant() {
20538 let mock = mock_with_response("hello");
20539 let yaml = r#"
20540name: NarrowToolsAgent
20541system_prompt: "You are helpful."
20542tools:
20543 - calculator
20544states:
20545 initial: current
20546 states:
20547 current:
20548 tools: [datetime]
20549"#;
20550 let agent = AgentBuilder::from_yaml(yaml)
20551 .unwrap()
20552 .llm(Arc::new(mock))
20553 .auto_configure_features()
20554 .unwrap()
20555 .build()
20556 .unwrap();
20557
20558 let available = agent.get_available_tool_ids().await.unwrap();
20559 assert!(available.is_empty());
20560 }
20561
20562 #[tokio::test]
20564 async fn test_integration_tool_execution() {
20565 let mock = mock_with_responses(vec![
20567 r#"I'll calculate that for you.
20569{"tool": "calculator", "arguments": {"expression": "2+2"}}"#,
20570 "The answer is 4.",
20572 ]);
20573 let observed = mock.clone();
20574 let mut tools = ai_agents_tools::ToolRegistry::new();
20575 tools
20576 .register(Arc::new(ai_agents_tools::CalculatorTool))
20577 .unwrap();
20578
20579 let agent = AgentBuilder::new()
20580 .system_prompt("You are a calculator assistant.")
20581 .llm(Arc::new(mock))
20582 .tools(tools)
20583 .build()
20584 .unwrap();
20585
20586 let response = agent.chat("What is 2+2?").await.unwrap();
20587
20588 assert_eq!(response.content, "The answer is 4.");
20589 assert_eq!(response.tool_calls.as_ref().map(Vec::len), Some(1));
20590 assert_eq!(
20591 observed.call_count(),
20592 2,
20593 "tool result must trigger a second LLM call"
20594 );
20595 let history = agent.tool_call_history();
20596 assert_eq!(history.len(), 1);
20597 assert_eq!(history[0].tool_id, "calculator");
20598 assert_eq!(
20599 history[0].result.get("result"),
20600 Some(&serde_json::json!(4.0)),
20601 "{:?}",
20602 history[0].result
20603 );
20604 }
20605
20606 #[test]
20609 fn legacy_tool_call_marker_is_plain_text() {
20610 let agent = AgentBuilder::new()
20611 .system_prompt("x")
20612 .llm(Arc::new(mock_with_response("x")))
20613 .build()
20614 .unwrap();
20615 let parsed = agent
20616 .parse_tool_calls(
20617 r#"[TOOL_CALL: {"name": "calculator", "arguments": {"expression": "2+2"}}]"#,
20618 )
20619 .unwrap();
20620 assert!(parsed.is_none());
20621 }
20622
20623 #[tokio::test]
20624 async fn test_tool_hitl_rejection_finalizes_blocking_turn() {
20625 let responses = Arc::new(std::sync::atomic::AtomicUsize::new(0));
20626 let hooks = Arc::new(ResponseCountingHooks {
20627 responses: Arc::clone(&responses),
20628 });
20629 let mock = mock_with_response(r#"{"tool":"echo","arguments":{"message":"hello"}}"#);
20630 let yaml = r#"
20631name: ToolRejectAgent
20632system_prompt: "You use tools when requested."
20633tools:
20634 - echo
20635hitl:
20636 tools:
20637 echo:
20638 require_approval: true
20639 approval_message: "Approve echo?"
20640"#;
20641 let agent = AgentBuilder::from_yaml(yaml)
20642 .unwrap()
20643 .llm(Arc::new(mock))
20644 .auto_configure_features()
20645 .unwrap()
20646 .hooks(hooks)
20647 .build()
20648 .unwrap();
20649
20650 let response = agent.chat("echo hello").await.unwrap();
20651
20652 assert!(
20653 response.content.contains("Operation cancelled"),
20654 "unexpected response: {}",
20655 response.content
20656 );
20657 assert_eq!(responses.load(Ordering::SeqCst), 1);
20658 let messages = agent.memory.get_messages(None).await.unwrap();
20659 assert_eq!(messages.len(), 3);
20660 assert_eq!(messages[0].content, "echo hello");
20661 assert!(messages[1].content.contains("\"tool\":\"echo\""));
20662 assert!(messages[2].content.contains("rejected by the approver"));
20663 }
20664
20665 #[tokio::test]
20666 async fn test_tool_hitl_rejection_finalizes_streaming_turn() {
20667 use futures::StreamExt;
20668
20669 let responses = Arc::new(std::sync::atomic::AtomicUsize::new(0));
20670 let hooks = Arc::new(ResponseCountingHooks {
20671 responses: Arc::clone(&responses),
20672 });
20673 let mock = mock_with_response(r#"{"tool":"echo","arguments":{"message":"hello"}}"#);
20674 let yaml = r#"
20675name: ToolRejectStreamingAgent
20676system_prompt: "You use tools when requested."
20677tools:
20678 - echo
20679streaming:
20680 enabled: true
20681hitl:
20682 tools:
20683 echo:
20684 require_approval: true
20685 approval_message: "Approve echo?"
20686"#;
20687 let agent = AgentBuilder::from_yaml(yaml)
20688 .unwrap()
20689 .llm(Arc::new(mock))
20690 .auto_configure_features()
20691 .unwrap()
20692 .hooks(hooks)
20693 .build()
20694 .unwrap();
20695
20696 let mut stream = agent.chat_stream("echo hello").await.unwrap();
20697 let mut terminal_error = String::new();
20698 let mut done = false;
20699 while let Some(chunk) = stream.next().await {
20700 match chunk {
20701 StreamChunk::Error { message } => terminal_error = message,
20702 StreamChunk::Done {} => {
20703 done = true;
20704 break;
20705 }
20706 _ => {}
20707 }
20708 }
20709
20710 assert!(done);
20711 assert!(
20712 terminal_error.contains("Operation cancelled"),
20713 "unexpected terminal error: {}",
20714 terminal_error
20715 );
20716 assert_eq!(responses.load(Ordering::SeqCst), 1);
20717 let messages = agent.memory.get_messages(None).await.unwrap();
20718 assert_eq!(messages.len(), 3);
20719 assert_eq!(messages[0].content, "echo hello");
20720 assert!(messages[1].content.contains("\"tool\":\"echo\""));
20721 assert!(messages[2].content.contains("rejected by the approver"));
20722 }
20723
20724 #[tokio::test]
20725 async fn tool_hitl_rejection_preserves_legacy_error_but_finalizes_event_stream() {
20726 let mock = mock_with_response(r#"{"tool":"echo","arguments":{"message":"hello"}}"#);
20727 let yaml = r#"
20728name: ToolRejectEventAgent
20729system_prompt: "You use tools when requested."
20730tools:
20731 - echo
20732streaming:
20733 enabled: true
20734hitl:
20735 tools:
20736 echo:
20737 require_approval: true
20738 approval_message: "Approve echo?"
20739"#;
20740 let agent = AgentBuilder::from_yaml(yaml)
20741 .unwrap()
20742 .llm(Arc::new(mock))
20743 .auto_configure_features()
20744 .unwrap()
20745 .build()
20746 .unwrap();
20747
20748 let mut stream = agent.chat_stream_events("echo hello").await.unwrap();
20749 let mut error_seen = false;
20750 let mut final_response = None;
20751 while let Some(event) = stream.next().await {
20752 match event {
20753 AgentStreamEvent::Chunk(StreamChunk::Error { .. }) => error_seen = true,
20754 AgentStreamEvent::Final(response) => final_response = Some(response),
20755 AgentStreamEvent::Chunk(_) => {}
20756 }
20757 }
20758
20759 assert!(!error_seen);
20760 assert!(
20761 final_response
20762 .is_some_and(|response| { response.content.contains("Operation cancelled") })
20763 );
20764 }
20765
20766 #[tokio::test]
20767 async fn test_pre_response_guard_transition_skips_old_state_llm() {
20768 let mock = mock_with_response("Billing state response");
20769 let call_counter = mock.clone();
20770 let yaml = r#"
20771name: OptimizedStateAgent
20772system_prompt: "You route before answering."
20773runtime:
20774 optimization:
20775 enabled: true
20776 pre_response_deterministic_transitions: true
20777states:
20778 initial: greeting
20779 states:
20780 greeting:
20781 prompt: "Old state prompt that should be skipped."
20782 transitions:
20783 - to: billing
20784 guard:
20785 context:
20786 topic:
20787 eq: billing
20788 timing: pre_response
20789 billing:
20790 prompt: "Answer from the billing state."
20791"#;
20792 let agent = AgentBuilder::from_yaml(yaml)
20793 .unwrap()
20794 .llm(Arc::new(mock))
20795 .build()
20796 .unwrap();
20797 agent
20798 .set_context("topic", serde_json::json!("billing"))
20799 .unwrap();
20800
20801 let response = agent.chat("I need billing help").await.unwrap();
20802
20803 assert_eq!(agent.current_state().as_deref(), Some("billing"));
20804 assert_eq!(response.content, "Billing state response");
20805 assert_eq!(call_counter.call_count(), 1);
20806 assert_eq!(agent.actor_facts().len(), 0);
20807 }
20808
20809 #[tokio::test]
20810 async fn test_set_context_supports_dotted_paths_for_pre_response_guards() {
20811 let mock = mock_with_response("Billing state response");
20812 let call_counter = mock.clone();
20813 let yaml = r#"
20814name: OptimizedStateAgent
20815system_prompt: "You route before answering."
20816runtime:
20817 optimization:
20818 enabled: true
20819 pre_response_deterministic_transitions: true
20820context:
20821 request:
20822 type: runtime
20823 default:
20824 topic: general
20825states:
20826 initial: greeting
20827 states:
20828 greeting:
20829 prompt: "Old state prompt that should be skipped."
20830 transitions:
20831 - to: billing
20832 guard:
20833 context:
20834 request.topic:
20835 eq: billing
20836 timing: pre_response
20837 billing:
20838 prompt: "Answer from the billing state."
20839"#;
20840 let agent = AgentBuilder::from_yaml(yaml)
20841 .unwrap()
20842 .llm(Arc::new(mock))
20843 .build()
20844 .unwrap();
20845 agent
20846 .set_context("request.topic", serde_json::json!("billing"))
20847 .unwrap();
20848
20849 let response = agent.chat("I need billing help").await.unwrap();
20850
20851 assert_eq!(agent.current_state().as_deref(), Some("billing"));
20852 assert_eq!(response.content, "Billing state response");
20853 assert_eq!(call_counter.call_count(), 1);
20854 assert_eq!(
20855 agent.get_context().get("request"),
20856 Some(&serde_json::json!({"topic": "billing"}))
20857 );
20858 }
20859
20860 #[tokio::test]
20861 async fn test_pre_response_rejection_does_not_commit_staged_context_or_user() {
20862 let mock = mock_with_response("billing");
20863 let yaml = r#"
20864name: OptimizedStateAgent
20865system_prompt: "You route before answering."
20866runtime:
20867 optimization:
20868 enabled: true
20869 pre_response_deterministic_transitions: true
20870hitl:
20871 states:
20872 billing:
20873 on_enter: require_approval
20874 approval_message: "Approve billing route?"
20875states:
20876 initial: greeting
20877 states:
20878 greeting:
20879 prompt: "Old state prompt."
20880 extract:
20881 - key: topic
20882 description: "Support topic"
20883 transitions:
20884 - to: billing
20885 guard:
20886 context:
20887 topic:
20888 eq: billing
20889 timing: pre_response
20890 run_extractors: true
20891 billing:
20892 prompt: "Billing state."
20893"#;
20894 let agent = AgentBuilder::from_yaml(yaml)
20895 .unwrap()
20896 .llm(Arc::new(mock))
20897 .build()
20898 .unwrap();
20899
20900 let response = agent
20901 .try_pre_response_transition("billing please")
20902 .await
20903 .unwrap();
20904
20905 assert!(response.is_none());
20906 assert_eq!(agent.current_state().as_deref(), Some("greeting"));
20907 assert!(!agent.get_context().contains_key("topic"));
20908 assert_eq!(agent.memory.get_messages(None).await.unwrap().len(), 0);
20909 }
20910
20911 #[tokio::test]
20912 async fn test_pre_response_extractor_commits_context_on_winning_path() {
20913 let mock = mock_with_responses(vec!["billing", "Billing response"]);
20914 let yaml = r#"
20915name: OptimizedStateAgent
20916system_prompt: "You route before answering."
20917runtime:
20918 optimization:
20919 enabled: true
20920 pre_response_deterministic_transitions: true
20921states:
20922 initial: greeting
20923 states:
20924 greeting:
20925 prompt: "Old state prompt."
20926 extract:
20927 - key: topic
20928 description: "Support topic"
20929 transitions:
20930 - to: billing
20931 guard:
20932 context:
20933 topic:
20934 eq: billing
20935 timing: pre_response
20936 run_extractors: true
20937 billing:
20938 prompt: "Billing state."
20939"#;
20940 let agent = AgentBuilder::from_yaml(yaml)
20941 .unwrap()
20942 .llm(Arc::new(mock))
20943 .build()
20944 .unwrap();
20945
20946 let response = agent.chat("billing please").await.unwrap();
20947
20948 assert_eq!(agent.current_state().as_deref(), Some("billing"));
20949 assert_eq!(response.content, "Billing response");
20950 assert_eq!(
20951 agent.get_context().get("topic"),
20952 Some(&serde_json::json!("billing"))
20953 );
20954 }
20955
20956 #[tokio::test]
20957 async fn test_pre_response_extractor_miss_does_not_mutate_context() {
20958 let mock = mock_with_response("__NONE__");
20959 let yaml = r#"
20960name: OptimizedStateAgent
20961system_prompt: "You route before answering."
20962runtime:
20963 optimization:
20964 enabled: true
20965 pre_response_deterministic_transitions: true
20966states:
20967 initial: greeting
20968 states:
20969 greeting:
20970 prompt: "Old state prompt."
20971 extract:
20972 - key: topic
20973 description: "Support topic"
20974 transitions:
20975 - to: billing
20976 guard:
20977 context:
20978 topic:
20979 eq: billing
20980 timing: pre_response
20981 run_extractors: true
20982 billing:
20983 prompt: "Billing state."
20984"#;
20985 let agent = AgentBuilder::from_yaml(yaml)
20986 .unwrap()
20987 .llm(Arc::new(mock))
20988 .build()
20989 .unwrap();
20990
20991 let response = agent.try_pre_response_transition("hello").await.unwrap();
20992
20993 assert!(response.is_none());
20994 assert_eq!(agent.current_state().as_deref(), Some("greeting"));
20995 assert!(!agent.get_context().contains_key("topic"));
20996 }
20997
20998 #[tokio::test]
20999 async fn test_default_guard_transition_stays_post_response() {
21000 let mock = mock_with_responses(vec!["Greeting response", "Billing response"]);
21001 let call_counter = mock.clone();
21002 let yaml = r#"
21003name: TimingAgent
21004system_prompt: "You route carefully."
21005runtime:
21006 optimization:
21007 enabled: true
21008 pre_response_deterministic_transitions: true
21009states:
21010 initial: greeting
21011 states:
21012 greeting:
21013 prompt: "Old state prompt."
21014 transitions:
21015 - to: billing
21016 guard:
21017 context:
21018 topic:
21019 eq: billing
21020 billing:
21021 prompt: "Billing state."
21022"#;
21023 let agent = AgentBuilder::from_yaml(yaml)
21024 .unwrap()
21025 .llm(Arc::new(mock))
21026 .build()
21027 .unwrap();
21028 agent
21029 .set_context("topic", serde_json::json!("billing"))
21030 .unwrap();
21031
21032 let response = agent.chat("billing please").await.unwrap();
21033
21034 assert_eq!(agent.current_state().as_deref(), Some("billing"));
21035 assert_eq!(response.content, "Billing response");
21036 assert_eq!(call_counter.call_count(), 2);
21037 }
21038
21039 #[tokio::test]
21040 async fn test_explicit_post_response_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 timing: post_response
21062 billing:
21063 prompt: "Billing state."
21064"#;
21065 let agent = AgentBuilder::from_yaml(yaml)
21066 .unwrap()
21067 .llm(Arc::new(mock))
21068 .build()
21069 .unwrap();
21070 agent
21071 .set_context("topic", serde_json::json!("billing"))
21072 .unwrap();
21073
21074 let response = agent.chat("billing please").await.unwrap();
21075
21076 assert_eq!(agent.current_state().as_deref(), Some("billing"));
21077 assert_eq!(response.content, "Billing response");
21078 assert_eq!(call_counter.call_count(), 2);
21079 }
21080
21081 #[tokio::test]
21082 async fn test_pre_response_extractors_are_transition_scoped() {
21083 let mock = mock_with_responses(vec!["billing", "Billing response"]);
21084 let yaml = r#"
21085name: ScopedExtractorAgent
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 extract:
21097 - key: topic
21098 description: "Support topic"
21099 transitions:
21100 - to: wrong
21101 guard:
21102 context:
21103 topic:
21104 eq: billing
21105 timing: pre_response
21106 - to: billing
21107 guard:
21108 context:
21109 topic:
21110 eq: billing
21111 timing: pre_response
21112 run_extractors: true
21113 wrong:
21114 prompt: "Wrong state."
21115 billing:
21116 prompt: "Billing state."
21117"#;
21118 let agent = AgentBuilder::from_yaml(yaml)
21119 .unwrap()
21120 .llm(Arc::new(mock))
21121 .build()
21122 .unwrap();
21123
21124 let response = agent.chat("billing please").await.unwrap();
21125
21126 assert_eq!(agent.current_state().as_deref(), Some("billing"));
21127 assert_eq!(response.content, "Billing response");
21128 }
21129
21130 #[tokio::test]
21131 async fn test_pre_response_resolved_intent_routes_early() {
21132 let mock = mock_with_response("Billing response");
21133 let yaml = r#"
21134name: IntentAgent
21135system_prompt: "You route carefully."
21136runtime:
21137 optimization:
21138 enabled: true
21139 pre_response_deterministic_transitions: true
21140states:
21141 initial: greeting
21142 states:
21143 greeting:
21144 prompt: "Old state prompt."
21145 transitions:
21146 - to: billing
21147 intent: billing
21148 timing: pre_response
21149 billing:
21150 prompt: "Billing state."
21151"#;
21152 let agent = AgentBuilder::from_yaml(yaml)
21153 .unwrap()
21154 .llm(Arc::new(mock))
21155 .build()
21156 .unwrap();
21157 agent
21158 .set_context("resolved_intent", serde_json::json!("billing"))
21159 .unwrap();
21160
21161 let response = agent
21162 .try_pre_response_transition("I need billing help")
21163 .await
21164 .unwrap()
21165 .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_background_overflow_error_surfaces() {
21173 let mut config = RuntimeConfig::default();
21174 config.optimization.enabled = true;
21175 config.optimization.post_turn.max_background_tasks = 1;
21176 config.optimization.post_turn.on_background_overflow = BackgroundOverflowPolicy::Error;
21177 let policy = crate::optimization::MaintenanceTaskPolicy {
21178 mode: MaintenanceMode::Background,
21179 await_before_next_turn: AwaitBeforeNextTurn::Always,
21180 };
21181 let agent = AgentBuilder::new()
21182 .system_prompt("You are helpful.")
21183 .llm(Arc::new(mock_with_response("ok")))
21184 .build()
21185 .unwrap()
21186 .with_runtime_config(config);
21187 agent
21188 .background_maintenance
21189 .spawn(None, async { std::future::pending::<Result<()>>().await })
21190 .unwrap();
21191
21192 let result = agent
21193 .spawn_or_handle_background(None, async { Ok(()) }, "facts", &policy)
21194 .await;
21195
21196 assert!(result.is_err());
21197 }
21198
21199 #[tokio::test]
21200 async fn test_speculative_reasoning_low_cap_uses_serial_reasoning() {
21201 let default_mock = mock_with_response("Plain draft response");
21202 let router_mock = mock_with_response("cot");
21203 let router_counter = router_mock.clone();
21204 let yaml = r#"
21205name: ReasoningReservationAgent
21206system_prompt: "You answer plainly unless reasoning wins."
21207llm:
21208 default: default
21209 router: router
21210observability:
21211 enabled: true
21212 export:
21213 write_raw_events: true
21214reasoning:
21215 mode: auto
21216 judge_llm: router
21217runtime:
21218 optimization:
21219 enabled: true
21220 max_speculative_llm_calls_per_turn: 1
21221 speculative_reasoning_auto: true
21222 max_parallel_runtime_tasks: 2
21223"#;
21224 let agent = AgentBuilder::from_yaml(yaml)
21225 .unwrap()
21226 .llm_alias("default", Arc::new(default_mock))
21227 .llm_alias("router", Arc::new(router_mock))
21228 .build()
21229 .unwrap();
21230
21231 let response = agent.chat("hello").await.unwrap();
21232
21233 assert_eq!(response.content, "Plain draft response");
21234 assert_eq!(router_counter.call_count(), 1);
21235 let events = agent.observability().unwrap().raw_events();
21236 assert!(!events.iter().any(|event| {
21237 event.dimensions.get("commit_behavior") == Some(&"reasoning_decision".to_string())
21238 }));
21239 }
21240
21241 #[tokio::test]
21242 async fn test_forced_reasoning_skips_plain_speculative_draft() {
21243 let mock = mock_with_response("Reasoned response");
21244 let yaml = r#"
21245name: ForcedReasoningAgent
21246system_prompt: "You reason before answering."
21247observability:
21248 enabled: true
21249 export:
21250 write_raw_events: true
21251reasoning:
21252 mode: cot
21253runtime:
21254 optimization:
21255 enabled: true
21256 max_speculative_llm_calls_per_turn: 2
21257 speculative_state_transitions: true
21258 max_parallel_runtime_tasks: 2
21259states:
21260 initial: triage
21261 states:
21262 triage:
21263 prompt: "Answer from triage."
21264 transitions:
21265 - to: billing
21266 guard:
21267 context:
21268 route:
21269 eq: billing
21270 timing: parallel
21271 billing:
21272 prompt: "Billing state."
21273"#;
21274 let agent = AgentBuilder::from_yaml(yaml)
21275 .unwrap()
21276 .llm(Arc::new(mock))
21277 .build()
21278 .unwrap();
21279
21280 let response = agent.chat("hello").await.unwrap();
21281
21282 assert_eq!(response.content, "Reasoned response");
21283 let events = agent.observability().unwrap().raw_events();
21284 assert!(
21285 !events
21286 .iter()
21287 .any(|event| event.dimensions.contains_key("branch_status"))
21288 );
21289 }
21290
21291 #[tokio::test]
21292 async fn test_speculative_skill_low_cap_uses_serial_skill_route() {
21293 let default_mock = mock_with_response("Skill committed response");
21294 let router_mock = mock_with_response("helper");
21295 let router_counter = router_mock.clone();
21296 let yaml = r#"
21297name: SkillReservationAgent
21298system_prompt: "Use skills when they match."
21299llm:
21300 default: default
21301 router: router
21302observability:
21303 enabled: true
21304 export:
21305 write_raw_events: true
21306runtime:
21307 optimization:
21308 enabled: true
21309 max_speculative_llm_calls_per_turn: 1
21310 speculative_skill_routing: true
21311 max_parallel_runtime_tasks: 2
21312skills:
21313 - id: helper
21314 description: "Answer helper requests"
21315 trigger: "User asks for helper"
21316 steps:
21317 - prompt: "Answer the helper request: {{ user_input }}"
21318"#;
21319 let agent = AgentBuilder::from_yaml(yaml)
21320 .unwrap()
21321 .llm_alias("default", Arc::new(default_mock))
21322 .llm_alias("router", Arc::new(router_mock))
21323 .build()
21324 .unwrap();
21325
21326 let response = agent.chat("please use helper").await.unwrap();
21327
21328 assert_eq!(response.content, "Skill committed response");
21329 assert_eq!(router_counter.call_count(), 1);
21330 let events = agent.observability().unwrap().raw_events();
21331 assert!(
21332 !events
21333 .iter()
21334 .any(|event| event.dimensions.contains_key("branch_status"))
21335 );
21336 }
21337
21338 #[tokio::test]
21339 async fn test_parallel_transition_low_cap_allows_deterministic_route() {
21340 let mock = mock_with_response("unused");
21341 let call_counter = mock.clone();
21342 let yaml = r#"
21343name: ParallelTransitionLowCapAgent
21344system_prompt: "Route before stale responses when safe."
21345runtime:
21346 optimization:
21347 enabled: true
21348 max_speculative_llm_calls_per_turn: 1
21349 speculative_state_transitions: true
21350 max_parallel_runtime_tasks: 2
21351states:
21352 initial: triage
21353 states:
21354 triage:
21355 prompt: "Triage state."
21356 transitions:
21357 - to: billing
21358 guard:
21359 context:
21360 route:
21361 eq: billing
21362 timing: parallel
21363 billing:
21364 prompt: "Billing state."
21365"#;
21366 let agent = AgentBuilder::from_yaml(yaml)
21367 .unwrap()
21368 .llm(Arc::new(mock))
21369 .build()
21370 .unwrap();
21371 agent
21372 .set_context("route", serde_json::json!("billing"))
21373 .unwrap();
21374 agent.update_active_turn_context("billing help", HashMap::new());
21375 assert!(
21376 agent.reserve_active_speculative_llm_call(
21377 RuntimeOptimizationKind::ParallelStateTransition
21378 )
21379 );
21380
21381 let selection = agent
21382 .select_parallel_transition_candidate("billing help")
21383 .await
21384 .unwrap();
21385 agent.end_root_turn();
21386
21387 match selection {
21388 ParallelTransitionSelection::Candidate(candidate) => {
21389 assert_eq!(candidate.target(), "billing");
21390 }
21391 ParallelTransitionSelection::NoMatch => panic!("deterministic route did not match"),
21392 ParallelTransitionSelection::ReservationExhausted => {
21393 panic!("deterministic route consumed LLM budget")
21394 }
21395 }
21396 assert_eq!(call_counter.call_count(), 0);
21397 }
21398
21399 #[tokio::test]
21400 async fn speculative_transition_drops_loser_before_state_actions() {
21401 let lock = Arc::new(tokio::sync::Mutex::new(()));
21402 let first_started = Arc::new(tokio::sync::Notify::new());
21403 let first_dropped = Arc::new(AtomicBool::new(false));
21404 let committed_after_drop = Arc::new(AtomicBool::new(false));
21405 let default = Arc::new(FirstCallLockingProvider {
21406 lock,
21407 first_started: Arc::clone(&first_started),
21408 first_dropped: Arc::clone(&first_dropped),
21409 committed_after_drop: Arc::clone(&committed_after_drop),
21410 calls: AtomicU64::new(0),
21411 });
21412 let router = Arc::new(RoutingAfterProviderStart {
21413 provider_started: first_started,
21414 });
21415 let yaml = r#"
21416name: SpeculativeCancellationAgent
21417system_prompt: "Route before committed work."
21418llm:
21419 default: default
21420 router: router
21421runtime:
21422 optimization:
21423 enabled: true
21424 max_speculative_llm_calls_per_turn: 2
21425 speculative_state_transitions: true
21426 max_parallel_runtime_tasks: 2
21427states:
21428 initial: triage
21429 states:
21430 triage:
21431 prompt: "Triage state."
21432 transitions:
21433 - to: technical
21434 when: "The request needs technical support"
21435 timing: parallel
21436 technical:
21437 prompt: "Technical state."
21438 on_enter:
21439 - prompt: "Prepare technical context."
21440 llm: default
21441 store_as: preparation
21442"#;
21443 let agent = AgentBuilder::from_yaml(yaml)
21444 .unwrap()
21445 .llm_alias("default", default)
21446 .llm_alias("router", router)
21447 .build()
21448 .unwrap();
21449
21450 let response = tokio::time::timeout(
21451 std::time::Duration::from_secs(2),
21452 agent.chat("I cannot log in because of AUTH-17."),
21453 )
21454 .await
21455 .expect("committed work must not wait on the losing provider future")
21456 .unwrap();
21457
21458 assert_eq!(response.content, "Committed technical response.");
21459 assert_eq!(agent.current_state().as_deref(), Some("technical"));
21460 assert!(first_dropped.load(Ordering::SeqCst));
21461 assert!(committed_after_drop.load(Ordering::SeqCst));
21462 }
21463
21464 #[tokio::test]
21465 async fn buffered_transition_drops_stale_stream_before_redispatch() {
21466 use futures::StreamExt;
21467
21468 let lock = Arc::new(tokio::sync::Mutex::new(()));
21469 let stream_started = Arc::new(tokio::sync::Notify::new());
21470 let stream_dropped = Arc::new(AtomicBool::new(false));
21471 let committed_after_drop = Arc::new(AtomicBool::new(false));
21472 let default = Arc::new(BufferedLockingProvider {
21473 lock,
21474 stream_started: Arc::clone(&stream_started),
21475 stream_dropped: Arc::clone(&stream_dropped),
21476 committed_after_drop: Arc::clone(&committed_after_drop),
21477 });
21478 let router = Arc::new(RoutingAfterProviderStart {
21479 provider_started: stream_started,
21480 });
21481 let yaml = r#"
21482name: BufferedCancellationAgent
21483system_prompt: "Hide stale streamed output."
21484llm:
21485 default: default
21486 router: router
21487streaming:
21488 enabled: true
21489 buffer_size: 8
21490runtime:
21491 optimization:
21492 enabled: true
21493 max_speculative_llm_calls_per_turn: 2
21494 speculative_state_transitions: true
21495 streaming_policy: buffer_until_routing_done
21496 max_parallel_runtime_tasks: 2
21497states:
21498 initial: triage
21499 states:
21500 triage:
21501 prompt: "Triage state."
21502 transitions:
21503 - to: technical
21504 when: "The request needs technical support"
21505 timing: parallel
21506 technical:
21507 prompt: "Technical state."
21508"#;
21509 let agent = AgentBuilder::from_yaml(yaml)
21510 .unwrap()
21511 .llm_alias("default", default)
21512 .llm_alias("router", router)
21513 .build()
21514 .unwrap();
21515
21516 let content = tokio::time::timeout(std::time::Duration::from_secs(2), async {
21517 let mut stream = agent
21518 .chat_stream("AUTH-17 needs technical help.")
21519 .await
21520 .unwrap();
21521 let mut content = String::new();
21522 while let Some(chunk) = stream.next().await {
21523 match chunk {
21524 StreamChunk::Content { text } => content.push_str(&text),
21525 StreamChunk::Done {} => break,
21526 StreamChunk::Error { message } => panic!("unexpected stream error: {message}"),
21527 _ => {}
21528 }
21529 }
21530 content
21531 })
21532 .await
21533 .expect("redispatch must not wait on the stale streaming future");
21534
21535 assert_eq!(content, "Committed technical response.");
21536 assert_eq!(agent.current_state().as_deref(), Some("technical"));
21537 assert!(stream_dropped.load(Ordering::SeqCst));
21538 assert!(committed_after_drop.load(Ordering::SeqCst));
21539 }
21540
21541 #[tokio::test]
21542 async fn buffered_transition_drops_established_stream_before_redispatch() {
21543 use futures::StreamExt;
21544
21545 let stream_started = Arc::new(tokio::sync::Notify::new());
21546 let stream_dropped = Arc::new(AtomicBool::new(false));
21547 let stream_dropped_notify = Arc::new(tokio::sync::Notify::new());
21548 let committed_after_drop = Arc::new(AtomicBool::new(false));
21549 let default = Arc::new(EstablishedStreamProvider {
21550 stream_started: Arc::clone(&stream_started),
21551 stream_dropped: Arc::clone(&stream_dropped),
21552 stream_dropped_notify,
21553 committed_after_drop: Arc::clone(&committed_after_drop),
21554 });
21555 let router = Arc::new(RoutingAfterProviderStart {
21556 provider_started: stream_started,
21557 });
21558 let yaml = r#"
21559name: EstablishedStreamCancellationAgent
21560system_prompt: "Hide stale streamed output."
21561llm:
21562 default: default
21563 router: router
21564streaming:
21565 enabled: true
21566 buffer_size: 8
21567runtime:
21568 optimization:
21569 enabled: true
21570 max_speculative_llm_calls_per_turn: 2
21571 speculative_state_transitions: true
21572 streaming_policy: buffer_until_routing_done
21573 max_parallel_runtime_tasks: 2
21574states:
21575 initial: triage
21576 states:
21577 triage:
21578 prompt: "Triage state."
21579 transitions:
21580 - to: technical
21581 when: "The request needs technical support"
21582 timing: parallel
21583 technical:
21584 prompt: "Technical state."
21585"#;
21586 let agent = AgentBuilder::from_yaml(yaml)
21587 .unwrap()
21588 .llm_alias("default", default)
21589 .llm_alias("router", router)
21590 .build()
21591 .unwrap();
21592
21593 let content = tokio::time::timeout(std::time::Duration::from_secs(2), async {
21594 let mut stream = agent
21595 .chat_stream("AUTH-17 needs technical help.")
21596 .await
21597 .unwrap();
21598 let mut content = String::new();
21599 while let Some(chunk) = stream.next().await {
21600 match chunk {
21601 StreamChunk::Content { text } => content.push_str(&text),
21602 StreamChunk::Done {} => break,
21603 StreamChunk::Error { message } => panic!("unexpected stream error: {message}"),
21604 _ => {}
21605 }
21606 }
21607 content
21608 })
21609 .await
21610 .expect("redispatch must wait for the established stale stream to be dropped");
21611
21612 assert_eq!(content, "Committed technical response.");
21613 assert_eq!(agent.current_state().as_deref(), Some("technical"));
21614 assert!(stream_dropped.load(Ordering::SeqCst));
21615 assert!(committed_after_drop.load(Ordering::SeqCst));
21616 }
21617
21618 #[tokio::test]
21619 async fn test_buffered_streaming_transition_reservation_falls_back() {
21620 use futures::StreamExt;
21621
21622 let mock = mock_with_responses(vec![
21623 "Serial streaming response",
21624 "Serial streaming response",
21625 ]);
21626 let router_mock = mock_with_response("1");
21627 let router_counter = router_mock.clone();
21628 let yaml = r#"
21629name: BufferedReservationFallbackAgent
21630system_prompt: "Stream normally if speculative routing cannot be evaluated."
21631llm:
21632 default: default
21633 router: router
21634observability:
21635 enabled: true
21636 export:
21637 write_raw_events: true
21638streaming:
21639 enabled: true
21640 buffer_size: 8
21641runtime:
21642 optimization:
21643 enabled: true
21644 max_speculative_llm_calls_per_turn: 1
21645 speculative_state_transitions: true
21646 streaming_policy: buffer_until_routing_done
21647 max_parallel_runtime_tasks: 2
21648states:
21649 initial: triage
21650 states:
21651 triage:
21652 prompt: "Triage state."
21653 transitions:
21654 - to: billing
21655 guard:
21656 context:
21657 route:
21658 eq: billing
21659 when: "User asks about billing"
21660 timing: parallel
21661 billing:
21662 prompt: "Billing state."
21663"#;
21664 let agent = AgentBuilder::from_yaml(yaml)
21665 .unwrap()
21666 .llm_alias("default", Arc::new(mock))
21667 .llm_alias("router", Arc::new(router_mock))
21668 .build()
21669 .unwrap();
21670
21671 let mut stream = agent.chat_stream("hello").await.unwrap();
21672 let mut content = String::new();
21673 let mut error = None;
21674 while let Some(chunk) = stream.next().await {
21675 match chunk {
21676 StreamChunk::Content { text } => content.push_str(&text),
21677 StreamChunk::Error { message } => error = Some(message),
21678 StreamChunk::Done {} => break,
21679 _ => {}
21680 }
21681 }
21682
21683 assert_eq!(error, None);
21684 assert_eq!(content, "Serial streaming response");
21685 assert_eq!(router_counter.call_count(), 0);
21686 let events = agent.observability().unwrap().raw_events();
21687 assert!(events.iter().any(|event| {
21688 event.dimensions.get("branch_status") == Some(&"cancelled".to_string())
21689 && event.dimensions.get("commit_behavior")
21690 == Some(&"transition_decision".to_string())
21691 }));
21692 }
21693
21694 #[tokio::test]
21695 async fn test_blocking_error_cleanup_resets_root_turn_for_next_chat() {
21696 let mut mock = mock_with_response("Recovered response");
21697 mock.set_error("boom");
21698 let mut handle = mock.clone();
21699 let agent = AgentBuilder::new()
21700 .system_prompt("You are helpful.")
21701 .llm(Arc::new(mock))
21702 .build()
21703 .unwrap();
21704
21705 assert!(agent.chat("first").await.is_err());
21706 handle.clear_error();
21707 let response = agent.chat("second").await.unwrap();
21708
21709 assert_eq!(response.content, "Recovered response");
21710 let messages = agent.memory.get_messages(None).await.unwrap();
21711 let user_count = messages
21712 .iter()
21713 .filter(|message| message.role == ai_agents_core::Role::User)
21714 .count();
21715 assert_eq!(user_count, 2);
21716 }
21717
21718 #[tokio::test]
21719 async fn test_streaming_error_cleanup_resets_root_turn_for_next_chat() {
21720 use futures::StreamExt;
21721
21722 let mut mock = mock_with_response("Recovered response");
21723 mock.set_error("stream boom");
21724 let mut handle = mock.clone();
21725 let agent = AgentBuilder::new()
21726 .system_prompt("You are helpful.")
21727 .llm(Arc::new(mock))
21728 .build()
21729 .unwrap();
21730
21731 let mut stream = agent.chat_stream("first").await.unwrap();
21732 let mut saw_error = false;
21733 while let Some(chunk) = stream.next().await {
21734 if matches!(chunk, StreamChunk::Error { .. }) {
21735 saw_error = true;
21736 }
21737 }
21738 assert!(saw_error);
21739
21740 handle.clear_error();
21741 let response = agent.chat("second").await.unwrap();
21742
21743 assert_eq!(response.content, "Recovered response");
21744 let messages = agent.memory.get_messages(None).await.unwrap();
21745 let user_count = messages
21746 .iter()
21747 .filter(|message| message.role == ai_agents_core::Role::User)
21748 .count();
21749 assert_eq!(user_count, 2);
21750 }
21751
21752 #[tokio::test]
21753 async fn test_buffered_streaming_route_miss_releases_buffer_limit() {
21754 use futures::StreamExt;
21755
21756 let mut mock = mock_with_response("one two three");
21757 mock.set_latency(10);
21758 let yaml = r#"
21759name: BufferedMissAgent
21760system_prompt: "You stream safely."
21761llm:
21762 default: default
21763streaming:
21764 enabled: true
21765 buffer_size: 1
21766runtime:
21767 optimization:
21768 enabled: true
21769 max_speculative_llm_calls_per_turn: 2
21770 speculative_state_transitions: true
21771 streaming_policy: buffer_until_routing_done
21772 max_parallel_runtime_tasks: 2
21773states:
21774 initial: triage
21775 states:
21776 triage:
21777 prompt: "Answer from triage."
21778 transitions:
21779 - to: billing
21780 guard:
21781 context:
21782 route:
21783 eq: billing
21784 timing: parallel
21785 billing:
21786 prompt: "Billing state."
21787"#;
21788 let agent = AgentBuilder::from_yaml(yaml)
21789 .unwrap()
21790 .llm_alias("default", Arc::new(mock))
21791 .build()
21792 .unwrap();
21793
21794 let mut stream = agent.chat_stream("hello").await.unwrap();
21795 let mut content = String::new();
21796 let mut error = None;
21797 while let Some(chunk) = stream.next().await {
21798 match chunk {
21799 StreamChunk::Content { text } => content.push_str(&text),
21800 StreamChunk::Error { message } => error = Some(message),
21801 StreamChunk::Done {} => break,
21802 _ => {}
21803 }
21804 }
21805
21806 assert_eq!(error, None);
21807 assert_eq!(content, "one two three");
21808 }
21809
21810 #[tokio::test]
21811 async fn test_buffered_streaming_main_failure_finalizes_branch() {
21812 use futures::StreamExt;
21813
21814 let mock = mock_with_response("one two");
21815 let mut router_mock = mock_with_response("0");
21816 router_mock.set_latency(50);
21817 let yaml = r#"
21818name: BufferedFailureAgent
21819system_prompt: "You stream safely."
21820llm:
21821 default: default
21822 router: router
21823observability:
21824 enabled: true
21825 export:
21826 write_raw_events: true
21827streaming:
21828 enabled: true
21829 buffer_size: 1
21830runtime:
21831 optimization:
21832 enabled: true
21833 max_speculative_llm_calls_per_turn: 2
21834 speculative_state_transitions: true
21835 streaming_policy: buffer_until_routing_done
21836 max_parallel_runtime_tasks: 2
21837states:
21838 initial: triage
21839 states:
21840 triage:
21841 prompt: "Ask for the category."
21842 transitions:
21843 - to: billing
21844 when: "User asks about billing"
21845 timing: parallel
21846 billing:
21847 prompt: "Billing state."
21848"#;
21849 let agent = AgentBuilder::from_yaml(yaml)
21850 .unwrap()
21851 .llm_alias("default", Arc::new(mock))
21852 .llm_alias("router", Arc::new(router_mock))
21853 .build()
21854 .unwrap();
21855
21856 let mut stream = agent.chat_stream("hello").await.unwrap();
21857 let mut error = String::new();
21858 while let Some(chunk) = stream.next().await {
21859 if let StreamChunk::Error { message } = chunk {
21860 error = message;
21861 }
21862 }
21863
21864 assert!(
21865 error.contains("stream buffer filled"),
21866 "unexpected stream error: {}",
21867 error
21868 );
21869 let events = agent.observability().unwrap().raw_events();
21870 assert!(events.iter().any(|event| {
21871 event.dimensions.get("branch_status") == Some(&"failed".to_string())
21872 && event.dimensions.get("commit_behavior") == Some(&"final_response".to_string())
21873 && event.dimensions.get("optimization")
21874 == Some(&"buffered_streaming_routing".to_string())
21875 }));
21876 }
21877
21878 #[tokio::test]
21879 async fn test_streaming_preflight_does_not_emit_old_state_content() {
21880 use futures::StreamExt;
21881
21882 let mock = mock_with_response("Billing streamed response");
21883 let yaml = r#"
21884name: StreamingOptimizedAgent
21885system_prompt: "You route before streaming."
21886runtime:
21887 optimization:
21888 enabled: true
21889 pre_response_deterministic_transitions: true
21890streaming:
21891 enabled: true
21892states:
21893 initial: greeting
21894 states:
21895 greeting:
21896 prompt: "OLD_STATE_SENTINEL"
21897 transitions:
21898 - to: billing
21899 guard:
21900 context:
21901 topic:
21902 eq: billing
21903 timing: pre_response
21904 billing:
21905 prompt: "Billing state."
21906"#;
21907 let agent = AgentBuilder::from_yaml(yaml)
21908 .unwrap()
21909 .llm(Arc::new(mock))
21910 .build()
21911 .unwrap();
21912 agent
21913 .set_context("topic", serde_json::json!("billing"))
21914 .unwrap();
21915
21916 let mut stream = agent.chat_stream("billing please").await.unwrap();
21917 let mut content = String::new();
21918 while let Some(chunk) = stream.next().await {
21919 match chunk {
21920 StreamChunk::Content { text } => content.push_str(&text),
21921 StreamChunk::Error { message } => panic!("stream error: {}", message),
21922 StreamChunk::Done {} => break,
21923 _ => {}
21924 }
21925 }
21926
21927 assert_eq!(agent.current_state().as_deref(), Some("billing"));
21928 assert!(content.contains("Billing streamed response"));
21929 assert!(!content.contains("OLD_STATE_SENTINEL"));
21930 }
21931
21932 #[tokio::test]
21934 async fn test_integration_state_machine_basic() {
21935 let yaml = r#"
21936name: StateAgent
21937system_prompt: "You are a support agent."
21938states:
21939 initial: greeting
21940 states:
21941 greeting:
21942 prompt: "Welcome the user warmly."
21943 transitions:
21944 - to: support
21945 when: "User needs help"
21946 auto: true
21947 support:
21948 prompt: "Help solve the user's problem."
21949"#;
21950 let mock = mock_with_responses(vec![
21951 "Welcome! How can I help?", "1", "I'll help you with that.", ]);
21955 let builder = AgentBuilder::from_yaml(yaml).unwrap();
21956 let agent = builder.llm(Arc::new(mock)).build().unwrap();
21957
21958 assert_eq!(agent.current_state(), Some("greeting".to_string()));
21959 let _ = agent.chat("I need help").await.unwrap();
21960 }
21963
21964 #[tokio::test]
21966 async fn test_integration_state_on_enter_set_context() {
21967 let yaml = r#"
21968name: ActionAgent
21969system_prompt: "You are helpful."
21970states:
21971 initial: step1
21972 states:
21973 step1:
21974 prompt: "Step 1"
21975 on_exit:
21976 - set_context:
21977 step1_exited: true
21978 transitions:
21979 - to: step2
21980 when: "always"
21981 auto: true
21982 step2:
21983 prompt: "Step 2"
21984 on_enter:
21985 - set_context:
21986 step2_entered: true
21987"#;
21988 let mock = mock_with_responses(vec![
21990 "Processing step 1.",
21991 "0", ]);
21993 let builder = AgentBuilder::from_yaml(yaml).unwrap();
21994 let agent = builder.llm(Arc::new(mock)).build().unwrap();
21995
21996 assert_eq!(agent.current_state(), Some("step1".to_string()));
21997
21998 agent.transition_to("step2").await.unwrap();
22000
22001 assert_eq!(agent.current_state(), Some("step2".to_string()));
22002
22003 let ctx = agent.get_context();
22005 assert_eq!(ctx.get("step1_exited"), Some(&serde_json::json!(true)));
22006 assert_eq!(ctx.get("step2_entered"), Some(&serde_json::json!(true)));
22007 }
22008
22009 #[tokio::test]
22010 async fn state_action_tool_preserves_source_in_stored_record() {
22011 let yaml = r#"
22012name: StateActionToolAgent
22013system_prompt: "You are helpful."
22014tools:
22015 - context_echo
22016states:
22017 initial: idle
22018 states:
22019 idle:
22020 prompt: "Idle"
22021 active:
22022 prompt: "Active"
22023 on_enter:
22024 - set_context:
22025 action_started: true
22026 - tool: context_echo
22027 args: {}
22028"#;
22029 let agent = AgentBuilder::from_yaml(yaml)
22030 .unwrap()
22031 .llm(Arc::new(mock_with_response("unused")))
22032 .tool(Arc::new(ContextEchoTool))
22033 .build()
22034 .unwrap();
22035
22036 agent.transition_to("active").await.unwrap();
22037
22038 let record: ToolExecutionRecord = serde_json::from_value(
22039 agent
22040 .get_context()
22041 .get("last_tool_record")
22042 .cloned()
22043 .expect("successful state action must store its execution record"),
22044 )
22045 .unwrap();
22046 assert!(record.executed);
22047 assert!(record.success);
22048 assert_eq!(record.canonical_id, "context_echo");
22049 assert!(matches!(
22050 &record.source,
22051 ToolCallSource::StateAction {
22052 state: Some(state),
22053 action_index: 1,
22054 } if state == "active"
22055 ));
22056 }
22057
22058 #[tokio::test]
22059 async fn test_ordinary_transition_uses_on_enter_then_on_reenter() {
22060 let yaml = r#"
22061name: OrdinaryLifecycleAgent
22062system_prompt: "You are helpful."
22063states:
22064 initial: intake
22065 regenerate_on_transition: false
22066 states:
22067 intake:
22068 prompt: "Intake"
22069 transitions:
22070 - to: drafting
22071 guard:
22072 context:
22073 route:
22074 eq: drafting
22075 drafting:
22076 prompt: "Drafting"
22077 on_enter:
22078 - set_context:
22079 draft_version: 1
22080 on_reenter:
22081 - set_context:
22082 draft_version: 2
22083 transitions:
22084 - to: review
22085 guard:
22086 context:
22087 route:
22088 eq: review
22089 review:
22090 prompt: "Review"
22091 on_enter:
22092 - set_context:
22093 review_entry: first
22094 transitions:
22095 - to: drafting
22096 guard:
22097 context:
22098 route:
22099 eq: drafting
22100"#;
22101 let agent = AgentBuilder::from_yaml(yaml)
22102 .unwrap()
22103 .llm(Arc::new(mock_with_responses(vec![
22104 "Intake response",
22105 "Draft response",
22106 "Review response",
22107 ])))
22108 .build()
22109 .unwrap();
22110
22111 agent
22112 .set_context("route", serde_json::json!("drafting"))
22113 .unwrap();
22114 agent.chat("Start a draft").await.unwrap();
22115 assert_eq!(agent.current_state().as_deref(), Some("drafting"));
22116 assert_eq!(
22117 agent.get_context().get("draft_version"),
22118 Some(&serde_json::json!(1))
22119 );
22120
22121 agent
22122 .set_context("route", serde_json::json!("review"))
22123 .unwrap();
22124 agent.chat("Review this").await.unwrap();
22125 assert_eq!(agent.current_state().as_deref(), Some("review"));
22126 assert_eq!(
22127 agent.get_context().get("review_entry"),
22128 Some(&serde_json::json!("first"))
22129 );
22130
22131 agent
22132 .set_context("route", serde_json::json!("drafting"))
22133 .unwrap();
22134 agent.chat("Revise this").await.unwrap();
22135 assert_eq!(agent.current_state().as_deref(), Some("drafting"));
22136 assert_eq!(
22137 agent.get_context().get("draft_version"),
22138 Some(&serde_json::json!(2))
22139 );
22140 }
22141
22142 #[tokio::test]
22143 async fn test_manual_transition_uses_on_enter_then_on_reenter() {
22144 let yaml = r#"
22145name: ManualLifecycleAgent
22146system_prompt: "You are helpful."
22147states:
22148 initial: intake
22149 states:
22150 intake:
22151 prompt: "Intake"
22152 drafting:
22153 prompt: "Drafting"
22154 on_enter:
22155 - set_context:
22156 draft_version: 1
22157 on_reenter:
22158 - set_context:
22159 draft_version: 2
22160 review:
22161 prompt: "Review"
22162"#;
22163 let agent = AgentBuilder::from_yaml(yaml)
22164 .unwrap()
22165 .llm(Arc::new(mock_with_response("unused")))
22166 .build()
22167 .unwrap();
22168
22169 assert!(!agent.get_context().contains_key("draft_version"));
22170 agent.transition_to("drafting").await.unwrap();
22171 assert_eq!(agent.current_state().as_deref(), Some("drafting"));
22172 assert_eq!(
22173 agent.get_context().get("draft_version"),
22174 Some(&serde_json::json!(1))
22175 );
22176
22177 agent.transition_to("review").await.unwrap();
22178 agent.transition_to("drafting").await.unwrap();
22179 assert_eq!(agent.current_state().as_deref(), Some("drafting"));
22180 assert_eq!(
22181 agent.get_context().get("draft_version"),
22182 Some(&serde_json::json!(2))
22183 );
22184 }
22185
22186 #[tokio::test]
22187 async fn test_timeout_transition_uses_on_enter_then_on_reenter() {
22188 let yaml = r#"
22189name: TimeoutLifecycleAgent
22190system_prompt: "You are helpful."
22191states:
22192 initial: intake
22193 regenerate_on_transition: false
22194 states:
22195 intake:
22196 prompt: "Intake"
22197 max_turns: 1
22198 timeout_to: drafting
22199 drafting:
22200 prompt: "Drafting"
22201 max_turns: 1
22202 timeout_to: review
22203 on_enter:
22204 - set_context:
22205 draft_version: 1
22206 on_reenter:
22207 - set_context:
22208 draft_version: 2
22209 review:
22210 prompt: "Review"
22211 max_turns: 1
22212 timeout_to: drafting
22213 on_enter:
22214 - set_context:
22215 review_entry: first
22216"#;
22217 let agent = AgentBuilder::from_yaml(yaml)
22218 .unwrap()
22219 .llm(Arc::new(mock_with_responses(vec![
22220 "Intake",
22221 "First draft",
22222 "Review",
22223 "Revised draft",
22224 ])))
22225 .build()
22226 .unwrap();
22227
22228 agent.chat("First turn").await.unwrap();
22229 assert_eq!(agent.current_state().as_deref(), Some("intake"));
22230 assert!(!agent.get_context().contains_key("draft_version"));
22231
22232 agent.chat("Second turn").await.unwrap();
22233 assert_eq!(agent.current_state().as_deref(), Some("drafting"));
22234 assert_eq!(
22235 agent.get_context().get("draft_version"),
22236 Some(&serde_json::json!(1))
22237 );
22238
22239 agent.chat("Third turn").await.unwrap();
22240 assert_eq!(agent.current_state().as_deref(), Some("review"));
22241 assert_eq!(
22242 agent.get_context().get("review_entry"),
22243 Some(&serde_json::json!("first"))
22244 );
22245
22246 agent.chat("Fourth turn").await.unwrap();
22247 assert_eq!(agent.current_state().as_deref(), Some("drafting"));
22248 assert_eq!(
22249 agent.get_context().get("draft_version"),
22250 Some(&serde_json::json!(2))
22251 );
22252 }
22253
22254 #[tokio::test]
22256 async fn test_integration_process_normalize() {
22257 let yaml = r#"
22258name: ProcessAgent
22259system_prompt: "You are helpful."
22260process:
22261 input:
22262 - type: normalize
22263 config:
22264 trim: true
22265 collapse_whitespace: true
22266"#;
22267 let mock = mock_with_response("Got your message.");
22268 let builder = AgentBuilder::from_yaml(yaml).unwrap();
22269 let agent = builder.llm(Arc::new(mock.clone())).build().unwrap();
22270
22271 let _ = agent.chat(" hello world ").await.unwrap();
22272
22273 let history = mock.call_history();
22275 assert!(!history.is_empty());
22276 let last_call = history.last().unwrap();
22278 let user_msg = last_call
22279 .messages
22280 .iter()
22281 .find(|m| m.role == ai_agents_core::Role::User)
22282 .unwrap();
22283 assert_eq!(user_msg.content, "hello world");
22284 }
22285
22286 #[tokio::test]
22290 async fn test_integration_memory_compression() {
22291 let yaml = r#"
22292name: MemoryAgent
22293system_prompt: "You are helpful."
22294memory:
22295 type: compacting
22296 max_messages: 100
22297 compress_threshold: 5
22298 max_recent_messages: 3
22299 summarize_batch_size: 2
22300"#;
22301 let responses: Vec<&str> = (0..8).map(|_| "Response from assistant.").collect();
22303 let mock = mock_with_responses(responses);
22304 let builder = AgentBuilder::from_yaml(yaml).unwrap();
22305 let agent = builder.llm(Arc::new(mock)).build().unwrap();
22306
22307 for i in 0..6 {
22309 let _ = agent.chat(&format!("Message {}", i)).await.unwrap();
22310 }
22311
22312 let messages = agent.memory.get_messages(None).await.unwrap();
22315 assert!(messages.len() <= 12); }
22319
22320 #[tokio::test]
22322 async fn test_integration_multi_llm_registry() {
22323 let mut mock_default = MockLLMProvider::new("default");
22324 mock_default.set_response("Default LLM response.");
22325 let mut mock_router = MockLLMProvider::new("router");
22326 mock_router.set_response("Router response.");
22327
22328 let agent = AgentBuilder::new()
22329 .system_prompt("You are helpful.")
22330 .llm_alias("default", Arc::new(mock_default))
22331 .llm_alias("router", Arc::new(mock_router))
22332 .build()
22333 .unwrap();
22334
22335 let response = agent.chat("Hello").await.unwrap();
22336 assert_eq!(response.content, "Default LLM response.");
22337 }
22338
22339 #[tokio::test]
22341 async fn test_integration_agent_reset() {
22342 let mock = mock_with_responses(vec!["Hello!", "Hello again!"]);
22343 let agent = AgentBuilder::new()
22344 .system_prompt("You are helpful.")
22345 .llm(Arc::new(mock))
22346 .build()
22347 .unwrap();
22348
22349 let _ = agent.chat("Hi").await.unwrap();
22350 let messages = agent.memory.get_messages(None).await.unwrap();
22351 assert_eq!(messages.len(), 2); agent.reset().await.unwrap();
22354 let messages = agent.memory.get_messages(None).await.unwrap();
22355 assert_eq!(messages.len(), 0);
22356 }
22357
22358 #[tokio::test]
22360 async fn test_integration_process_validate_reject() {
22361 use ai_agents_process::{ProcessConfig, ProcessProcessor};
22362
22363 let validate_config = ai_agents_process::ValidateStage {
22364 id: Some("length_check".to_string()),
22365 condition: None,
22366 config: ai_agents_process::ValidateConfig {
22367 rules: vec![ai_agents_process::ValidationRule::MinLength {
22368 min_length: 10,
22369 on_fail: ai_agents_process::ValidationAction {
22370 action: ai_agents_process::ValidationActionType::Reject,
22371 message: None,
22372 },
22373 }],
22374 ..Default::default()
22375 },
22376 };
22377 let process_config = ProcessConfig {
22378 input: vec![ai_agents_process::ProcessStage::Validate(validate_config)],
22379 ..Default::default()
22380 };
22381 let processor = ProcessProcessor::new(process_config);
22382
22383 let mock = mock_with_response("Should not reach here.");
22384 let agent = AgentBuilder::new()
22385 .system_prompt("You are helpful.")
22386 .llm(Arc::new(mock))
22387 .process_processor(processor)
22388 .build()
22389 .unwrap();
22390
22391 let response = agent.chat("Hi").await.unwrap();
22392 assert!(
22394 response.content.contains("rejected")
22395 || response.content.contains("Input rejected")
22396 || response.content.contains("too short")
22397 || response.content.contains("Too short")
22398 || response.content.len() < 50, "Expected rejection response, got: {}",
22400 response.content
22401 );
22402 }
22403
22404 #[tokio::test]
22406 async fn test_llm_fallback_on_failure() {
22407 use ai_agents_recovery::{ErrorRecoveryConfig, LLMFailureAction, LLMRecoveryConfig};
22408
22409 let mut primary = MockLLMProvider::new("primary");
22410 primary.set_error("Primary LLM is unavailable");
22411
22412 let mut fallback = MockLLMProvider::new("fallback");
22413 fallback.set_response("Fallback response works!");
22414
22415 let agent = AgentBuilder::new()
22416 .system_prompt("You are helpful.")
22417 .llm_alias("default", Arc::new(primary))
22418 .llm_alias("backup", Arc::new(fallback))
22419 .recovery_manager(RecoveryManager::new(ErrorRecoveryConfig {
22420 llm: LLMRecoveryConfig {
22421 on_failure: LLMFailureAction::FallbackLlm {
22422 fallback_llm: "backup".to_string(),
22423 },
22424 ..Default::default()
22425 },
22426 ..Default::default()
22427 }))
22428 .build()
22429 .unwrap();
22430
22431 let response = agent.chat("Hello").await.unwrap();
22432 assert!(
22433 response.content.contains("Fallback response"),
22434 "Expected fallback response, got: {}",
22435 response.content
22436 );
22437 }
22438
22439 #[tokio::test]
22441 async fn test_llm_fallback_response_static_message() {
22442 use ai_agents_recovery::{ErrorRecoveryConfig, LLMFailureAction, LLMRecoveryConfig};
22443
22444 let mut primary = MockLLMProvider::new("primary");
22445 primary.set_error("Primary LLM is unavailable");
22446
22447 let agent = AgentBuilder::new()
22448 .system_prompt("You are helpful.")
22449 .llm(Arc::new(primary))
22450 .recovery_manager(RecoveryManager::new(ErrorRecoveryConfig {
22451 llm: LLMRecoveryConfig {
22452 on_failure: LLMFailureAction::FallbackResponse {
22453 message: "I am temporarily unavailable. Please try again later."
22454 .to_string(),
22455 },
22456 ..Default::default()
22457 },
22458 ..Default::default()
22459 }))
22460 .build()
22461 .unwrap();
22462
22463 let response = agent.chat("Hello").await.unwrap();
22464 assert!(
22465 response.content.contains("temporarily unavailable"),
22466 "Expected static fallback message, got: {}",
22467 response.content
22468 );
22469 }
22470
22471 #[tokio::test]
22474 async fn test_tool_failure_skip() {
22475 use ai_agents_recovery::{
22476 ErrorRecoveryConfig, ToolFailureAction, ToolRecoveryConfig, ToolRetryConfig,
22477 };
22478
22479 let mock = mock_with_responses(vec![
22480 r#"{"tool": "calculator", "arguments": {"expression": "not a number +"}}"#,
22481 "The calculation was skipped, but I can still help you.",
22482 ]);
22483 let observed = mock.clone();
22484 let mut tools = ai_agents_tools::ToolRegistry::new();
22485 tools
22486 .register(Arc::new(ai_agents_tools::CalculatorTool))
22487 .unwrap();
22488
22489 let agent = AgentBuilder::new()
22490 .system_prompt("You are helpful.")
22491 .llm(Arc::new(mock))
22492 .tools(tools)
22493 .recovery_manager(RecoveryManager::new(ErrorRecoveryConfig {
22494 tools: ToolRecoveryConfig {
22495 default: ToolRetryConfig {
22496 max_retries: 0,
22497 timeout_ms: None,
22498 on_failure: ToolFailureAction::Skip,
22499 },
22500 ..Default::default()
22501 },
22502 ..Default::default()
22503 }))
22504 .build()
22505 .unwrap();
22506
22507 let response = agent.chat("Compute this").await.unwrap();
22508
22509 assert_eq!(
22510 response.content,
22511 "The calculation was skipped, but I can still help you."
22512 );
22513 assert_eq!(observed.call_count(), 2);
22514 let history = agent.tool_call_history();
22516 assert_eq!(history.len(), 1);
22517 assert_eq!(history[0].tool_id, "calculator");
22518 assert_eq!(
22519 history[0].result.get("skipped"),
22520 Some(&serde_json::json!(true)),
22521 "{:?}",
22522 history[0].result
22523 );
22524 }
22525
22526 #[tokio::test]
22528 async fn test_unregistered_tool_call_records_unavailable_and_continues() {
22529 let mock = mock_with_responses(vec![
22530 r#"{"tool": "nonexistent_tool", "arguments": {}}"#,
22531 "The tool was unavailable, but I can still help you.",
22532 ]);
22533 let observed = mock.clone();
22534
22535 let agent = AgentBuilder::new()
22536 .system_prompt("You are helpful.")
22537 .llm(Arc::new(mock))
22538 .build()
22539 .unwrap();
22540
22541 let response = agent.chat("Use the nonexistent tool").await.unwrap();
22542
22543 assert_eq!(
22544 response.content,
22545 "The tool was unavailable, but I can still help you."
22546 );
22547 assert_eq!(observed.call_count(), 2);
22548 let history = agent.tool_call_history();
22549 assert_eq!(history.len(), 1);
22550 assert_eq!(history[0].tool_id, "nonexistent_tool");
22551 assert_eq!(
22552 history[0].result.pointer("/error/kind"),
22553 Some(&serde_json::json!("tool_unavailable")),
22554 "{:?}",
22555 history[0].result
22556 );
22557 }
22558
22559 fn fallback_llm_recovery(fallback_llm: &str) -> RecoveryManager {
22564 use ai_agents_recovery::{ErrorRecoveryConfig, LLMFailureAction, LLMRecoveryConfig};
22565 RecoveryManager::new(ErrorRecoveryConfig {
22566 llm: LLMRecoveryConfig {
22567 on_failure: LLMFailureAction::FallbackLlm {
22568 fallback_llm: fallback_llm.to_string(),
22569 },
22570 ..Default::default()
22571 },
22572 ..Default::default()
22573 })
22574 }
22575
22576 #[tokio::test]
22577 async fn test_stream_llm_fallback_on_open_failure() {
22578 let mut primary = MockLLMProvider::new("primary");
22579 primary.set_error("Primary LLM is unavailable");
22580 let mut fallback = MockLLMProvider::new("fallback");
22581 fallback.set_response("Fallback response works!");
22582
22583 let agent = AgentBuilder::new()
22584 .system_prompt("You are helpful.")
22585 .llm_alias("default", Arc::new(primary))
22586 .llm_alias("backup", Arc::new(fallback))
22587 .recovery_manager(fallback_llm_recovery("backup"))
22588 .build()
22589 .unwrap();
22590
22591 let (content, chunks, final_response) = collect_stream_events(&agent, "Hello").await;
22592 assert!(
22593 !chunks.iter().any(StreamChunk::is_error),
22594 "fallback must not surface as a stream error: {chunks:?}"
22595 );
22596 let final_response = final_response.expect("Final must be emitted after fallback");
22597 assert!(content.contains("Fallback response"));
22598 assert!(final_response.content.contains("Fallback response"));
22599 }
22600
22601 #[tokio::test]
22602 async fn test_stream_llm_fallback_response_static_message() {
22603 use ai_agents_recovery::{ErrorRecoveryConfig, LLMFailureAction, LLMRecoveryConfig};
22604
22605 let mut primary = MockLLMProvider::new("primary");
22606 primary.set_error("Primary LLM is unavailable");
22607
22608 let agent = AgentBuilder::new()
22609 .system_prompt("You are helpful.")
22610 .llm(Arc::new(primary))
22611 .recovery_manager(RecoveryManager::new(ErrorRecoveryConfig {
22612 llm: LLMRecoveryConfig {
22613 on_failure: LLMFailureAction::FallbackResponse {
22614 message: "Service is temporarily unavailable.".to_string(),
22615 },
22616 ..Default::default()
22617 },
22618 ..Default::default()
22619 }))
22620 .build()
22621 .unwrap();
22622
22623 let (content, chunks, final_response) = collect_stream_events(&agent, "Hello").await;
22624 assert!(!chunks.iter().any(StreamChunk::is_error));
22625 let content_chunks = chunks.iter().filter(|c| c.is_content()).count();
22626 assert_eq!(content_chunks, 1, "static fallback is one content chunk");
22627 assert_eq!(content, "Service is temporarily unavailable.");
22628 assert_eq!(
22629 final_response.expect("Final").content,
22630 "Service is temporarily unavailable."
22631 );
22632 }
22633
22634 struct FailOnceStreamProvider {
22636 remaining_failures: Arc<std::sync::atomic::AtomicUsize>,
22637 open_attempts: Arc<std::sync::atomic::AtomicUsize>,
22638 }
22639
22640 #[async_trait]
22641 impl LLMProvider for FailOnceStreamProvider {
22642 async fn complete(
22643 &self,
22644 _messages: &[ChatMessage],
22645 _config: Option<&LLMConfig>,
22646 ) -> std::result::Result<LLMResponse, LLMError> {
22647 Ok(LLMResponse::new("blocking path", FinishReason::Stop))
22648 }
22649
22650 async fn complete_stream(
22651 &self,
22652 _messages: &[ChatMessage],
22653 _config: Option<&LLMConfig>,
22654 ) -> std::result::Result<
22655 Box<dyn Stream<Item = std::result::Result<LLMChunk, LLMError>> + Unpin + Send>,
22656 LLMError,
22657 > {
22658 self.open_attempts.fetch_add(1, Ordering::SeqCst);
22659 if self
22660 .remaining_failures
22661 .fetch_update(Ordering::SeqCst, Ordering::SeqCst, |n| n.checked_sub(1))
22662 .is_ok()
22663 {
22664 return Err(LLMError::Network("connection reset".to_string()));
22665 }
22666 Ok(Box::new(futures::stream::iter(vec![Ok(LLMChunk::new(
22667 "Recovered after retry",
22668 true,
22669 ))])))
22670 }
22671
22672 fn provider_name(&self) -> &str {
22673 "fail-once-stream"
22674 }
22675
22676 fn supports(&self, feature: LLMFeature) -> bool {
22677 matches!(feature, LLMFeature::Streaming)
22678 }
22679 }
22680
22681 #[tokio::test]
22682 async fn test_stream_llm_retry_then_success() {
22683 use ai_agents_recovery::{BackoffConfig, ErrorRecoveryConfig, RetryConfig};
22684
22685 let open_attempts = Arc::new(std::sync::atomic::AtomicUsize::new(0));
22686 let provider = FailOnceStreamProvider {
22687 remaining_failures: Arc::new(std::sync::atomic::AtomicUsize::new(1)),
22688 open_attempts: Arc::clone(&open_attempts),
22689 };
22690
22691 let agent = AgentBuilder::new()
22692 .system_prompt("You are helpful.")
22693 .llm(Arc::new(provider))
22694 .recovery_manager(RecoveryManager::new(ErrorRecoveryConfig {
22695 default: RetryConfig {
22696 max_retries: 1,
22697 backoff: BackoffConfig {
22698 initial_ms: 1,
22699 max_ms: 1,
22700 ..Default::default()
22701 },
22702 ..Default::default()
22703 },
22704 ..Default::default()
22705 }))
22706 .build()
22707 .unwrap();
22708
22709 let (content, chunks, final_response) = collect_stream_events(&agent, "Hello").await;
22710 assert!(!chunks.iter().any(StreamChunk::is_error), "{chunks:?}");
22711 assert_eq!(open_attempts.load(Ordering::SeqCst), 2);
22712 assert_eq!(content, "Recovered after retry");
22713 assert_eq!(
22714 final_response.expect("Final").content,
22715 "Recovered after retry"
22716 );
22717 }
22718
22719 #[tokio::test]
22720 async fn test_stream_llm_error_action_error_emits_terminal_error() {
22721 let mut primary = MockLLMProvider::new("primary");
22722 primary.set_error("Primary LLM is unavailable");
22723
22724 let agent = AgentBuilder::new()
22725 .system_prompt("You are helpful.")
22726 .llm(Arc::new(primary))
22727 .build()
22728 .unwrap();
22729
22730 let (_, chunks, final_response) = collect_stream_events(&agent, "Hello").await;
22731 assert!(
22732 final_response.is_none(),
22733 "default Error action must not produce Final"
22734 );
22735 assert!(
22736 chunks.iter().any(StreamChunk::is_error),
22737 "default Error action must surface a stream error"
22738 );
22739 }
22740
22741 struct MidStreamFailureProvider;
22743
22744 #[async_trait]
22745 impl LLMProvider for MidStreamFailureProvider {
22746 async fn complete(
22747 &self,
22748 _messages: &[ChatMessage],
22749 _config: Option<&LLMConfig>,
22750 ) -> std::result::Result<LLMResponse, LLMError> {
22751 Ok(LLMResponse::new("blocking path", FinishReason::Stop))
22752 }
22753
22754 async fn complete_stream(
22755 &self,
22756 _messages: &[ChatMessage],
22757 _config: Option<&LLMConfig>,
22758 ) -> std::result::Result<
22759 Box<dyn Stream<Item = std::result::Result<LLMChunk, LLMError>> + Unpin + Send>,
22760 LLMError,
22761 > {
22762 Ok(Box::new(futures::stream::iter(vec![
22763 Ok(LLMChunk::new("Partial ", false)),
22764 Err(LLMError::Network("connection dropped".to_string())),
22765 ])))
22766 }
22767
22768 fn provider_name(&self) -> &str {
22769 "mid-stream-failure"
22770 }
22771
22772 fn supports(&self, feature: LLMFeature) -> bool {
22773 matches!(feature, LLMFeature::Streaming)
22774 }
22775 }
22776
22777 #[tokio::test]
22778 async fn test_stream_mid_stream_failure_is_terminal() {
22779 let mut fallback = MockLLMProvider::new("fallback");
22780 fallback.set_response("Fallback must not run");
22781 let fallback_calls = fallback.clone();
22782
22783 let agent = AgentBuilder::new()
22784 .system_prompt("You are helpful.")
22785 .llm_alias("default", Arc::new(MidStreamFailureProvider))
22786 .llm_alias("backup", Arc::new(fallback))
22787 .recovery_manager(fallback_llm_recovery("backup"))
22788 .build()
22789 .unwrap();
22790
22791 let (content, chunks, final_response) = collect_stream_events(&agent, "Hello").await;
22792 assert_eq!(content, "Partial ");
22793 assert!(chunks.iter().any(StreamChunk::is_error));
22794 assert!(final_response.is_none());
22795 assert_eq!(
22796 fallback_calls.call_count(),
22797 0,
22798 "fallback must not run after a visible delta"
22799 );
22800 }
22801
22802 #[tokio::test]
22803 async fn test_buffered_streaming_draft_uses_fallback_llm() {
22804 use futures::StreamExt;
22805
22806 let mut primary = MockLLMProvider::new("primary");
22807 primary.set_error("Primary LLM is unavailable");
22808 let fallback = mock_with_response("fallback one two");
22809 let yaml = r#"
22810name: BufferedFallbackAgent
22811system_prompt: "You stream safely."
22812llm:
22813 default: default
22814streaming:
22815 enabled: true
22816 buffer_size: 8
22817runtime:
22818 optimization:
22819 enabled: true
22820 max_speculative_llm_calls_per_turn: 2
22821 speculative_state_transitions: true
22822 streaming_policy: buffer_until_routing_done
22823 max_parallel_runtime_tasks: 2
22824states:
22825 initial: triage
22826 states:
22827 triage:
22828 prompt: "Answer from triage."
22829 transitions:
22830 - to: billing
22831 guard:
22832 context:
22833 route:
22834 eq: billing
22835 timing: parallel
22836 billing:
22837 prompt: "Billing state."
22838"#;
22839 let agent = AgentBuilder::from_yaml(yaml)
22840 .unwrap()
22841 .llm_alias("default", Arc::new(primary))
22842 .llm_alias("backup", Arc::new(fallback))
22843 .recovery_manager(fallback_llm_recovery("backup"))
22844 .build()
22845 .unwrap();
22846
22847 let mut stream = agent.chat_stream("hello").await.unwrap();
22848 let mut content = String::new();
22849 let mut error = None;
22850 while let Some(chunk) = stream.next().await {
22851 match chunk {
22852 StreamChunk::Content { text } => content.push_str(&text),
22853 StreamChunk::Error { message } => error = Some(message),
22854 StreamChunk::Done {} => break,
22855 _ => {}
22856 }
22857 }
22858
22859 assert_eq!(error, None);
22860 assert_eq!(content, "fallback one two");
22861 }
22862
22863 #[tokio::test]
22864 async fn parity_llm_fallback_llm() {
22865 let build = || {
22866 let mut primary = MockLLMProvider::new("primary");
22867 primary.set_error("Primary LLM is unavailable");
22868 let mut fallback = MockLLMProvider::new("fallback");
22869 fallback.set_response("Fallback response works!");
22870 AgentBuilder::new()
22871 .system_prompt("You are helpful.")
22872 .llm_alias("default", Arc::new(primary))
22873 .llm_alias("backup", Arc::new(fallback))
22874 .recovery_manager(fallback_llm_recovery("backup"))
22875 .build()
22876 .unwrap()
22877 };
22878 assert_blocking_streaming_parity(build, "Hello").await;
22879 }
22880
22881 #[tokio::test]
22882 async fn parity_llm_fallback_response() {
22883 use ai_agents_recovery::{ErrorRecoveryConfig, LLMFailureAction, LLMRecoveryConfig};
22884 let build = || {
22885 let mut primary = MockLLMProvider::new("primary");
22886 primary.set_error("Primary LLM is unavailable");
22887 AgentBuilder::new()
22888 .system_prompt("You are helpful.")
22889 .llm(Arc::new(primary))
22890 .recovery_manager(RecoveryManager::new(ErrorRecoveryConfig {
22891 llm: LLMRecoveryConfig {
22892 on_failure: LLMFailureAction::FallbackResponse {
22893 message: "Service is temporarily unavailable.".to_string(),
22894 },
22895 ..Default::default()
22896 },
22897 ..Default::default()
22898 }))
22899 .build()
22900 .unwrap()
22901 };
22902 assert_blocking_streaming_parity(build, "Hello").await;
22903 }
22904
22905 #[tokio::test]
22906 async fn parity_basic_chat() {
22907 let build = || {
22908 AgentBuilder::new()
22909 .system_prompt("You are helpful.")
22910 .llm(Arc::new(mock_with_response("Plain answer")))
22911 .build()
22912 .unwrap()
22913 };
22914 assert_blocking_streaming_parity(build, "Hello").await;
22915 }
22916
22917 fn skills_with_parallel_transition_yaml(extra_optimization: &str, streaming: &str) -> String {
22923 format!(
22924 r#"
22925name: SkillsBesideTransitionAgent
22926system_prompt: "Use skills when they match."
22927llm:
22928 default: default
22929 router: router
22930observability:
22931 enabled: true
22932 export:
22933 write_raw_events: true
22934{streaming}
22935runtime:
22936 optimization:
22937 enabled: true
22938 speculative_state_transitions: true
22939{extra_optimization}
22940states:
22941 initial: triage
22942 states:
22943 triage:
22944 prompt: "Triage state."
22945 transitions:
22946 - to: billing
22947 guard:
22948 context:
22949 route:
22950 eq: billing
22951 timing: parallel
22952 billing:
22953 prompt: "Billing state."
22954skills:
22955 - id: helper
22956 description: "Answer helper requests"
22957 trigger: "User asks for helper"
22958 steps:
22959 - prompt: "Answer the helper request: {{{{ user_input }}}}"
22960 llm: skill
22961"#
22962 )
22963 }
22964
22965 struct RoleMocks {
22969 main: MockLLMProvider,
22970 router: MockLLMProvider,
22971 skill: MockLLMProvider,
22972 }
22973
22974 fn role_mocks(main: MockLLMProvider, router: MockLLMProvider) -> RoleMocks {
22975 RoleMocks {
22976 main,
22977 router,
22978 skill: mock_with_response("Skill step response"),
22979 }
22980 }
22981
22982 fn build_skills_beside_transition_agent(yaml: &str, mocks: RoleMocks) -> RuntimeAgent {
22983 AgentBuilder::from_yaml(yaml)
22984 .unwrap()
22985 .llm_alias("default", Arc::new(mocks.main))
22986 .llm_alias("router", Arc::new(mocks.router))
22987 .llm_alias("skill", Arc::new(mocks.skill))
22988 .build()
22989 .unwrap()
22990 }
22991
22992 fn branch_events_with_commit_behavior(agent: &RuntimeAgent, behavior: &str) -> usize {
22993 agent
22994 .observability()
22995 .unwrap()
22996 .raw_events()
22997 .iter()
22998 .filter(|event| event.dimensions.get("commit_behavior") == Some(&behavior.to_string()))
22999 .count()
23000 }
23001
23002 #[tokio::test]
23003 async fn test_speculative_transition_with_skills_and_no_skill_branch_routes_skill_serially() {
23004 let default_mock = mock_with_response("Draft response");
23005 let router_mock = mock_with_response("helper");
23006 let router_counter = router_mock.clone();
23007 let yaml = skills_with_parallel_transition_yaml(
23008 " max_speculative_llm_calls_per_turn: 2\n max_parallel_runtime_tasks: 2",
23009 "",
23010 );
23011 let agent =
23012 build_skills_beside_transition_agent(&yaml, role_mocks(default_mock, router_mock));
23013
23014 let response = agent.chat("please use helper").await.unwrap();
23015
23016 assert_eq!(
23017 response.metadata.as_ref().and_then(|m| m.get("skill_id")),
23018 Some(&serde_json::json!("helper")),
23019 "skill must route even without a skill branch: {response:?}"
23020 );
23021 assert_eq!(router_counter.call_count(), 1);
23022 assert!(branch_events_with_commit_behavior(&agent, "transition_decision") > 0);
23024 assert_eq!(
23025 branch_events_with_commit_behavior(&agent, "skill_selection"),
23026 0
23027 );
23028 }
23029
23030 #[tokio::test]
23031 async fn test_speculative_transition_with_skills_no_match_commits_draft() {
23032 let default_mock = mock_with_response("Draft response");
23033 let router_mock = mock_with_response("none");
23034 let router_counter = router_mock.clone();
23035 let yaml = skills_with_parallel_transition_yaml(
23036 " max_speculative_llm_calls_per_turn: 2\n max_parallel_runtime_tasks: 2",
23037 "",
23038 );
23039 let agent =
23040 build_skills_beside_transition_agent(&yaml, role_mocks(default_mock, router_mock));
23041
23042 let response = agent.chat("just chat").await.unwrap();
23043
23044 assert_eq!(response.content, "Draft response");
23045 assert!(
23046 response
23047 .metadata
23048 .as_ref()
23049 .is_none_or(|m| !m.contains_key("skill_id"))
23050 );
23051 assert_eq!(router_counter.call_count(), 1);
23052 assert!(branch_events_with_commit_behavior(&agent, "final_response") > 0);
23053 }
23054
23055 #[tokio::test]
23056 async fn test_speculative_transition_win_skips_serial_skill_selection() {
23057 let default_mock = mock_with_response("Billing answer");
23058 let router_mock = mock_with_response("none");
23059 let router_counter = router_mock.clone();
23060 let yaml = skills_with_parallel_transition_yaml(
23061 " max_speculative_llm_calls_per_turn: 2\n max_parallel_runtime_tasks: 2",
23062 "",
23063 );
23064 let agent =
23065 build_skills_beside_transition_agent(&yaml, role_mocks(default_mock, router_mock));
23066 agent
23067 .set_context("route", serde_json::json!("billing"))
23068 .unwrap();
23069
23070 let response = agent.chat("billing please").await.unwrap();
23071
23072 assert_eq!(agent.current_state().as_deref(), Some("billing"));
23073 assert_eq!(response.content, "Billing answer");
23074 assert_eq!(router_counter.call_count(), 1);
23076 }
23077
23078 #[tokio::test]
23079 async fn test_speculative_skill_capacity_exhausted_still_routes_skill_serially() {
23080 let default_mock = mock_with_response("Draft response");
23081 let router_mock = mock_with_response("helper");
23082 let router_counter = router_mock.clone();
23083 let yaml = skills_with_parallel_transition_yaml(
23085 " speculative_skill_routing: true\n max_speculative_llm_calls_per_turn: 2\n max_parallel_runtime_tasks: 2",
23086 "",
23087 );
23088 let agent =
23089 build_skills_beside_transition_agent(&yaml, role_mocks(default_mock, router_mock));
23090
23091 let response = agent.chat("please use helper").await.unwrap();
23092
23093 assert_eq!(
23094 response.metadata.as_ref().and_then(|m| m.get("skill_id")),
23095 Some(&serde_json::json!("helper"))
23096 );
23097 assert_eq!(router_counter.call_count(), 1);
23098 assert_eq!(
23099 branch_events_with_commit_behavior(&agent, "skill_selection"),
23100 0
23101 );
23102 }
23103
23104 #[tokio::test]
23105 async fn test_speculative_transition_and_skill_both_enabled_unchanged() {
23106 let default_mock = mock_with_response("Draft response");
23107 let router_mock = mock_with_response("helper");
23108 let router_counter = router_mock.clone();
23109 let yaml = skills_with_parallel_transition_yaml(
23110 " speculative_skill_routing: true\n max_speculative_llm_calls_per_turn: 3\n max_parallel_runtime_tasks: 3",
23111 "",
23112 );
23113 let agent =
23114 build_skills_beside_transition_agent(&yaml, role_mocks(default_mock, router_mock));
23115
23116 let response = agent.chat("please use helper").await.unwrap();
23117
23118 assert_eq!(
23119 response.metadata.as_ref().and_then(|m| m.get("skill_id")),
23120 Some(&serde_json::json!("helper"))
23121 );
23122 assert_eq!(router_counter.call_count(), 1);
23123 assert!(branch_events_with_commit_behavior(&agent, "skill_selection") > 0);
23125 }
23126
23127 const BUFFERED_STREAMING_YAML_FRAGMENT: &str = "streaming:\n enabled: true\n buffer_size: 16";
23128 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";
23129
23130 #[tokio::test]
23131 async fn test_buffered_streaming_skill_wins_after_transition_miss() {
23132 let mut default_mock = mock_with_response("draft one two");
23133 default_mock.set_latency(10);
23134 let router_mock = mock_with_response("helper");
23135 let router_counter = router_mock.clone();
23136 let yaml = skills_with_parallel_transition_yaml(
23137 BUFFERED_OPTIMIZATION_FRAGMENT,
23138 BUFFERED_STREAMING_YAML_FRAGMENT,
23139 );
23140 let agent =
23141 build_skills_beside_transition_agent(&yaml, role_mocks(default_mock, router_mock));
23142
23143 let (content, chunks, final_response) =
23144 collect_stream_events(&agent, "please use helper").await;
23145
23146 assert!(!chunks.iter().any(StreamChunk::is_error), "{chunks:?}");
23147 assert!(
23148 !content.contains("draft"),
23149 "buffered draft must be discarded when a skill wins: {content:?}"
23150 );
23151 let final_response = final_response.expect("Final");
23152 assert_eq!(
23153 final_response
23154 .metadata
23155 .as_ref()
23156 .and_then(|m| m.get("skill_id")),
23157 Some(&serde_json::json!("helper"))
23158 );
23159 assert_eq!(content, final_response.content);
23160 assert_eq!(router_counter.call_count(), 1);
23161 }
23162
23163 #[tokio::test]
23164 async fn test_buffered_streaming_skill_miss_releases_buffer_and_commits_draft() {
23165 let mut default_mock = mock_with_response("draft one two");
23166 default_mock.set_latency(10);
23167 let router_mock = mock_with_response("none");
23168 let router_counter = router_mock.clone();
23169 let yaml = skills_with_parallel_transition_yaml(
23170 BUFFERED_OPTIMIZATION_FRAGMENT,
23171 BUFFERED_STREAMING_YAML_FRAGMENT,
23172 );
23173 let agent =
23174 build_skills_beside_transition_agent(&yaml, role_mocks(default_mock, router_mock));
23175
23176 let (content, chunks, final_response) = collect_stream_events(&agent, "just chat").await;
23177
23178 assert!(!chunks.iter().any(StreamChunk::is_error), "{chunks:?}");
23179 assert_eq!(content, "draft one two");
23180 assert_eq!(final_response.expect("Final").content, "draft one two");
23181 assert_eq!(router_counter.call_count(), 1);
23182 }
23183
23184 #[tokio::test]
23185 async fn test_buffered_streaming_transition_win_skips_skill_selection() {
23186 let default_mock = mock_with_response("Billing answer");
23187 let router_mock = mock_with_response("none");
23188 let router_counter = router_mock.clone();
23189 let yaml = skills_with_parallel_transition_yaml(
23190 BUFFERED_OPTIMIZATION_FRAGMENT,
23191 BUFFERED_STREAMING_YAML_FRAGMENT,
23192 );
23193 let agent =
23194 build_skills_beside_transition_agent(&yaml, role_mocks(default_mock, router_mock));
23195 agent
23196 .set_context("route", serde_json::json!("billing"))
23197 .unwrap();
23198
23199 let (content, chunks, final_response) =
23200 collect_stream_events(&agent, "billing please").await;
23201
23202 assert!(!chunks.iter().any(StreamChunk::is_error), "{chunks:?}");
23203 assert_eq!(agent.current_state().as_deref(), Some("billing"));
23204 assert_eq!(content, "Billing answer");
23205 assert_eq!(final_response.expect("Final").content, "Billing answer");
23206 assert_eq!(router_counter.call_count(), 1);
23208 }
23209
23210 #[tokio::test]
23211 async fn parity_buffered_policy_with_skills() {
23212 let yaml = skills_with_parallel_transition_yaml(
23213 BUFFERED_OPTIMIZATION_FRAGMENT,
23214 BUFFERED_STREAMING_YAML_FRAGMENT,
23215 );
23216 let build = || {
23217 build_skills_beside_transition_agent(
23218 &yaml,
23219 role_mocks(
23220 mock_with_response("draft one two"),
23221 mock_with_response("helper"),
23222 ),
23223 )
23224 };
23225 let (blocking, _, _) = assert_blocking_streaming_parity(build, "please use helper").await;
23226 assert_eq!(
23227 blocking.metadata.as_ref().and_then(|m| m.get("skill_id")),
23228 Some(&serde_json::json!("helper"))
23229 );
23230 }
23231
23232 #[tokio::test]
23233 async fn parity_buffered_policy_with_cot() {
23234 let yaml = format!(
23235 r#"
23236name: BufferedCotAgent
23237system_prompt: "Think first."
23238llm:
23239 default: default
23240streaming:
23241 enabled: true
23242 buffer_size: 16
23243reasoning:
23244 mode: cot
23245runtime:
23246 optimization:
23247 enabled: true
23248 speculative_state_transitions: true
23249{BUFFERED_OPTIMIZATION_FRAGMENT}
23250states:
23251 initial: triage
23252 states:
23253 triage:
23254 prompt: "Triage state."
23255 transitions:
23256 - to: billing
23257 guard:
23258 context:
23259 route:
23260 eq: billing
23261 timing: parallel
23262 billing:
23263 prompt: "Billing state."
23264"#
23265 );
23266 let build = || {
23267 AgentBuilder::from_yaml(&yaml)
23268 .unwrap()
23269 .llm_alias(
23270 "default",
23271 Arc::new(mock_with_response(
23272 "<thinking>step by step</thinking>Reasoned answer",
23273 )),
23274 )
23275 .build()
23276 .unwrap()
23277 };
23278 let (blocking, streamed, _) = assert_blocking_streaming_parity(build, "hello").await;
23279 assert_eq!(blocking.content, "Reasoned answer");
23280 let mode = streamed
23281 .metadata
23282 .as_ref()
23283 .and_then(|m| m.get("reasoning"))
23284 .and_then(|r| r.get("mode_used"))
23285 .cloned();
23286 assert_eq!(
23288 mode,
23289 Some(serde_json::to_value(ReasoningMode::CoT).unwrap())
23290 );
23291 }
23292
23293 fn post_response_transition_yaml(states_extra: &str, billing_extra: &str) -> String {
23300 format!(
23301 r#"
23302name: PostResponseTransitionAgent
23303system_prompt: "You are helpful."
23304streaming:
23305 enabled: true
23306states:
23307 initial: intake
23308{states_extra}
23309 states:
23310 intake:
23311 prompt: "Intake"
23312 transitions:
23313 - to: billing
23314 guard:
23315 context:
23316 route:
23317 eq: billing
23318 billing:
23319 prompt: "Billing"
23320{billing_extra}
23321"#
23322 )
23323 }
23324
23325 fn build_post_response_transition_agent(yaml: &str, mock: MockLLMProvider) -> RuntimeAgent {
23326 let agent = AgentBuilder::from_yaml(yaml)
23327 .unwrap()
23328 .llm(Arc::new(mock))
23329 .build()
23330 .unwrap();
23331 agent
23332 .set_context("route", serde_json::json!("billing"))
23333 .unwrap();
23334 agent
23335 }
23336
23337 fn count_occurrences(haystack: &str, needle: &str) -> usize {
23338 haystack.matches(needle).count()
23339 }
23340
23341 #[tokio::test]
23342 async fn test_stream_transition_without_regeneration_emits_content_once() {
23343 let yaml = post_response_transition_yaml(" regenerate_on_transition: false", "");
23344 let agent =
23345 build_post_response_transition_agent(&yaml, mock_with_response("Intake answer"));
23346
23347 let (content, chunks, final_response) = collect_stream_events(&agent, "hello").await;
23348
23349 assert!(!chunks.iter().any(StreamChunk::is_error), "{chunks:?}");
23350 assert_eq!(agent.current_state().as_deref(), Some("billing"));
23351 assert_eq!(
23352 count_occurrences(&content, "Intake answer"),
23353 1,
23354 "committed content must not be emitted twice: {content:?}"
23355 );
23356 assert_eq!(final_response.expect("Final").content, content);
23357 assert!(
23358 chunks
23359 .iter()
23360 .any(|c| matches!(c, StreamChunk::StateTransition { .. }))
23361 );
23362 }
23363
23364 #[tokio::test]
23365 async fn test_stream_transition_without_regeneration_buffered_emits_content_once() {
23366 let yaml = post_response_transition_yaml(" regenerate_on_transition: false", "");
23367 let mut mock = mock_with_response("Intake answer");
23368 mock.set_tool_choice(Some(ToolChoice::Auto));
23370 let agent = build_post_response_transition_agent(&yaml, mock);
23371
23372 let (content, chunks, final_response) = collect_stream_events(&agent, "hello").await;
23373
23374 assert!(!chunks.iter().any(StreamChunk::is_error), "{chunks:?}");
23375 assert_eq!(agent.current_state().as_deref(), Some("billing"));
23376 assert_eq!(
23377 count_occurrences(&content, "Intake answer"),
23378 1,
23379 "{content:?}"
23380 );
23381 assert_eq!(final_response.expect("Final").content, content);
23382 }
23383
23384 #[tokio::test]
23385 async fn test_stream_state_regenerate_on_enter_false_emits_content_once() {
23386 let yaml = post_response_transition_yaml("", " regenerate_on_enter: false");
23387 let agent =
23388 build_post_response_transition_agent(&yaml, mock_with_response("Intake answer"));
23389
23390 let (content, chunks, final_response) = collect_stream_events(&agent, "hello").await;
23391
23392 assert!(!chunks.iter().any(StreamChunk::is_error), "{chunks:?}");
23393 assert_eq!(agent.current_state().as_deref(), Some("billing"));
23394 assert_eq!(
23395 count_occurrences(&content, "Intake answer"),
23396 1,
23397 "{content:?}"
23398 );
23399 assert_eq!(final_response.expect("Final").content, content);
23400 }
23401
23402 #[tokio::test]
23403 async fn test_stream_transition_with_regeneration_emits_replacement() {
23404 let yaml = post_response_transition_yaml("", "");
23405 let agent = build_post_response_transition_agent(
23406 &yaml,
23407 mock_with_responses(vec!["Intake answer", "Billing answer"]),
23408 );
23409
23410 let (content, chunks, final_response) = collect_stream_events(&agent, "hello").await;
23411
23412 assert!(!chunks.iter().any(StreamChunk::is_error), "{chunks:?}");
23413 assert_eq!(agent.current_state().as_deref(), Some("billing"));
23414 assert_eq!(count_occurrences(&content, "Intake answer"), 1);
23416 assert_eq!(count_occurrences(&content, "Billing answer"), 1);
23417 assert_eq!(final_response.expect("Final").content, "Billing answer");
23418 }
23419
23420 #[tokio::test]
23421 async fn test_blocking_transition_without_regeneration_unchanged() {
23422 let yaml = post_response_transition_yaml(" regenerate_on_transition: false", "");
23423 let agent =
23424 build_post_response_transition_agent(&yaml, mock_with_response("Intake answer"));
23425
23426 let response = agent.chat("hello").await.unwrap();
23427
23428 assert_eq!(response.content, "Intake answer");
23429 assert_eq!(agent.current_state().as_deref(), Some("billing"));
23430 }
23431
23432 #[tokio::test]
23433 async fn parity_transition_regenerate_off() {
23434 let yaml = post_response_transition_yaml(" regenerate_on_transition: false", "");
23435 let build =
23436 || build_post_response_transition_agent(&yaml, mock_with_response("Intake answer"));
23437 assert_blocking_streaming_parity(build, "hello").await;
23438 }
23439
23440 #[tokio::test]
23441 async fn parity_transition_regenerate_on() {
23442 let yaml = post_response_transition_yaml("", "");
23443 let build = || {
23444 build_post_response_transition_agent(
23445 &yaml,
23446 mock_with_responses(vec!["Intake answer", "Billing answer"]),
23447 )
23448 };
23449 let (blocking, _, _) = assert_blocking_streaming_parity(build, "hello").await;
23450 assert_eq!(blocking.content, "Billing answer");
23451 }
23452
23453 fn rejecting_process_processor() -> ProcessProcessor {
23458 use ai_agents_process::ProcessConfig;
23459 let validate_config = ai_agents_process::ValidateStage {
23460 id: Some("length_check".to_string()),
23461 condition: None,
23462 config: ai_agents_process::ValidateConfig {
23463 rules: vec![ai_agents_process::ValidationRule::MinLength {
23464 min_length: 10,
23465 on_fail: ai_agents_process::ValidationAction {
23466 action: ai_agents_process::ValidationActionType::Reject,
23467 message: None,
23468 },
23469 }],
23470 ..Default::default()
23471 },
23472 };
23473 ProcessProcessor::new(ProcessConfig {
23474 input: vec![ai_agents_process::ProcessStage::Validate(validate_config)],
23475 ..Default::default()
23476 })
23477 }
23478
23479 fn looks_like_rejection(content: &str) -> bool {
23481 content.contains("rejected")
23482 || content.contains("Input rejected")
23483 || content.contains("too short")
23484 || content.contains("Too short")
23485 || content.len() < 50
23486 }
23487
23488 #[tokio::test]
23489 async fn test_stream_input_rejection_is_final_response() {
23490 let responses = Arc::new(std::sync::atomic::AtomicUsize::new(0));
23491 let hooks = Arc::new(ResponseCountingHooks {
23492 responses: Arc::clone(&responses),
23493 });
23494 let mock = mock_with_response("Should not reach here.");
23495 let llm_calls = mock.clone();
23496 let agent = AgentBuilder::new()
23497 .system_prompt("You are helpful.")
23498 .llm(Arc::new(mock))
23499 .process_processor(rejecting_process_processor())
23500 .hooks(hooks.clone())
23501 .build()
23502 .unwrap();
23503
23504 let (content, chunks, final_response) = collect_stream_events(&agent, "Hi").await;
23505
23506 assert!(
23507 !chunks.iter().any(StreamChunk::is_error),
23508 "rejection is a response, not a stream error: {chunks:?}"
23509 );
23510 let final_response = final_response.expect("rejection must finalize as Final");
23511 assert!(
23512 looks_like_rejection(&final_response.content),
23513 "Expected rejection response, got: {}",
23514 final_response.content
23515 );
23516 assert_eq!(content, final_response.content);
23517 assert_eq!(
23518 llm_calls.call_count(),
23519 0,
23520 "rejected input must not reach the LLM"
23521 );
23522 assert_eq!(responses.load(Ordering::SeqCst), 1, "on_response must fire");
23523 }
23524
23525 #[tokio::test]
23526 async fn parity_input_rejection() {
23527 let build = || {
23528 AgentBuilder::new()
23529 .system_prompt("You are helpful.")
23530 .llm(Arc::new(mock_with_response("Should not reach here.")))
23531 .process_processor(rejecting_process_processor())
23532 .build()
23533 .unwrap()
23534 };
23535 let (blocking, _, _) = assert_blocking_streaming_parity(build, "Hi").await;
23536 assert!(
23537 looks_like_rejection(&blocking.content),
23538 "{}",
23539 blocking.content
23540 );
23541 }
23542
23543 fn pre_response_transition_yaml(streaming_policy: &str) -> String {
23544 format!(
23545 r#"
23546name: StreamingPreflightAgent
23547system_prompt: "You route before streaming."
23548runtime:
23549 optimization:
23550 enabled: true
23551 pre_response_deterministic_transitions: true
23552 streaming_policy: {streaming_policy}
23553streaming:
23554 enabled: true
23555 buffer_size: 16
23556states:
23557 initial: greeting
23558 states:
23559 greeting:
23560 prompt: "OLD_STATE_SENTINEL"
23561 transitions:
23562 - to: billing
23563 guard:
23564 context:
23565 topic:
23566 eq: billing
23567 timing: pre_response
23568 billing:
23569 prompt: "Billing state."
23570"#
23571 )
23572 }
23573
23574 fn build_pre_response_transition_agent(yaml: &str) -> RuntimeAgent {
23575 let agent = AgentBuilder::from_yaml(yaml)
23576 .unwrap()
23577 .llm(Arc::new(mock_with_response("Billing streamed response")))
23578 .build()
23579 .unwrap();
23580 agent
23581 .set_context("topic", serde_json::json!("billing"))
23582 .unwrap();
23583 agent
23584 }
23585
23586 #[tokio::test]
23587 async fn test_stream_buffered_policy_runs_pre_response_deterministic_transition() {
23588 let yaml = pre_response_transition_yaml("buffer_until_routing_done");
23589 let agent = build_pre_response_transition_agent(&yaml);
23590
23591 let (content, chunks, final_response) =
23592 collect_stream_events(&agent, "billing please").await;
23593
23594 assert!(!chunks.iter().any(StreamChunk::is_error), "{chunks:?}");
23595 assert_eq!(agent.current_state().as_deref(), Some("billing"));
23596 assert!(content.contains("Billing streamed response"));
23597 assert!(!content.contains("OLD_STATE_SENTINEL"));
23598 assert_eq!(final_response.expect("Final").content, content);
23599 }
23600
23601 #[tokio::test]
23602 async fn test_stream_disabled_policy_skips_preflight() {
23603 let yaml = pre_response_transition_yaml("disabled");
23604 let agent = build_pre_response_transition_agent(&yaml);
23605
23606 let (_, chunks, final_response) = collect_stream_events(&agent, "billing please").await;
23607
23608 assert!(!chunks.iter().any(StreamChunk::is_error), "{chunks:?}");
23609 assert!(final_response.is_some());
23610 assert_eq!(agent.current_state().as_deref(), Some("greeting"));
23614 }
23615
23616 #[tokio::test]
23617 async fn parity_pre_response_transition_buffered_policy() {
23618 let yaml = pre_response_transition_yaml("buffer_until_routing_done");
23619 let build = || build_pre_response_transition_agent(&yaml);
23620 assert_blocking_streaming_parity(build, "billing please").await;
23621 }
23622
23623 fn calculator_agent_with(mock: MockLLMProvider) -> RuntimeAgent {
23628 let mut tools = ai_agents_tools::ToolRegistry::new();
23629 tools
23630 .register(Arc::new(ai_agents_tools::CalculatorTool))
23631 .unwrap();
23632 AgentBuilder::new()
23633 .system_prompt("You are a calculator assistant.")
23634 .llm(Arc::new(mock))
23635 .tools(tools)
23636 .build()
23637 .unwrap()
23638 }
23639
23640 #[tokio::test]
23641 async fn test_stream_tool_start_events_precede_results_for_batch() {
23642 let mock = mock_with_responses(vec![
23643 r#"[{"tool": "calculator", "arguments": {"expression": "1+1"}}, {"tool": "calculator", "arguments": {"expression": "2+2"}}]"#,
23644 "Both answers are ready.",
23645 ]);
23646 let agent = calculator_agent_with(mock);
23647
23648 let (_, chunks, final_response) = collect_stream_events(&agent, "compute both").await;
23649
23650 assert!(!chunks.iter().any(StreamChunk::is_error), "{chunks:?}");
23651 let final_response = final_response.expect("Final");
23652 assert_eq!(final_response.tool_calls.as_ref().map(Vec::len), Some(2));
23653
23654 let tool_events: Vec<&StreamChunk> = chunks
23655 .iter()
23656 .filter(|c| {
23657 matches!(
23658 c,
23659 StreamChunk::ToolCallStart { .. }
23660 | StreamChunk::ToolResult { .. }
23661 | StreamChunk::ToolCallEnd { .. }
23662 )
23663 })
23664 .collect();
23665 assert_eq!(tool_events.len(), 6, "{tool_events:?}");
23666 assert!(matches!(tool_events[0], StreamChunk::ToolCallStart { .. }));
23668 assert!(matches!(tool_events[1], StreamChunk::ToolCallStart { .. }));
23669 assert!(matches!(
23670 tool_events[2],
23671 StreamChunk::ToolResult { success: true, .. }
23672 ));
23673 assert!(matches!(tool_events[3], StreamChunk::ToolCallEnd { .. }));
23674 assert!(matches!(
23675 tool_events[4],
23676 StreamChunk::ToolResult { success: true, .. }
23677 ));
23678 assert!(matches!(tool_events[5], StreamChunk::ToolCallEnd { .. }));
23679 }
23680
23681 #[tokio::test]
23682 async fn test_stream_clarification_final_carries_options_and_detection() {
23683 let responses = || {
23684 vec![
23685 r#"{"is_ambiguous":true,"confidence":0.2,"ambiguity_type":"missing_target","reasoning":"target missing","what_is_unclear":["target"],"detected_language":"en"}"#,
23686 r#"{"question":"What should I send?","options":["report","invoice"]}"#,
23687 ]
23688 };
23689 let (blocking_agent, _) = state_disambiguation_agent(responses(), true, None, true);
23690 let (streaming_agent, _) = state_disambiguation_agent(responses(), true, None, true);
23691
23692 let blocking = blocking_agent.chat("Send it").await.unwrap();
23693 let (_, chunks, streamed) = collect_stream_events(&streaming_agent, "Send it").await;
23694 assert!(!chunks.iter().any(StreamChunk::is_error), "{chunks:?}");
23695 let streamed = streamed.expect("clarification must finalize as Final");
23696
23697 assert_eq!(streamed.content, "What should I send?");
23698 let streamed_meta = streamed
23699 .metadata
23700 .as_ref()
23701 .and_then(|m| m.get("disambiguation"))
23702 .cloned()
23703 .expect("disambiguation metadata");
23704 for key in ["status", "options", "clarifying", "detection"] {
23705 assert!(
23706 streamed_meta.get(key).is_some(),
23707 "missing {key}: {streamed_meta}"
23708 );
23709 }
23710 assert_eq!(
23711 streamed_meta.get("detection").and_then(|d| d.get("type")),
23712 Some(&serde_json::json!("missing_target"))
23713 );
23714 assert_eq!(
23715 blocking
23716 .metadata
23717 .as_ref()
23718 .and_then(|m| m.get("disambiguation")),
23719 Some(&streamed_meta),
23720 "blocking and streaming clarification metadata must be identical"
23721 );
23722 }
23723
23724 struct FailingMemory {
23726 messages: parking_lot::RwLock<Vec<ChatMessage>>,
23727 fail_on_add: usize,
23728 adds: std::sync::atomic::AtomicUsize,
23729 }
23730
23731 #[async_trait]
23732 impl ai_agents_core::Memory for FailingMemory {
23733 async fn add_message(&self, message: ChatMessage) -> Result<()> {
23734 let n = self.adds.fetch_add(1, Ordering::SeqCst) + 1;
23735 if n == self.fail_on_add {
23736 return Err(AgentError::Other(format!(
23737 "simulated memory failure on add #{n}"
23738 )));
23739 }
23740 self.messages.write().push(message);
23741 Ok(())
23742 }
23743
23744 async fn get_messages(&self, limit: Option<usize>) -> Result<Vec<ChatMessage>> {
23745 let messages = self.messages.read();
23746 Ok(match limit {
23747 Some(n) if n < messages.len() => messages[messages.len() - n..].to_vec(),
23748 _ => messages.clone(),
23749 })
23750 }
23751
23752 async fn clear(&self) -> Result<()> {
23753 self.messages.write().clear();
23754 Ok(())
23755 }
23756
23757 fn len(&self) -> usize {
23758 self.messages.read().len()
23759 }
23760
23761 async fn restore(&self, snapshot: ai_agents_core::MemorySnapshot) -> Result<()> {
23762 *self.messages.write() = snapshot.messages;
23763 Ok(())
23764 }
23765 }
23766
23767 impl ai_agents_memory::Memory for FailingMemory {}
23768
23769 #[tokio::test]
23770 async fn test_stream_memory_write_failure_surfaces_as_error() {
23771 let yaml = r#"
23774name: TransitionOnToolCallAgent
23775system_prompt: "You are helpful."
23776streaming:
23777 enabled: true
23778states:
23779 initial: intake
23780 states:
23781 intake:
23782 prompt: "Intake"
23783 transitions:
23784 - to: billing
23785 guard:
23786 context:
23787 route:
23788 eq: billing
23789 billing:
23790 prompt: "Billing"
23791"#;
23792 let build = |fail_on_add: usize| {
23793 let mut tools = ai_agents_tools::ToolRegistry::new();
23794 tools
23795 .register(Arc::new(ai_agents_tools::CalculatorTool))
23796 .unwrap();
23797 let agent = AgentBuilder::from_yaml(yaml)
23798 .unwrap()
23799 .llm(Arc::new(mock_with_responses(vec![
23800 r#"{"tool": "calculator", "arguments": {"expression": "1+1"}}"#,
23801 "Billing answer",
23802 ])))
23803 .tools(tools)
23804 .memory(Arc::new(FailingMemory {
23805 messages: parking_lot::RwLock::new(Vec::new()),
23806 fail_on_add,
23807 adds: std::sync::atomic::AtomicUsize::new(0),
23808 }))
23809 .build()
23810 .unwrap();
23811 agent
23812 .set_context("route", serde_json::json!("billing"))
23813 .unwrap();
23814 agent
23815 };
23816
23817 let blocking = build(2).chat("compute").await;
23818 assert!(
23819 blocking.is_err(),
23820 "blocking must surface the memory failure"
23821 );
23822
23823 let (_, chunks, final_response) = collect_stream_events(&build(2), "compute").await;
23824 assert!(
23825 final_response.is_none(),
23826 "streaming must not finalize after a memory failure"
23827 );
23828 assert!(
23829 chunks.iter().any(|c| matches!(c, StreamChunk::Error { message } if message.contains("simulated memory failure"))),
23830 "streaming must surface the memory failure: {chunks:?}"
23831 );
23832
23833 assert!(build(usize::MAX).chat("compute").await.is_ok());
23835 }
23836
23837 #[tokio::test]
23838 async fn parity_tool_execution() {
23839 let build = || {
23840 calculator_agent_with(mock_with_responses(vec![
23841 r#"{"tool": "calculator", "arguments": {"expression": "2+2"}}"#,
23842 "The answer is 4.",
23843 ]))
23844 };
23845 let (blocking, _, chunks) = assert_blocking_streaming_parity(build, "What is 2+2?").await;
23846 assert_eq!(blocking.content, "The answer is 4.");
23847 assert!(
23848 chunks
23849 .iter()
23850 .any(|c| matches!(c, StreamChunk::ToolResult { .. }))
23851 );
23852 }
23853
23854 #[tokio::test]
23855 async fn parity_disambiguation_clarification() {
23856 let build = || {
23857 state_disambiguation_agent(
23858 vec![
23859 r#"{"is_ambiguous":true,"confidence":0.2,"ambiguity_type":"missing_target","reasoning":"target missing","what_is_unclear":["target"],"detected_language":"en"}"#,
23860 r#"{"question":"What should I send?","options":null}"#,
23861 ],
23862 true,
23863 None,
23864 true,
23865 )
23866 .0
23867 };
23868 let (blocking, _, _) = assert_blocking_streaming_parity(build, "Send it").await;
23869 assert_eq!(blocking.content, "What should I send?");
23870 }
23871
23872 #[tokio::test]
23873 async fn runtime_disambiguation_uses_configured_recent_history_projection() {
23874 let yaml = r#"
23875name: DisambiguationContextAgent
23876system_prompt: "Help."
23877llm:
23878 default: default
23879 router: router
23880disambiguation:
23881 enabled: true
23882 detection:
23883 llm: router
23884 context:
23885 recent_messages: 1
23886 include_state: false
23887 include_available_tools: false
23888states:
23889 initial: private_state
23890 states:
23891 private_state:
23892 prompt: "PRIVATE_STATE_PROMPT"
23893"#;
23894 let main = mock_with_responses(vec!["FIRST_MAIN_MARKER", "SECOND_MAIN_MARKER"]);
23895 let router = mock_with_responses(vec![
23896 r#"{"is_ambiguous":false,"confidence":0.9,"ambiguity_type":null,"reasoning":"clear","what_is_unclear":[],"detected_language":"en"}"#,
23897 r#"{"is_ambiguous":false,"confidence":0.9,"ambiguity_type":null,"reasoning":"clear","what_is_unclear":[],"detected_language":"en"}"#,
23898 ]);
23899 let router_calls = router.clone();
23900 let agent = AgentBuilder::from_yaml(yaml)
23901 .unwrap()
23902 .llm_alias("default", Arc::new(main))
23903 .llm_alias("router", Arc::new(router))
23904 .build()
23905 .unwrap();
23906
23907 agent.chat("FIRST_USER_MARKER").await.unwrap();
23908 agent.chat("SECOND_USER_MARKER").await.unwrap();
23909
23910 let calls = router_calls.call_history();
23911 assert_eq!(calls.len(), 2);
23912 let second_prompt = &calls[1].messages.last().unwrap().content;
23913 assert!(
23914 second_prompt.contains("FIRST_MAIN_MARKER"),
23915 "{second_prompt}"
23916 );
23917 assert!(
23918 !second_prompt.contains("FIRST_USER_MARKER"),
23919 "{second_prompt}"
23920 );
23921 assert!(
23922 !second_prompt.contains("PRIVATE_STATE_PROMPT"),
23923 "{second_prompt}"
23924 );
23925 }
23926
23927 #[tokio::test]
23928 async fn parity_reflection_enabled() {
23929 let yaml = r#"
23930name: ReflectionAgent
23931system_prompt: "You are careful."
23932reflection:
23933 enabled: true
23934 criteria:
23935 - "Is the answer helpful?"
23936"#;
23937 let build = || {
23938 AgentBuilder::from_yaml(yaml)
23939 .unwrap()
23940 .llm(Arc::new(mock_with_responses(vec![
23941 "Main answer",
23942 "OVERALL: PASS\nCONFIDENCE: 0.9",
23943 ])))
23944 .build()
23945 .unwrap()
23946 };
23947 let (blocking, streamed, _) = assert_blocking_streaming_parity(build, "hello").await;
23948 assert_eq!(blocking.content, "Main answer");
23949 assert!(metadata_keys(&streamed).contains("reflection"));
23950 }
23951
23952 fn state_reflection_agent(
23953 global_retries: u32,
23954 state_retries: u32,
23955 main: MockLLMProvider,
23956 evaluator: MockLLMProvider,
23957 ) -> RuntimeAgent {
23958 let yaml = format!(
23959 r#"
23960name: StateReflectionAgent
23961system_prompt: "You are careful."
23962llm:
23963 default: default
23964 router: evaluator
23965reflection:
23966 enabled: true
23967 evaluator_llm: evaluator
23968 max_retries: {global_retries}
23969 criteria:
23970 - "Global criterion"
23971states:
23972 initial: active
23973 states:
23974 active:
23975 prompt: "Handle the active state."
23976 reflection:
23977 enabled: true
23978 evaluator_llm: evaluator
23979 max_retries: {state_retries}
23980 criteria:
23981 - "State criterion"
23982"#
23983 );
23984 AgentBuilder::from_yaml(&yaml)
23985 .unwrap()
23986 .llm_alias("default", Arc::new(main))
23987 .llm_alias("evaluator", Arc::new(evaluator))
23988 .build()
23989 .unwrap()
23990 }
23991
23992 #[tokio::test]
23993 async fn state_reflection_zero_retries_overrides_global_limit() {
23994 let main = mock_with_response("First answer");
23995 let main_calls = main.clone();
23996 let evaluator = mock_with_response("OVERALL: FAIL\nCONFIDENCE: 0.1");
23997 let evaluator_calls = evaluator.clone();
23998 let agent = state_reflection_agent(2, 0, main, evaluator);
23999
24000 let response = agent.chat("hello").await.unwrap();
24001 let reflection = response
24002 .metadata
24003 .as_ref()
24004 .and_then(|metadata| metadata.get("reflection"))
24005 .expect("reflection metadata");
24006
24007 assert_eq!(response.content, "First answer");
24008 assert_eq!(reflection["attempts"], 1);
24009 assert_eq!(main_calls.call_count(), 1);
24010 assert_eq!(evaluator_calls.call_count(), 1);
24011 }
24012
24013 #[tokio::test]
24014 async fn state_reflection_retry_limit_overrides_zero_global_limit() {
24015 let main = mock_with_responses(vec!["First answer", "Second answer", "Third answer"]);
24016 let main_calls = main.clone();
24017 let evaluator = mock_with_responses(vec![
24018 "OVERALL: FAIL\nCONFIDENCE: 0.1",
24019 "OVERALL: FAIL\nCONFIDENCE: 0.1",
24020 "OVERALL: FAIL\nCONFIDENCE: 0.1",
24021 ]);
24022 let evaluator_calls = evaluator.clone();
24023 let agent = state_reflection_agent(0, 2, main, evaluator);
24024
24025 let response = agent.chat("hello").await.unwrap();
24026 let reflection = response
24027 .metadata
24028 .as_ref()
24029 .and_then(|metadata| metadata.get("reflection"))
24030 .expect("reflection metadata");
24031
24032 assert_eq!(response.content, "Third answer");
24033 assert_eq!(reflection["attempts"], 3);
24034 assert_eq!(main_calls.call_count(), 3);
24035 assert_eq!(evaluator_calls.call_count(), 3);
24036 }
24037
24038 #[tokio::test]
24039 async fn state_reflection_override_is_preserved_in_event_stream_metadata() {
24040 let main = mock_with_responses(vec!["First answer", "Second answer", "Third answer"]);
24041 let main_calls = main.clone();
24042 let evaluator = mock_with_responses(vec![
24043 "OVERALL: FAIL\nCONFIDENCE: 0.1",
24044 "OVERALL: FAIL\nCONFIDENCE: 0.1",
24045 "OVERALL: FAIL\nCONFIDENCE: 0.1",
24046 ]);
24047 let evaluator_calls = evaluator.clone();
24048 let agent = state_reflection_agent(0, 2, main, evaluator);
24049
24050 let (content, chunks, final_response) = collect_stream_events(&agent, "hello").await;
24051 let final_response = final_response.expect("successful stream Final");
24052 let reflection = final_response
24053 .metadata
24054 .as_ref()
24055 .and_then(|metadata| metadata.get("reflection"))
24056 .expect("reflection metadata");
24057
24058 assert_eq!(content, "Third answer");
24059 assert!(!chunks.iter().any(StreamChunk::is_error));
24060 assert_eq!(reflection["attempts"], 3);
24061 assert_eq!(main_calls.call_count(), 3);
24062 assert_eq!(evaluator_calls.call_count(), 3);
24063 }
24064
24065 #[tokio::test]
24066 async fn state_reflection_override_preserves_legacy_stream_completion() {
24067 use futures::StreamExt;
24068
24069 let main = mock_with_response("First answer");
24070 let main_calls = main.clone();
24071 let evaluator = mock_with_response("OVERALL: FAIL\nCONFIDENCE: 0.1");
24072 let evaluator_calls = evaluator.clone();
24073 let agent = state_reflection_agent(2, 0, main, evaluator);
24074 let mut stream = agent.chat_stream("hello").await.unwrap();
24075 let mut content = String::new();
24076 let mut done = 0;
24077 while let Some(chunk) = stream.next().await {
24078 match chunk {
24079 StreamChunk::Content { text } => content.push_str(&text),
24080 StreamChunk::Done {} => done += 1,
24081 StreamChunk::Error { message } => panic!("unexpected error: {message}"),
24082 _ => {}
24083 }
24084 }
24085
24086 assert_eq!(content, "First answer");
24087 assert_eq!(done, 1);
24088 assert_eq!(main_calls.call_count(), 1);
24089 assert_eq!(evaluator_calls.call_count(), 1);
24090 }
24091
24092 #[tokio::test]
24093 async fn reflection_evaluator_error_emits_no_event_stream_final() {
24094 let main = mock_with_response("First answer");
24095 let mut evaluator = MockLLMProvider::new("evaluator");
24096 evaluator.set_error("judge failed");
24097 let agent = state_reflection_agent(0, 2, main, evaluator);
24098
24099 let (_, chunks, final_response) = collect_stream_events(&agent, "hello").await;
24100
24101 assert!(final_response.is_none());
24102 assert!(chunks.iter().any(StreamChunk::is_error));
24103 assert!(!chunks.iter().any(StreamChunk::is_done));
24104 }
24105
24106 #[tokio::test]
24107 async fn parity_cot_hidden_thinking() {
24108 let yaml = r#"
24109name: CotHiddenAgent
24110system_prompt: "Think first."
24111reasoning:
24112 mode: cot
24113 output: hidden
24114"#;
24115 let build = || {
24116 AgentBuilder::from_yaml(yaml)
24117 .unwrap()
24118 .llm(Arc::new(mock_with_response(
24119 "<thinking>step by step</thinking>Visible answer",
24120 )))
24121 .build()
24122 .unwrap()
24123 };
24124 let (blocking, streamed, chunks) = assert_blocking_streaming_parity(build, "hello").await;
24125 assert_eq!(blocking.content, "Visible answer");
24126 assert_eq!(content_chunks(&chunks).concat(), streamed.content);
24128 }
24129
24130 fn content_chunks(chunks: &[StreamChunk]) -> Vec<String> {
24135 chunks
24136 .iter()
24137 .filter_map(|c| match c {
24138 StreamChunk::Content { text } => Some(text.clone()),
24139 _ => None,
24140 })
24141 .collect()
24142 }
24143
24144 fn reflection_auto_agent(main: MockLLMProvider, judge: MockLLMProvider) -> RuntimeAgent {
24145 let yaml = r#"
24146name: ReflectionAutoAgent
24147system_prompt: "You are careful."
24148llm:
24149 default: default
24150 router: router
24151reflection:
24152 enabled: auto
24153 evaluator_llm: router
24154 criteria:
24155 - "Is the answer helpful?"
24156"#;
24157 AgentBuilder::from_yaml(yaml)
24158 .unwrap()
24159 .llm_alias("default", Arc::new(main))
24160 .llm_alias("router", Arc::new(judge))
24161 .build()
24162 .unwrap()
24163 }
24164
24165 #[tokio::test]
24166 async fn test_stream_reflection_auto_buffers_and_calls_judge_once_per_iteration() {
24167 let judge = mock_with_responses(vec!["YES", "OVERALL: PASS\nCONFIDENCE: 0.9"]);
24168 let judge_calls = judge.clone();
24169 let agent = reflection_auto_agent(mock_with_response("Main answer one two"), judge);
24170
24171 let (content, chunks, final_response) = collect_stream_events(&agent, "hello").await;
24172
24173 assert!(!chunks.iter().any(StreamChunk::is_error), "{chunks:?}");
24174 assert_eq!(
24175 content_chunks(&chunks).len(),
24176 1,
24177 "auto reflection must buffer the main response: {chunks:?}"
24178 );
24179 assert_eq!(content, "Main answer one two");
24180 assert_eq!(final_response.expect("Final").content, content);
24181 assert_eq!(judge_calls.call_count(), 2);
24183 }
24184
24185 #[tokio::test]
24186 async fn test_stream_reflection_auto_rewrite_is_streamed() {
24187 let judge = mock_with_responses(vec![
24188 "YES",
24189 "OVERALL: FAIL\nCONFIDENCE: 0.1",
24190 "OVERALL: PASS\nCONFIDENCE: 0.9",
24191 ]);
24192 let agent = reflection_auto_agent(
24193 mock_with_responses(vec!["First attempt", "Improved answer"]),
24194 judge,
24195 );
24196
24197 let (content, chunks, final_response) = collect_stream_events(&agent, "hello").await;
24198
24199 assert!(!chunks.iter().any(StreamChunk::is_error), "{chunks:?}");
24200 assert_eq!(content_chunks(&chunks).len(), 1, "{chunks:?}");
24201 assert_eq!(
24202 content, "Improved answer",
24203 "the rewritten answer is what streams"
24204 );
24205 assert_eq!(final_response.expect("Final").content, "Improved answer");
24206 }
24207
24208 fn reasoning_agent(mode: &str, output: &str) -> RuntimeAgent {
24209 let yaml = format!(
24210 r#"
24211name: ReasoningStreamAgent
24212system_prompt: "Think first."
24213reasoning:
24214 mode: {mode}
24215 output: {output}
24216"#
24217 );
24218 AgentBuilder::from_yaml(&yaml)
24219 .unwrap()
24220 .llm(Arc::new(mock_with_response(
24221 "<thinking>step by step</thinking>Visible answer",
24222 )))
24223 .build()
24224 .unwrap()
24225 }
24226
24227 #[tokio::test]
24228 async fn test_stream_cot_hidden_emits_no_thinking_tags() {
24229 let agent = reasoning_agent("cot", "hidden");
24230 let (content, chunks, final_response) = collect_stream_events(&agent, "hello").await;
24231 assert!(!chunks.iter().any(StreamChunk::is_error), "{chunks:?}");
24232 assert_eq!(content_chunks(&chunks).len(), 1, "{chunks:?}");
24233 assert!(!content.contains("<thinking>"), "{content:?}");
24234 assert_eq!(content, "Visible answer");
24235 assert_eq!(final_response.expect("Final").content, content);
24236 }
24237
24238 #[tokio::test]
24239 async fn test_stream_cot_visible_matches_final_format() {
24240 let agent = reasoning_agent("cot", "visible");
24241 let (content, chunks, final_response) = collect_stream_events(&agent, "hello").await;
24242 assert!(!chunks.iter().any(StreamChunk::is_error), "{chunks:?}");
24243 assert!(content.starts_with("Thinking:"), "{content:?}");
24244 assert!(content.contains("Answer:\nVisible answer"), "{content:?}");
24245 assert_eq!(final_response.expect("Final").content, content);
24246 }
24247
24248 #[tokio::test]
24249 async fn test_stream_react_buffers() {
24250 let agent = reasoning_agent("react", "hidden");
24251 let (content, chunks, final_response) = collect_stream_events(&agent, "hello").await;
24252 assert!(!chunks.iter().any(StreamChunk::is_error), "{chunks:?}");
24253 assert_eq!(content_chunks(&chunks).len(), 1, "{chunks:?}");
24254 assert_eq!(content, "Visible answer");
24255 assert_eq!(final_response.expect("Final").content, content);
24256 }
24257
24258 #[tokio::test]
24259 async fn test_stream_plain_mode_still_streams_deltas() {
24260 let agent = AgentBuilder::new()
24261 .system_prompt("You are helpful.")
24262 .llm(Arc::new(mock_with_response("one two three")))
24263 .build()
24264 .unwrap();
24265 let (content, chunks, final_response) = collect_stream_events(&agent, "hello").await;
24266 assert!(
24267 content_chunks(&chunks).len() >= 2,
24268 "plain turns must keep token-level streaming: {chunks:?}"
24269 );
24270 assert_eq!(content, "one two three");
24271 assert_eq!(final_response.expect("Final").content, content);
24272 }
24273
24274 struct ActorProbeHooks {
24276 seen: parking_lot::Mutex<Option<crate::TurnActorContext>>,
24277 }
24278
24279 #[async_trait]
24280 impl AgentHooks for ActorProbeHooks {
24281 async fn on_message_received(&self, _input: &str) {
24282 *self.seen.lock() = current_turn_actor_context();
24283 }
24284 }
24285
24286 fn actor_probe_agent(hooks: Arc<ActorProbeHooks>) -> RuntimeAgent {
24287 let yaml = r#"
24288name: ActorStreamAgent
24289system_prompt: "You are helpful."
24290observability:
24291 enabled: true
24292 export:
24293 write_raw_events: true
24294"#;
24295 AgentBuilder::from_yaml(yaml)
24296 .unwrap()
24297 .llm(Arc::new(mock_with_response("Hello actor")))
24298 .hooks(hooks)
24299 .build()
24300 .unwrap()
24301 }
24302
24303 async fn collect_actor_stream_final(
24304 agent: &RuntimeAgent,
24305 input: &str,
24306 actor_context: crate::TurnActorContext,
24307 ) -> AgentResponse {
24308 use futures::StreamExt;
24309 let mut events = agent
24310 .chat_stream_events_with_actor_context(input, actor_context)
24311 .await
24312 .expect("stream opens");
24313 let mut final_response = None;
24314 while let Some(event) = events.next().await {
24315 match event {
24316 AgentStreamEvent::Final(response) => final_response = Some(response),
24317 AgentStreamEvent::Chunk(StreamChunk::Error { message }) => {
24318 panic!("unexpected stream error: {message}")
24319 }
24320 AgentStreamEvent::Chunk(_) => {}
24321 }
24322 }
24323 final_response.expect("Final")
24324 }
24325
24326 #[tokio::test]
24327 async fn test_stream_events_with_actor_context_scopes_actor_for_turn() {
24328 let hooks = Arc::new(ActorProbeHooks {
24329 seen: parking_lot::Mutex::new(None),
24330 });
24331 let agent = actor_probe_agent(Arc::clone(&hooks));
24332 let actor_context = crate::TurnActorContext::new().with_origin_actor("customer_42");
24333
24334 let final_response = collect_actor_stream_final(&agent, "hi", actor_context).await;
24335
24336 assert_eq!(final_response.content, "Hello actor");
24337 assert_eq!(
24338 hooks
24339 .seen
24340 .lock()
24341 .as_ref()
24342 .and_then(|context| context.effective_actor_id().map(str::to_string)),
24343 Some("customer_42".to_string()),
24344 "the actor context must be visible inside the streaming turn"
24345 );
24346 assert!(
24347 agent.actor_id().is_none(),
24348 "a turn-scoped actor must not mutate the global actor ID"
24349 );
24350 let events = agent.observability().unwrap().raw_events();
24351 assert!(
24352 events
24353 .iter()
24354 .any(|event| event.dimensions.get("actor") == Some(&"customer_42".to_string())),
24355 "observation events must carry the actor dimension"
24356 );
24357 }
24358
24359 #[tokio::test]
24360 async fn test_stream_events_with_actor_context_matches_blocking_actor_context() {
24361 let actor_context = crate::TurnActorContext::new()
24362 .with_origin_actor("customer_42")
24363 .with_sender_agent("coordinator");
24364
24365 let blocking_hooks = Arc::new(ActorProbeHooks {
24366 seen: parking_lot::Mutex::new(None),
24367 });
24368 let blocking_agent = actor_probe_agent(Arc::clone(&blocking_hooks));
24369 let blocking = blocking_agent
24370 .chat_with_actor_context("hi", actor_context.clone())
24371 .await
24372 .unwrap();
24373
24374 let streaming_hooks = Arc::new(ActorProbeHooks {
24375 seen: parking_lot::Mutex::new(None),
24376 });
24377 let streaming_agent = actor_probe_agent(Arc::clone(&streaming_hooks));
24378 let streamed =
24379 collect_actor_stream_final(&streaming_agent, "hi", actor_context.clone()).await;
24380
24381 assert_eq!(blocking.content, streamed.content);
24382 assert_eq!(metadata_keys(&blocking), metadata_keys(&streamed));
24383 assert_eq!(
24384 *blocking_hooks.seen.lock(),
24385 *streaming_hooks.seen.lock(),
24386 "both entry points must expose the same turn actor context"
24387 );
24388 assert_eq!(*streaming_hooks.seen.lock(), Some(actor_context));
24389 }
24390
24391 #[tokio::test]
24392 async fn test_stream_events_with_actor_context_releases_root_turn_on_drop() {
24393 use futures::StreamExt;
24394 let agent = AgentBuilder::new()
24395 .system_prompt("You are helpful.")
24396 .llm(Arc::new(mock_with_response("one two three")))
24397 .build()
24398 .unwrap();
24399 {
24400 let mut events = agent
24401 .chat_stream_events_with_actor_context(
24402 "hi",
24403 crate::TurnActorContext::new().with_origin_actor("customer_42"),
24404 )
24405 .await
24406 .unwrap();
24407 let _first = events.next().await;
24409 }
24410 let next = tokio::time::timeout(Duration::from_secs(5), agent.chat("next")).await;
24411 assert!(
24412 matches!(next, Ok(Ok(_))),
24413 "the root turn must be released when the actor stream is dropped: {next:?}"
24414 );
24415 }
24416
24417 fn skill_scope_agent_with_router(states: &str) -> (RuntimeAgent, MockLLMProvider) {
24418 let yaml = format!(
24419 r#"
24420name: SkillScopeAgent
24421system_prompt: "Route skills."
24422skills:
24423 - id: alpha
24424 description: "Alpha"
24425 trigger: "alpha"
24426 steps:
24427 - prompt: "alpha {{{{ user_input }}}}"
24428 - id: beta
24429 description: "Beta"
24430 trigger: "beta"
24431 steps:
24432 - prompt: "beta {{{{ user_input }}}}"
24433{states}
24434"#
24435 );
24436 let router = mock_with_response("none");
24437 let calls = router.clone();
24438 let agent = AgentBuilder::from_yaml(&yaml)
24439 .unwrap()
24440 .llm(Arc::new(router))
24441 .build()
24442 .unwrap();
24443 (agent, calls)
24444 }
24445
24446 fn skill_scope_agent(states: &str) -> RuntimeAgent {
24447 skill_scope_agent_with_router(states).0
24448 }
24449
24450 fn available_skill_ids(agent: &RuntimeAgent) -> Vec<String> {
24451 let mut ids: Vec<String> = agent
24452 .get_available_skills()
24453 .into_iter()
24454 .map(|skill| skill.id.clone())
24455 .collect();
24456 ids.sort();
24457 ids
24458 }
24459
24460 #[test]
24461 fn state_skill_scope_characterizes_empty_inheritance_and_unknown_ids() {
24462 assert_eq!(
24463 available_skill_ids(&skill_scope_agent("")),
24464 vec!["alpha", "beta"]
24465 );
24466 assert_eq!(
24467 available_skill_ids(&skill_scope_agent(
24468 "states:\n initial: current\n states:\n current:\n skills: []\n"
24469 )),
24470 vec!["alpha", "beta"]
24471 );
24472 assert!(
24473 available_skill_ids(&skill_scope_agent(
24474 "states:\n initial: current\n states:\n current:\n skills: [unknown]\n"
24475 ))
24476 .is_empty()
24477 );
24478 assert_eq!(
24479 available_skill_ids(&skill_scope_agent(
24480 "states:\n initial: parent\n states:\n parent:\n skills: [alpha]\n initial: child\n states:\n child:\n skills: []\n"
24481 )),
24482 vec!["alpha"]
24483 );
24484 assert_eq!(
24485 available_skill_ids(&skill_scope_agent(
24486 "states:\n initial: parent\n states:\n parent:\n skills: [alpha]\n initial: child\n states:\n child:\n skills: [beta]\n"
24487 )),
24488 vec!["alpha", "beta"]
24489 );
24490 assert_eq!(
24491 available_skill_ids(&skill_scope_agent(
24492 "states:\n initial: parent\n states:\n parent:\n skills: [alpha]\n initial: child\n states:\n child:\n inherit_parent: false\n skills: []\n"
24493 )),
24494 vec!["alpha", "beta"]
24495 );
24496 }
24497
24498 #[tokio::test]
24499 async fn state_skill_scope_reaches_the_router_candidate_prompt() {
24500 let (inherited, inherited_calls) = skill_scope_agent_with_router(
24501 "states:\n initial: parent\n states:\n parent:\n skills: [alpha]\n initial: child\n states:\n child:\n skills: []\n",
24502 );
24503 assert!(
24504 inherited
24505 .select_skill_candidate("route")
24506 .await
24507 .unwrap()
24508 .is_none()
24509 );
24510 let inherited_call = inherited_calls.last_call().unwrap();
24511 let inherited_prompt = &inherited_call.messages[0].content;
24512 assert!(inherited_prompt.contains("- alpha:"));
24513 assert!(!inherited_prompt.contains("- beta:"));
24514
24515 let (fallback_all, fallback_calls) = skill_scope_agent_with_router(
24516 "states:\n initial: current\n states:\n current:\n skills: []\n",
24517 );
24518 assert!(
24519 fallback_all
24520 .select_skill_candidate("route")
24521 .await
24522 .unwrap()
24523 .is_none()
24524 );
24525 let fallback_call = fallback_calls.last_call().unwrap();
24526 let fallback_prompt = &fallback_call.messages[0].content;
24527 assert!(fallback_prompt.contains("- alpha:"));
24528 assert!(fallback_prompt.contains("- beta:"));
24529
24530 let (unknown, unknown_calls) = skill_scope_agent_with_router(
24531 "states:\n initial: current\n states:\n current:\n skills: [unknown]\n",
24532 );
24533 assert!(
24534 unknown
24535 .select_skill_candidate("route")
24536 .await
24537 .unwrap()
24538 .is_none()
24539 );
24540 assert_eq!(unknown_calls.call_count(), 0);
24541 }
24542
24543 #[tokio::test]
24544 async fn parity_skill_route() {
24545 let yaml = skills_with_parallel_transition_yaml(
24546 " max_speculative_llm_calls_per_turn: 2\n max_parallel_runtime_tasks: 2",
24547 "",
24548 );
24549 let build = || {
24550 build_skills_beside_transition_agent(
24551 &yaml,
24552 role_mocks(
24553 mock_with_response("Draft response"),
24554 mock_with_response("helper"),
24555 ),
24556 )
24557 };
24558 assert_blocking_streaming_parity(build, "please use helper").await;
24559 }
24560
24561 fn required_context_agent(mock: MockLLMProvider, default: bool) -> RuntimeAgent {
24563 let default_yaml = if default {
24564 " default:\n brief: fallback\n"
24565 } else {
24566 ""
24567 };
24568 let yaml = format!(
24569 "name: RequiredContextAgent\nsystem_prompt: 'Voice: {{{{ context.voice.brief }}}}'\ncontext:\n voice:\n type: runtime\n required: true\n{default_yaml}"
24570 );
24571 AgentBuilder::from_yaml(&yaml)
24572 .unwrap()
24573 .llm(Arc::new(mock))
24574 .build()
24575 .unwrap()
24576 }
24577
24578 struct CountingContextProvider {
24579 marker: &'static str,
24580 calls: std::sync::atomic::AtomicUsize,
24581 }
24582
24583 #[async_trait]
24584 impl ContextProvider for CountingContextProvider {
24585 async fn get(&self, _key: &str, _current_context: &Value) -> Result<Value> {
24586 let call = self.calls.fetch_add(1, Ordering::SeqCst) + 1;
24587 Ok(serde_json::json!({"call": call, "marker": self.marker}))
24588 }
24589 }
24590
24591 struct FailOnceContextProvider {
24592 attempts: std::sync::atomic::AtomicUsize,
24593 }
24594
24595 #[async_trait]
24596 impl ContextProvider for FailOnceContextProvider {
24597 async fn get(&self, _key: &str, _current_context: &Value) -> Result<Value> {
24598 if self.attempts.fetch_add(1, Ordering::SeqCst) == 0 {
24599 return Err(AgentError::Other("context initialization failed".into()));
24600 }
24601 Ok(serde_json::json!({"brief": "ready"}))
24602 }
24603 }
24604
24605 fn session_context_agent(
24606 provider: Arc<CountingContextProvider>,
24607 refresh: &str,
24608 ) -> RuntimeAgent {
24609 let yaml = format!(
24610 "name: SessionContextAgent\nsystem_prompt: 'Call: {{{{ context.session_data.call }}}}'\ncontext:\n session_data:\n type: callback\n name: counter\n refresh: {refresh}\n"
24611 );
24612 let agent = AgentBuilder::from_yaml(&yaml)
24613 .unwrap()
24614 .llm(Arc::new(mock_with_response("ok")))
24615 .build()
24616 .unwrap();
24617 agent.register_context_provider("counter", provider);
24618 agent
24619 }
24620
24621 #[tokio::test]
24622 async fn session_context_characterizes_reset_and_restore_lifecycle() {
24623 let original_provider = Arc::new(CountingContextProvider {
24624 marker: "original",
24625 calls: std::sync::atomic::AtomicUsize::new(0),
24626 });
24627 let original = session_context_agent(original_provider.clone(), "per_session");
24628 original.chat("first").await.unwrap();
24629 original.chat("second").await.unwrap();
24630 assert_eq!(original_provider.calls.load(Ordering::SeqCst), 1);
24631 assert_eq!(original.get_context()["session_data"]["call"], 1);
24632 assert_eq!(original.get_context()["session_data"]["marker"], "original");
24633
24634 original.reset().await.unwrap();
24635 original.chat("after reset").await.unwrap();
24636 assert_eq!(original_provider.calls.load(Ordering::SeqCst), 1);
24637 assert_eq!(original.get_context()["session_data"]["call"], 1);
24638 assert_eq!(original.get_context()["session_data"]["marker"], "original");
24639 let snapshot = original.save_state().await.unwrap();
24640
24641 let fresh_provider = Arc::new(CountingContextProvider {
24642 marker: "fresh",
24643 calls: std::sync::atomic::AtomicUsize::new(0),
24644 });
24645 let fresh = session_context_agent(fresh_provider.clone(), "per_session");
24646 fresh.restore_state(snapshot.clone()).await.unwrap();
24647 fresh.chat("fresh restore").await.unwrap();
24648 assert_eq!(fresh_provider.calls.load(Ordering::SeqCst), 1);
24649 assert_eq!(fresh.get_context()["session_data"]["call"], 1);
24650 assert_eq!(fresh.get_context()["session_data"]["marker"], "fresh");
24651
24652 let warm_provider = Arc::new(CountingContextProvider {
24653 marker: "warm",
24654 calls: std::sync::atomic::AtomicUsize::new(0),
24655 });
24656 let warm = session_context_agent(warm_provider.clone(), "per_session");
24657 warm.chat("warmup").await.unwrap();
24658 assert_eq!(warm_provider.calls.load(Ordering::SeqCst), 1);
24659 warm.restore_state(snapshot).await.unwrap();
24660 warm.chat("warm restore").await.unwrap();
24661 assert_eq!(warm_provider.calls.load(Ordering::SeqCst), 1);
24662 assert_eq!(warm.get_context()["session_data"]["call"], 1);
24663 assert_eq!(warm.get_context()["session_data"]["marker"], "original");
24664 }
24665
24666 #[tokio::test]
24667 async fn once_context_is_not_refreshed_by_later_turns_or_reset() {
24668 let provider = Arc::new(CountingContextProvider {
24669 marker: "once",
24670 calls: std::sync::atomic::AtomicUsize::new(0),
24671 });
24672 let agent = session_context_agent(provider.clone(), "once");
24673
24674 agent.chat("first").await.unwrap();
24675 agent.chat("second").await.unwrap();
24676 agent.reset().await.unwrap();
24677 agent.chat("after reset").await.unwrap();
24678
24679 assert_eq!(provider.calls.load(Ordering::SeqCst), 1);
24680 assert_eq!(agent.get_context()["session_data"]["marker"], "once");
24681 }
24682
24683 #[tokio::test]
24684 async fn per_turn_context_refreshes_after_initialization_and_restore() {
24685 let original_provider = Arc::new(CountingContextProvider {
24686 marker: "original-per-turn",
24687 calls: std::sync::atomic::AtomicUsize::new(0),
24688 });
24689 let original = session_context_agent(original_provider.clone(), "per_turn");
24690 original.chat("first").await.unwrap();
24691 assert_eq!(original_provider.calls.load(Ordering::SeqCst), 2);
24692 assert_eq!(original.get_context()["session_data"]["call"], 2);
24693 original.chat("second").await.unwrap();
24694 assert_eq!(original_provider.calls.load(Ordering::SeqCst), 3);
24695 original.reset().await.unwrap();
24696 original.chat("after reset").await.unwrap();
24697 assert_eq!(original_provider.calls.load(Ordering::SeqCst), 4);
24698 let snapshot = original.save_state().await.unwrap();
24699
24700 let fresh_provider = Arc::new(CountingContextProvider {
24701 marker: "fresh-per-turn",
24702 calls: std::sync::atomic::AtomicUsize::new(0),
24703 });
24704 let fresh = session_context_agent(fresh_provider.clone(), "per_turn");
24705 fresh.restore_state(snapshot.clone()).await.unwrap();
24706 fresh.chat("fresh restore").await.unwrap();
24707 assert_eq!(fresh_provider.calls.load(Ordering::SeqCst), 2);
24708 assert_eq!(
24709 fresh.get_context()["session_data"]["marker"],
24710 "fresh-per-turn"
24711 );
24712 assert_eq!(fresh.get_context()["session_data"]["call"], 2);
24713
24714 let warm_provider = Arc::new(CountingContextProvider {
24715 marker: "warm-per-turn",
24716 calls: std::sync::atomic::AtomicUsize::new(0),
24717 });
24718 let warm = session_context_agent(warm_provider.clone(), "per_turn");
24719 warm.chat("warmup").await.unwrap();
24720 assert_eq!(warm_provider.calls.load(Ordering::SeqCst), 2);
24721 warm.restore_state(snapshot).await.unwrap();
24722 warm.chat("warm restore").await.unwrap();
24723 assert_eq!(warm_provider.calls.load(Ordering::SeqCst), 3);
24724 assert_eq!(
24725 warm.get_context()["session_data"]["marker"],
24726 "warm-per-turn"
24727 );
24728 assert_eq!(warm.get_context()["session_data"]["call"], 3);
24729 }
24730
24731 #[tokio::test]
24732 async fn test_context_initialization_retries_after_failure() {
24733 let mock = mock_with_response("Voice response");
24734 let calls = mock.clone();
24735 let yaml = "name: CallbackAgent\nsystem_prompt: 'Voice: {{ context.voice.brief }}'\ncontext:\n voice:\n type: callback\n name: flaky\n";
24736 let agent = AgentBuilder::from_yaml(yaml)
24737 .unwrap()
24738 .llm(Arc::new(mock))
24739 .build()
24740 .unwrap();
24741 let provider = Arc::new(FailOnceContextProvider {
24742 attempts: std::sync::atomic::AtomicUsize::new(0),
24743 });
24744 agent.register_context_provider("flaky", provider.clone());
24745
24746 assert!(agent.chat("first").await.is_err());
24747 assert_eq!(calls.call_count(), 0);
24748 assert_eq!(
24749 agent.chat("second").await.unwrap().content,
24750 "Voice response"
24751 );
24752 assert_eq!(provider.attempts.load(Ordering::SeqCst), 2);
24753 assert_eq!(calls.call_count(), 1);
24754 }
24755
24756 #[tokio::test]
24757 async fn test_required_context_blocks_chat_until_supplied_and_after_removal() {
24758 let mock = mock_with_response("Voice response");
24759 let calls = mock.clone();
24760 let agent = required_context_agent(mock, false);
24761
24762 let error = agent.chat("first").await.unwrap_err();
24763 assert!(
24764 error
24765 .to_string()
24766 .contains("Required context 'voice' not provided")
24767 );
24768 assert_eq!(calls.call_count(), 0);
24769 assert!(agent.memory.get_messages(None).await.unwrap().is_empty());
24770
24771 agent
24772 .set_context("voice.brief", serde_json::json!("ready"))
24773 .unwrap();
24774 assert_eq!(
24775 agent.chat("second").await.unwrap().content,
24776 "Voice response"
24777 );
24778 assert_eq!(calls.call_count(), 1);
24779
24780 agent.remove_context("voice");
24781 let error = agent.chat("third").await.unwrap_err();
24782 assert!(
24783 error
24784 .to_string()
24785 .contains("Required context 'voice' not provided")
24786 );
24787 assert_eq!(calls.call_count(), 1);
24788 }
24789
24790 #[tokio::test]
24791 async fn test_required_context_default_satisfies_presence_check() {
24792 let mock = mock_with_response("Fallback response");
24793 let calls = mock.clone();
24794 let agent = required_context_agent(mock, true);
24795
24796 assert_eq!(
24797 agent.chat("hello").await.unwrap().content,
24798 "Fallback response"
24799 );
24800 assert_eq!(
24801 agent.context_manager().get_path("voice.brief"),
24802 Some(serde_json::json!("fallback"))
24803 );
24804 assert_eq!(calls.call_count(), 1);
24805 }
24806
24807 #[tokio::test]
24808 async fn test_required_context_blocks_legacy_stream_before_model_call() {
24809 use futures::StreamExt;
24810
24811 let mock = mock_with_response("Voice response");
24812 let calls = mock.clone();
24813 let agent = required_context_agent(mock, false);
24814 let mut stream = agent.chat_stream("first").await.unwrap();
24815 assert!(
24816 matches!(stream.next().await, Some(StreamChunk::Error { message }) if message.contains("Required context 'voice' not provided"))
24817 );
24818 assert!(stream.next().await.is_none());
24819 drop(stream);
24820 assert_eq!(calls.call_count(), 0);
24821 assert!(agent.memory.get_messages(None).await.unwrap().is_empty());
24822
24823 agent
24824 .set_context("voice.brief", serde_json::json!("ready"))
24825 .unwrap();
24826 assert_eq!(
24827 agent.chat("second").await.unwrap().content,
24828 "Voice response"
24829 );
24830 }
24831
24832 #[tokio::test]
24833 async fn test_required_context_blocks_event_streams_without_final() {
24834 use futures::StreamExt;
24835
24836 for actor_scoped in [false, true] {
24837 let mock = mock_with_response("Voice response");
24838 let calls = mock.clone();
24839 let agent = required_context_agent(mock, false);
24840 let mut events = if actor_scoped {
24841 agent
24842 .chat_stream_events_with_actor_context(
24843 "first",
24844 crate::TurnActorContext::new().with_origin_actor("caller"),
24845 )
24846 .await
24847 .unwrap()
24848 } else {
24849 agent.chat_stream_events("first").await.unwrap()
24850 };
24851 assert!(
24852 matches!(events.next().await, Some(AgentStreamEvent::Chunk(StreamChunk::Error { message })) if message.contains("Required context 'voice' not provided"))
24853 );
24854 assert!(events.next().await.is_none());
24855 drop(events);
24856 assert_eq!(calls.call_count(), 0);
24857 assert!(agent.memory.get_messages(None).await.unwrap().is_empty());
24858
24859 agent
24860 .set_context("voice.brief", serde_json::json!("ready"))
24861 .unwrap();
24862 assert_eq!(
24863 agent.chat("second").await.unwrap().content,
24864 "Voice response"
24865 );
24866 }
24867 }
24868}