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, ReflectionMode, 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 routing_reflection_mode(&self) -> ReflectionMode {
3834 self.get_effective_reflection_config().enabled
3835 }
3836
3837 fn get_skill_reasoning_config(&self, skill: &SkillDefinition) -> ReasoningConfig {
3838 skill
3839 .reasoning
3840 .clone()
3841 .unwrap_or_else(|| self.get_effective_reasoning_config())
3842 }
3843
3844 fn get_skill_reflection_config(&self, skill: &SkillDefinition) -> ReflectionConfig {
3845 skill
3846 .reflection
3847 .clone()
3848 .unwrap_or_else(|| self.get_effective_reflection_config())
3849 }
3850
3851 async fn build_disambiguation_context(&self) -> Result<DisambiguationContext> {
3853 let context_config = self
3854 .disambiguation_manager
3855 .as_ref()
3856 .map(|manager| manager.config().context.clone())
3857 .unwrap_or_default();
3858 let recent_messages = if context_config.recent_messages == 0 {
3859 Vec::new()
3860 } else {
3861 Self::readable_native_messages(
3862 self.memory
3863 .get_messages(Some(context_config.recent_messages))
3864 .await?,
3865 )?
3866 .iter()
3867 .map(|message| format!("{:?}: {}", message.role, message.content))
3868 .collect()
3869 };
3870
3871 let current_state = self.current_state().map(|s| s.to_string());
3872
3873 let state_prompt: Option<String> = self
3876 .state_machine
3877 .as_ref()
3878 .and_then(|sm| sm.current_definition())
3879 .and_then(|def| def.prompt.clone());
3880
3881 let available_tools = if context_config.include_available_tools {
3882 self.get_available_tool_ids().await?
3883 } else {
3884 Vec::new()
3885 };
3886
3887 let available_skills: Vec<String> = self.skills.iter().map(|s| s.id.clone()).collect();
3888
3889 let mut user_context = self.build_context_with_overlays();
3890 user_context.remove(DISAMBIGUATION_STATE_GENERATION_KEY);
3891 if let Some(state_generation) = self
3892 .state_machine
3893 .as_ref()
3894 .map(|state_machine| state_machine.generation())
3895 {
3896 user_context.insert(
3897 DISAMBIGUATION_STATE_GENERATION_KEY.to_string(),
3898 serde_json::json!(state_generation),
3899 );
3900 }
3901
3902 let available_intents: Vec<String> = if let Some(ref sm) = self.state_machine {
3904 sm.current_definition()
3905 .map(|def| {
3906 def.transitions
3907 .iter()
3908 .filter_map(|t| t.intent.clone())
3909 .collect()
3910 })
3911 .unwrap_or_default()
3912 } else {
3913 Vec::new()
3914 };
3915
3916 Ok(DisambiguationContext::from_agent_state(
3917 recent_messages,
3918 current_state,
3919 state_prompt,
3920 available_tools,
3921 available_skills,
3922 available_intents,
3923 user_context,
3924 ))
3925 }
3926
3927 fn get_available_skills(&self) -> Vec<&SkillDefinition> {
3928 if let Some(ref sm) = self.state_machine
3929 && let Some(state_def) = sm.current_definition()
3930 {
3931 let parent_def = sm.get_parent_definition();
3932 let effective_skills = state_def.get_effective_skills(parent_def.as_ref());
3933 if !effective_skills.is_empty() {
3934 return self
3935 .skills
3936 .iter()
3937 .filter(|s| effective_skills.contains(&&s.id))
3938 .collect();
3939 }
3940 }
3941 self.skills.iter().collect()
3942 }
3943
3944 async fn build_messages(&self) -> Result<Vec<ChatMessage>> {
3945 self.build_messages_internal(true, None, true).await
3946 }
3947
3948 async fn build_messages_for_draft(&self, user_message: &str) -> Result<Vec<ChatMessage>> {
3949 self.build_messages_internal(false, Some(user_message), true)
3950 .await
3951 }
3952
3953 async fn build_messages_internal(
3954 &self,
3955 fire_persona_hooks: bool,
3956 ephemeral_user_message: Option<&str>,
3957 include_tool_prompt: bool,
3958 ) -> Result<Vec<ChatMessage>> {
3959 let system_prompt = self
3960 .get_effective_system_prompt_with_persona_hooks(fire_persona_hooks, include_tool_prompt)
3961 .await?;
3962 let mut messages = vec![ChatMessage::system(&system_prompt)];
3963
3964 let context = self.memory.get_context().await?;
3965 let history = if let Some(ref budget) = self.memory_token_budget {
3966 context.to_llm_messages_with_allocation(&budget.allocation)
3967 } else {
3968 context.to_llm_messages()
3969 };
3970 messages.extend(history);
3971 if let Some(user_message) = ephemeral_user_message {
3972 messages.push(ChatMessage::user(user_message));
3973 }
3974
3975 let total_tokens = self.estimate_total_tokens(&messages);
3976
3977 if total_tokens > self.max_context_tokens {
3978 debug!(
3979 total = total_tokens,
3980 limit = self.max_context_tokens,
3981 "Context overflow"
3982 );
3983
3984 match &self.recovery_manager.config().llm.on_context_overflow {
3985 ContextOverflowAction::Error => {
3986 return Err(AgentError::LLM(format!(
3987 "Context overflow: {} tokens > {} limit",
3988 total_tokens, self.max_context_tokens
3989 )));
3990 }
3991 ContextOverflowAction::Truncate { keep_recent } => {
3992 self.truncate_context(&mut messages, *keep_recent)?;
3993 }
3994 ContextOverflowAction::Summarize {
3995 summarizer_llm,
3996 max_summary_tokens,
3997 custom_prompt,
3998 keep_recent,
3999 filter,
4000 } => {
4001 self.summarize_context(
4002 &mut messages,
4003 summarizer_llm.as_deref(),
4004 *max_summary_tokens,
4005 custom_prompt.as_deref(),
4006 *keep_recent,
4007 filter.as_ref(),
4008 )
4009 .await?;
4010 }
4011 }
4012 }
4013
4014 self.validate_active_native_history(&messages, true)?;
4015 Ok(messages)
4016 }
4017
4018 async fn main_tool_protocol(
4019 &self,
4020 llm: &dyn LLMProvider,
4021 ephemeral_new_turn: bool,
4022 ) -> Result<MainToolProtocol> {
4023 let mut choice = llm.configured_tool_choice();
4024 if matches!(choice.as_ref(), Some(ToolChoice::None)) {
4025 return Ok(MainToolProtocol {
4026 choice,
4027 tool_ids: Vec::new(),
4028 definitions: Vec::new(),
4029 });
4030 }
4031
4032 let mut tool_ids = self.get_available_tool_ids().await?;
4033 tool_ids.sort();
4034 tool_ids.dedup();
4035 if let Some(ToolChoice::Specific(expected)) = choice.as_ref() {
4036 let canonical = self.tools.canonical_id(expected).ok_or_else(|| {
4037 AgentError::Config(format!(
4038 "specific tool choice '{expected}' is not registered"
4039 ))
4040 })?;
4041 if canonical != *expected {
4042 return Err(AgentError::Config(format!(
4043 "specific tool choice must use canonical ID '{canonical}', not '{expected}'"
4044 )));
4045 }
4046 if !tool_ids.iter().any(|tool_id| tool_id == expected) {
4047 return Err(AgentError::Config(format!(
4048 "specific tool choice '{expected}' is outside the effective tool grant"
4049 )));
4050 }
4051 }
4052 if matches!(
4053 choice.as_ref(),
4054 Some(ToolChoice::Required | ToolChoice::Specific(_))
4055 ) && tool_ids.is_empty()
4056 {
4057 return Err(AgentError::Config(
4058 "required tool choice has no tool inside the effective grant".to_string(),
4059 ));
4060 }
4061 if !ephemeral_new_turn
4062 && let Some(configured_choice) = choice.as_ref()
4063 && matches!(
4064 configured_choice,
4065 ToolChoice::Required | ToolChoice::Specific(_)
4066 )
4067 && self
4068 .tool_choice_satisfied_in_current_turn(configured_choice, &tool_ids)
4069 .await?
4070 {
4071 choice = Some(ToolChoice::Auto);
4072 }
4073 if let Some(ToolChoice::Specific(expected)) = choice.as_ref() {
4074 tool_ids.retain(|tool_id| tool_id == expected);
4075 }
4076
4077 let definitions = tool_ids
4078 .iter()
4079 .map(|tool_id| {
4080 let tool = self.tools.get(tool_id).ok_or_else(|| {
4081 AgentError::Config(format!(
4082 "effective tool '{tool_id}' disappeared before provider exposure"
4083 ))
4084 })?;
4085 Ok(LLMToolDefinition {
4086 name: tool_id.clone(),
4087 description: tool.description().to_string(),
4088 input_schema: tool.input_schema(),
4089 })
4090 })
4091 .collect::<Result<Vec<_>>>()?;
4092
4093 Ok(MainToolProtocol {
4097 choice,
4098 tool_ids,
4099 definitions,
4100 })
4101 }
4102
4103 async fn tool_choice_satisfied_in_current_turn(
4104 &self,
4105 choice: &ToolChoice,
4106 effective_tool_ids: &[String],
4107 ) -> Result<bool> {
4108 let messages = self.memory.get_messages(None).await?;
4109 let mut saw_tool_result = false;
4110 for message in messages.iter().rev() {
4111 match message.role {
4112 ai_agents_core::Role::Tool | ai_agents_core::Role::Function => {
4113 saw_tool_result = true;
4114 }
4115 ai_agents_core::Role::Assistant if saw_tool_result => {
4116 let Some(calls) = self.parse_tool_calls(&message.content)? else {
4117 continue;
4118 };
4119 let calls_are_effective = !calls.is_empty()
4120 && calls.iter().all(|call| {
4121 self.tools
4122 .canonical_id(&call.name)
4123 .is_some_and(|canonical| effective_tool_ids.contains(&canonical))
4124 });
4125 return Ok(calls_are_effective
4126 && match choice {
4127 ToolChoice::Required => true,
4128 ToolChoice::Specific(expected) => calls.iter().all(|call| {
4129 self.tools.canonical_id(&call.name).as_deref()
4130 == Some(expected.as_str())
4131 }),
4132 _ => false,
4133 });
4134 }
4135 ai_agents_core::Role::User => return Ok(false),
4136 _ => {}
4137 }
4138 }
4139 Ok(false)
4140 }
4141
4142 fn provider_can_use_native_tools(
4143 &self,
4144 llm: &dyn LLMProvider,
4145 protocol: &MainToolProtocol,
4146 ) -> bool {
4147 let Some(choice) = protocol.choice.as_ref() else {
4148 return false;
4149 };
4150 if matches!(choice, ToolChoice::None) || protocol.definitions.is_empty() {
4151 return false;
4152 }
4153 llm.supports_tool_choice(choice)
4154 && protocol.definitions.iter().all(|definition| {
4155 !definition.name.is_empty()
4156 && definition.name.len() <= 64
4157 && definition
4158 .name
4159 .bytes()
4160 .all(|byte| byte.is_ascii_alphanumeric() || matches!(byte, b'_' | b'-'))
4161 })
4162 }
4163
4164 fn prompt_messages_for_tool_protocol(
4165 &self,
4166 messages: &[ChatMessage],
4167 protocol: &MainToolProtocol,
4168 corrective: bool,
4169 ) -> Vec<ChatMessage> {
4170 let mut messages = messages.to_vec();
4171 let Some(choice) = protocol.choice.as_ref() else {
4172 return messages;
4173 };
4174 if matches!(choice, ToolChoice::None) || protocol.tool_ids.is_empty() {
4175 return messages;
4176 }
4177
4178 let mut tool_prompt = self.tools.generate_scoped_prompt_with_mode(
4179 &protocol.tool_ids,
4180 None,
4181 self.parallel_tools.enabled,
4182 self.runtime_config.tool_schema_prompt_mode,
4183 );
4184 match choice {
4185 ToolChoice::Required => tool_prompt.push_str(
4186 "\n\nYou must call at least one listed tool before giving a final answer.",
4187 ),
4188 ToolChoice::Specific(tool_id) => tool_prompt.push_str(&format!(
4189 "\n\nYou must call the '{tool_id}' tool before giving a final answer."
4190 )),
4191 ToolChoice::Auto => {}
4192 ToolChoice::None => return messages,
4193 _ => return messages,
4194 }
4195 if let Some(system) = messages
4196 .iter_mut()
4197 .find(|message| message.role == ai_agents_core::Role::System)
4198 {
4199 system.content.push_str("\n\n");
4200 system.content.push_str(&tool_prompt);
4201 } else {
4202 messages.insert(0, ChatMessage::system(tool_prompt));
4203 }
4204 if corrective {
4205 let instruction = match choice {
4206 ToolChoice::Required => {
4207 "Your previous response did not call a required tool. Call at least one listed tool now and return only the JSON tool call."
4208 }
4209 ToolChoice::Specific(tool_id) => {
4210 messages.push(ChatMessage::user(format!(
4211 "Your previous response did not call the required '{tool_id}' tool. Call it now and return only the JSON tool call."
4212 )));
4213 return messages;
4214 }
4215 _ => return messages,
4216 };
4217 messages.push(ChatMessage::user(instruction));
4218 }
4219 messages
4220 }
4221
4222 async fn invoke_main_provider(
4223 &self,
4224 llm: Arc<dyn LLMProvider>,
4225 messages: &[ChatMessage],
4226 protocol: &MainToolProtocol,
4227 corrective: bool,
4228 ) -> std::result::Result<MainProviderResponse, LLMError> {
4229 let use_native = self.provider_can_use_native_tools(llm.as_ref(), protocol);
4230 let response = if use_native {
4231 let request = LLMToolRequest {
4232 tools: protocol.definitions.clone(),
4233 choice: protocol
4234 .choice
4235 .clone()
4236 .expect("native tool requests require an explicit choice"),
4237 };
4238 self.observe_purpose(
4239 ObservationPurpose::MainResponse,
4240 llm.complete_with_tools(messages, None, &request),
4241 )
4242 .await?
4243 } else {
4244 let prompt_messages =
4245 self.prompt_messages_for_tool_protocol(messages, protocol, corrective);
4246 self.observe_purpose(
4247 ObservationPurpose::MainResponse,
4248 llm.complete(&prompt_messages, None),
4249 )
4250 .await?
4251 };
4252 Ok(MainProviderResponse {
4253 response,
4254 used_native_tools: use_native,
4255 })
4256 }
4257
4258 async fn complete_main_attempt_with_recovery(
4259 &self,
4260 llm: Arc<dyn LLMProvider>,
4261 messages: &[ChatMessage],
4262 protocol: &MainToolProtocol,
4263 corrective: bool,
4264 ) -> Result<MainProviderResponse> {
4265 let primary_result = self
4267 .recovery_manager
4268 .with_llm_retry(
4269 "llm_call",
4270 None,
4271 || {
4272 let llm = Arc::clone(&llm);
4273 async move {
4274 self.invoke_main_provider(llm, messages, protocol, corrective)
4275 .await
4276 }
4277 },
4278 |error| llm.is_terminal_error(error),
4279 )
4280 .await;
4281
4282 match primary_result {
4283 Ok(response) => Ok(response),
4284 Err(ai_agents_recovery::RetryFailure::Terminal { error, .. }) => {
4285 Err(AgentError::LLM(error.to_string()))
4286 }
4287 Err(failure) => {
4288 let primary_error = AgentError::LLM(failure.into_error().to_string());
4289 match &self.recovery_manager.config().llm.on_failure {
4290 LLMFailureAction::FallbackLlm { fallback_llm } => {
4291 let fallback = self.llm_registry.get(fallback_llm).map_err(|error| {
4292 AgentError::Config(format!(
4293 "Fallback LLM '{fallback_llm}' not found: {error}"
4294 ))
4295 })?;
4296 self.invoke_main_provider(fallback, messages, protocol, corrective)
4297 .await
4298 .map_err(|error| AgentError::LLM(error.to_string()))
4299 }
4300 LLMFailureAction::FallbackResponse { message } => {
4301 if matches!(
4302 protocol.choice.as_ref(),
4303 Some(ToolChoice::Required | ToolChoice::Specific(_))
4304 ) {
4305 Err(AgentError::LLM(format!(
4306 "Required tool selection failed and cannot be satisfied by a static fallback response: {primary_error}"
4307 )))
4308 } else {
4309 Ok(MainProviderResponse {
4310 response: LLMResponse::new(message.clone(), FinishReason::Stop),
4311 used_native_tools: false,
4312 })
4313 }
4314 }
4315 LLMFailureAction::Error => Err(primary_error),
4316 }
4317 }
4318 }
4319 }
4320
4321 fn normalize_main_provider_response(
4322 &self,
4323 mut response: LLMResponse,
4324 protocol: &MainToolProtocol,
4325 ) -> Result<(LLMResponse, bool)> {
4326 let provider_state = response
4327 .take_provider_state()
4328 .map_err(|error| AgentError::LLM(error.to_string()))?;
4329 let native_calls = response
4330 .tool_calls()
4331 .map_err(|error| AgentError::LLM(error.to_string()))?;
4332 let calls = match native_calls {
4333 Some(calls) => {
4334 response.content = encode_native_tool_call_markers(&calls, provider_state.as_ref())
4335 .map_err(|error| AgentError::LLM(error.to_string()))?;
4336 Some(calls)
4337 }
4338 None if provider_state.is_some() => {
4339 return Err(AgentError::LLM(
4340 "Provider returned replay state without native tool calls".to_string(),
4341 ));
4342 }
4343 None if !matches!(protocol.choice.as_ref(), Some(ToolChoice::None)) => {
4344 self.parse_tool_calls(response.content.trim())?
4345 }
4346 None => None,
4347 };
4348
4349 if protocol.choice.is_some()
4350 && let Some(calls) = calls.as_ref()
4351 && calls.iter().any(|call| {
4352 self.tools
4353 .canonical_id(&call.name)
4354 .is_none_or(|canonical| !protocol.tool_ids.contains(&canonical))
4355 })
4356 {
4357 return Err(AgentError::LLM(
4358 "Provider returned a tool call outside the effective grant".to_string(),
4359 ));
4360 }
4361
4362 let compliant = match protocol.choice.as_ref() {
4363 Some(ToolChoice::Required) => calls.as_ref().is_some_and(|calls| !calls.is_empty()),
4364 Some(ToolChoice::Specific(expected)) => calls.as_ref().is_some_and(|calls| {
4365 !calls.is_empty()
4366 && calls.iter().all(|call| {
4367 self.tools.canonical_id(&call.name).as_deref() == Some(expected.as_str())
4368 })
4369 }),
4370 _ => true,
4371 };
4372 Ok((response, compliant))
4373 }
4374
4375 async fn complete_main_llm_with_recovery(
4376 &self,
4377 llm: Arc<dyn LLMProvider>,
4378 messages: &[ChatMessage],
4379 protocol: &MainToolProtocol,
4380 ) -> Result<LLMResponse> {
4381 let first = self
4382 .complete_main_attempt_with_recovery(Arc::clone(&llm), messages, protocol, false)
4383 .await?;
4384 let (response, compliant) =
4385 self.normalize_main_provider_response(first.response, protocol)?;
4386 if compliant {
4387 return Ok(response);
4388 }
4389 if first.used_native_tools {
4390 return Err(AgentError::LLM(
4391 "Provider returned no compliant native call for required tool choice".to_string(),
4392 ));
4393 }
4394
4395 let corrected = self
4396 .complete_main_attempt_with_recovery(llm, messages, protocol, true)
4397 .await?;
4398 let (response, compliant) =
4399 self.normalize_main_provider_response(corrected.response, protocol)?;
4400 if compliant {
4401 return Ok(response);
4402 }
4403 Err(AgentError::LLM(
4404 "Provider returned no compliant tool call after one corrective retry".to_string(),
4405 ))
4406 }
4407
4408 async fn open_main_stream_with_recovery(
4415 &self,
4416 llm: Arc<dyn LLMProvider>,
4417 messages: &[ChatMessage],
4418 protocol: &MainToolProtocol,
4419 ) -> Result<MainStreamSource> {
4420 debug_assert!(
4421 protocol.choice.is_none(),
4422 "streaming raw path must not run with explicit tool choice"
4423 );
4424 let primary = self
4425 .recovery_manager
4426 .with_llm_retry(
4427 "llm_stream_open",
4428 None,
4429 || {
4430 let llm = Arc::clone(&llm);
4431 async move {
4432 self.observe_purpose(
4433 ObservationPurpose::MainResponse,
4434 llm.complete_stream(messages, None),
4435 )
4436 .await
4437 }
4438 },
4439 |error| llm.is_terminal_error(error),
4440 )
4441 .await;
4442
4443 match primary {
4444 Ok(stream) => Ok(MainStreamSource::Stream(stream)),
4445 Err(ai_agents_recovery::RetryFailure::Terminal { error, .. }) => {
4446 Err(AgentError::LLM(error.to_string()))
4447 }
4448 Err(failure) => {
4449 let primary_error = AgentError::LLM(failure.into_error().to_string());
4450 match &self.recovery_manager.config().llm.on_failure {
4451 LLMFailureAction::FallbackLlm { fallback_llm } => {
4452 let fallback = self.llm_registry.get(fallback_llm).map_err(|error| {
4453 AgentError::Config(format!(
4454 "Fallback LLM '{fallback_llm}' not found: {error}"
4455 ))
4456 })?;
4457 if fallback.supports(LLMFeature::Streaming) {
4458 let stream = self
4459 .observe_purpose(
4460 ObservationPurpose::MainResponse,
4461 fallback.complete_stream(messages, None),
4462 )
4463 .await
4464 .map_err(|error| AgentError::LLM(error.to_string()))?;
4465 Ok(MainStreamSource::Stream(stream))
4466 } else {
4467 let response = self
4468 .observe_purpose(
4469 ObservationPurpose::MainResponse,
4470 fallback.complete(messages, None),
4471 )
4472 .await
4473 .map_err(|error| AgentError::LLM(error.to_string()))?;
4474 Ok(MainStreamSource::StaticResponse(response.content))
4475 }
4476 }
4477 LLMFailureAction::FallbackResponse { message } => {
4478 Ok(MainStreamSource::StaticResponse(message.clone()))
4479 }
4480 LLMFailureAction::Error => Err(primary_error),
4481 }
4482 }
4483 }
4484 }
4485
4486 fn main_stream_must_buffer(
4495 &self,
4496 reasoning_mode: &ReasoningMode,
4497 protocol: &MainToolProtocol,
4498 ) -> bool {
4499 protocol.choice.is_some()
4500 || self.get_effective_reflection_config().requires_evaluation()
4501 || matches!(reasoning_mode, ReasoningMode::CoT | ReasoningMode::React)
4502 }
4503
4504 fn is_native_tool_call_content(content: &str) -> Result<bool> {
4506 decode_native_tool_call_markers(content)
4507 .map(|batch| batch.is_some())
4508 .map_err(|error| AgentError::LLM(error.to_string()))
4509 }
4510
4511 fn tool_result_message(
4513 tool_call: &ToolCall,
4514 output: &str,
4515 native_tool_call: bool,
4516 ) -> Result<ChatMessage> {
4517 if !native_tool_call {
4518 return Ok(ChatMessage::function(&tool_call.name, output));
4519 }
4520 let output = serde_json::from_str::<serde_json::Value>(output)
4521 .unwrap_or_else(|_| serde_json::Value::String(output.to_string()));
4522 let content = encode_native_tool_result_marker(tool_call, output)
4523 .map_err(|error| AgentError::LLM(error.to_string()))?;
4524 Ok(ChatMessage::function(&tool_call.name, content))
4525 }
4526
4527 fn remember_active_native_exchange(&self, content: &str) -> Result<()> {
4529 let Some(batch) = decode_native_tool_call_markers(content)
4530 .map_err(|error| AgentError::LLM(error.to_string()))?
4531 else {
4532 return Ok(());
4533 };
4534 let Some(state) = batch.provider_state() else {
4535 return Ok(());
4536 };
4537 let expected = ActiveNativeExchange {
4538 exchange_id: state.exchange_id().to_string(),
4539 call_ids: batch.calls().iter().map(|call| call.id.clone()).collect(),
4540 };
4541 let mut active = self.active_native_exchanges.write();
4542 if let Some(existing) = active
4543 .iter()
4544 .find(|existing| existing.exchange_id == expected.exchange_id)
4545 {
4546 if existing.call_ids != expected.call_ids {
4547 return Err(AgentError::LLM(format!(
4548 "Active native exchange '{}' changed its call identities",
4549 expected.exchange_id
4550 )));
4551 }
4552 } else {
4553 active.push(expected);
4554 }
4555 Ok(())
4556 }
4557
4558 fn validate_active_native_history(
4560 &self,
4561 messages: &[ChatMessage],
4562 require_complete: bool,
4563 ) -> Result<()> {
4564 let expected = self.active_native_exchanges.read().clone();
4565 if expected.is_empty() {
4566 return Ok(());
4567 }
4568 let inspection =
4569 inspect_native_history(messages).map_err(|error| AgentError::LLM(error.to_string()))?;
4570 let expected_count = expected.len();
4571 for (index, expected) in expected.iter().enumerate() {
4572 let Some(exchange) = inspection
4573 .exchanges()
4574 .iter()
4575 .find(|exchange| exchange.state().exchange_id() == expected.exchange_id)
4576 else {
4577 return Err(AgentError::LLM(format!(
4578 "Active native exchange '{}' was removed before provider continuation",
4579 expected.exchange_id
4580 )));
4581 };
4582 let must_be_complete = require_complete || index + 1 < expected_count;
4583 if exchange.call_ids() != expected.call_ids
4584 || (must_be_complete && !exchange.is_complete())
4585 {
4586 return Err(AgentError::LLM(format!(
4587 "Active native exchange '{}' is incomplete before provider continuation",
4588 expected.exchange_id
4589 )));
4590 }
4591 }
4592 Ok(())
4593 }
4594
4595 async fn remember_committed_native_exchange(&self, content: &str) -> Result<()> {
4597 self.remember_active_native_exchange(content)?;
4598 if !self.active_native_exchanges.read().is_empty() {
4599 let messages = self.memory.get_messages(None).await?;
4600 self.validate_active_native_history(&messages, false)?;
4601 }
4602 Ok(())
4603 }
4604
4605 fn readable_native_messages(mut messages: Vec<ChatMessage>) -> Result<Vec<ChatMessage>> {
4607 for message in &mut messages {
4608 if matches!(
4609 message.role,
4610 ai_agents_core::Role::Assistant
4611 | ai_agents_core::Role::Tool
4612 | ai_agents_core::Role::Function
4613 ) {
4614 message.content = native_readable_projection(&message.content)
4615 .map_err(|error| AgentError::LLM(error.to_string()))?;
4616 }
4617 }
4618 Ok(messages)
4619 }
4620
4621 fn parse_main_tool_calls(
4623 &self,
4624 content: &str,
4625 protocol: &MainToolProtocol,
4626 ) -> Result<Option<Vec<ToolCall>>> {
4627 if matches!(protocol.choice.as_ref(), Some(ToolChoice::None)) {
4628 Ok(None)
4629 } else {
4630 self.parse_tool_calls(content)
4631 }
4632 }
4633
4634 fn parse_tool_calls(&self, content: &str) -> Result<Option<Vec<ToolCall>>> {
4636 if let Some(batch) = decode_native_tool_call_markers(content)
4637 .map_err(|error| AgentError::LLM(error.to_string()))?
4638 {
4639 return Ok(Some(batch.into_parts().0));
4640 }
4641 if let Ok(parsed) = serde_json::from_str::<serde_json::Value>(content) {
4643 if let Some(arr) = parsed.as_array() {
4645 let calls: Vec<ToolCall> = arr
4646 .iter()
4647 .filter_map(|v| self.extract_tool_call_from_value(v))
4648 .collect();
4649 if !calls.is_empty() {
4650 return Ok(Some(calls));
4651 }
4652 }
4653 if let Some(tool_call) = self.extract_tool_call_from_value(&parsed) {
4655 return Ok(Some(vec![tool_call]));
4656 }
4657 }
4658
4659 if let Some(json_str) = self.extract_json_from_content(content)
4661 && let Ok(parsed) = serde_json::from_str::<serde_json::Value>(&json_str)
4662 {
4663 if let Some(arr) = parsed.as_array() {
4665 let calls: Vec<ToolCall> = arr
4666 .iter()
4667 .filter_map(|v| self.extract_tool_call_from_value(v))
4668 .collect();
4669 if !calls.is_empty() {
4670 return Ok(Some(calls));
4671 }
4672 }
4673 if let Some(tool_call) = self.extract_tool_call_from_value(&parsed) {
4675 return Ok(Some(vec![tool_call]));
4676 }
4677 }
4678
4679 Ok(None)
4680 }
4681
4682 fn extract_tool_call_from_value(&self, parsed: &serde_json::Value) -> Option<ToolCall> {
4683 if let Some(tool_name) = parsed.get("tool").and_then(|v| v.as_str()) {
4684 let arguments = parsed
4685 .get("arguments")
4686 .cloned()
4687 .unwrap_or(serde_json::json!({}));
4688 return Some(ToolCall {
4689 id: parsed
4690 .get("id")
4691 .and_then(|value| value.as_str())
4692 .filter(|id| !id.is_empty())
4693 .map(str::to_string)
4694 .unwrap_or_else(|| uuid::Uuid::new_v4().to_string()),
4695 name: tool_name.to_string(),
4696 arguments,
4697 });
4698 }
4699 None
4700 }
4701
4702 fn extract_json_from_content(&self, content: &str) -> Option<String> {
4704 if let Some(result) = self.extract_json_array_from_content(content) {
4706 return Some(result);
4707 }
4708 self.extract_json_object_from_content(content)
4709 }
4710
4711 fn extract_json_array_from_content(&self, content: &str) -> Option<String> {
4713 let start = content.find('[')?;
4714 let content_from_start = &content[start..];
4715
4716 let mut depth = 0;
4717 let mut end = 0;
4718 for (i, ch) in content_from_start.char_indices() {
4719 match ch {
4720 '[' => depth += 1,
4721 ']' => {
4722 depth -= 1;
4723 if depth == 0 {
4724 end = i + 1;
4725 break;
4726 }
4727 }
4728 _ => {}
4729 }
4730 }
4731
4732 if end > 0 {
4733 let json_str = &content_from_start[..end];
4734 if json_str.contains("\"tool\"") {
4736 return Some(json_str.to_string());
4737 }
4738 }
4739
4740 None
4741 }
4742
4743 fn extract_json_object_from_content(&self, content: &str) -> Option<String> {
4745 let start = content.find('{')?;
4746 let content_from_start = &content[start..];
4747
4748 let mut depth = 0;
4750 let mut end = 0;
4751 for (i, ch) in content_from_start.char_indices() {
4752 match ch {
4753 '{' => depth += 1,
4754 '}' => {
4755 depth -= 1;
4756 if depth == 0 {
4757 end = i + 1;
4758 break;
4759 }
4760 }
4761 _ => {}
4762 }
4763 }
4764
4765 if end > 0 {
4766 let json_str = &content_from_start[..end];
4767 if json_str.contains("\"tool\"") {
4769 return Some(json_str.to_string());
4770 }
4771 }
4772
4773 None
4774 }
4775
4776 #[allow(clippy::too_many_arguments)]
4780 fn record_from_parts(
4781 &self,
4782 request: &ToolExecutionRequest,
4783 canonical_id: String,
4784 executed_arguments: Value,
4785 started_at: chrono::DateTime<chrono::Utc>,
4786 start: Instant,
4787 executed: bool,
4788 success: bool,
4789 output: String,
4790 metadata: HashMap<String, Value>,
4791 policy: ToolPolicyDecisionRecord,
4792 approval: Option<ToolApprovalRecord>,
4793 timed_out: bool,
4794 output_truncated: bool,
4795 ) -> ToolExecutionRecord {
4796 let versions = ToolDecisionVersions {
4797 policy: self.active_tool_security().policy_version(),
4798 registry: self.tools.version(),
4799 runtime_control: self.runtime_control.version.load(Ordering::SeqCst),
4800 state: self
4801 .state_machine
4802 .as_ref()
4803 .map(|state_machine| state_machine.generation()),
4804 };
4805 self.record_from_parts_at(
4806 request,
4807 canonical_id,
4808 executed_arguments,
4809 started_at,
4810 start,
4811 executed,
4812 success,
4813 output,
4814 metadata,
4815 policy,
4816 approval,
4817 timed_out,
4818 output_truncated,
4819 versions,
4820 )
4821 }
4822
4823 #[allow(clippy::too_many_arguments)]
4825 fn record_from_parts_at(
4826 &self,
4827 request: &ToolExecutionRequest,
4828 canonical_id: String,
4829 executed_arguments: Value,
4830 started_at: chrono::DateTime<chrono::Utc>,
4831 start: Instant,
4832 executed: bool,
4833 success: bool,
4834 output: String,
4835 metadata: HashMap<String, Value>,
4836 policy: ToolPolicyDecisionRecord,
4837 approval: Option<ToolApprovalRecord>,
4838 timed_out: bool,
4839 output_truncated: bool,
4840 versions: ToolDecisionVersions,
4841 ) -> ToolExecutionRecord {
4842 ToolExecutionRecord {
4843 call_id: request.call_id.clone(),
4844 requested_name: request.requested_name.clone(),
4845 canonical_id,
4846 source: request.source.clone(),
4847 arguments: request.arguments.clone(),
4848 executed_arguments,
4849 policy_version: versions.policy,
4850 registry_version: versions.registry,
4851 runtime_config_version: versions.runtime_control,
4852 executed,
4853 success,
4854 output,
4855 metadata,
4856 policy,
4857 approval,
4858 started_at,
4859 duration_ms: start.elapsed().as_millis() as u64,
4860 timed_out,
4861 cancelled: false,
4862 cancellation_reason: None,
4863 output_truncated,
4864 }
4865 }
4866
4867 async fn finish_tool_record(&self, record: &ToolExecutionRecord) {
4869 let result = ToolResult {
4870 success: record.success,
4871 output: record.model_output_string(),
4872 metadata: if record.metadata.is_empty() {
4873 None
4874 } else {
4875 Some(record.metadata.clone())
4876 },
4877 };
4878 self.hooks
4879 .on_tool_complete(&record.canonical_id, &result, record.duration_ms)
4880 .await;
4881 self.hooks.on_tool_execution_record(record).await;
4882 self.record_tool_call(&record.canonical_id, record.model_output_value());
4883 if !record.success {
4884 self.hooks
4885 .on_error(&AgentError::Tool(record.output.clone()))
4886 .await;
4887 }
4888 }
4889
4890 async fn finish_tool_record_after_resource_guards(
4892 &self,
4893 resource_guards: ToolResourceGuards,
4894 record: &ToolExecutionRecord,
4895 ) {
4896 drop(resource_guards);
4897 self.finish_tool_record(record).await;
4898 }
4899
4900 fn validated_tool_timeout(timeout_ms: u64) -> Result<ValidatedToolTimeout> {
4904 if timeout_ms > MAX_TOOL_TIMEOUT_MS {
4905 return Err(AgentError::Config(format!(
4906 "effective tool timeout_ms must be no greater than {MAX_TOOL_TIMEOUT_MS} milliseconds"
4907 )));
4908 }
4909 let timer = Duration::from_millis(timeout_ms);
4910 let deadline_delta = chrono::Duration::from_std(timer).map_err(|_| {
4911 AgentError::Config(format!(
4912 "effective tool timeout_ms cannot be represented as a UTC deadline: {timeout_ms}"
4913 ))
4914 })?;
4915 Ok(ValidatedToolTimeout {
4916 timer,
4917 deadline_delta,
4918 })
4919 }
4920
4921 fn effective_tool_limits(
4925 security_engine: &ToolSecurityEngine,
4926 canonical_id: &str,
4927 safety: &ToolSafetyMetadata,
4928 classification: &ToolCallClassification,
4929 recovery_timeout_ms: Option<u64>,
4930 ) -> Result<(ToolExecutionLimits, ValidatedToolTimeout)> {
4931 if let Some(timeout_ms) = classification.timeout_ms {
4932 Self::validated_tool_timeout(timeout_ms)?;
4933 }
4934 if let Some(timeout_ms) = recovery_timeout_ms {
4935 Self::validated_tool_timeout(timeout_ms)?;
4936 }
4937
4938 let mut limits = security_engine.effective_limits(canonical_id, safety, classification);
4939 if let Some(recovery_timeout_ms) = recovery_timeout_ms {
4940 limits.timeout_ms = Some(limits.timeout_ms.map_or(recovery_timeout_ms, |timeout_ms| {
4941 timeout_ms.min(recovery_timeout_ms)
4942 }));
4943 }
4944 let timeout_ms = limits
4945 .timeout_ms
4946 .unwrap_or_else(|| security_engine.get_tool_timeout(canonical_id));
4947 let timeout = Self::validated_tool_timeout(timeout_ms)?;
4948 Ok((limits, timeout))
4949 }
4950
4951 async fn execute_resolved_tool_once(
4953 &self,
4954 tool: Arc<dyn ai_agents_core::Tool>,
4955 args: Value,
4956 mut ctx: ToolExecutionContext,
4957 timeout: ValidatedToolTimeout,
4958 ) -> Result<(ToolResult, bool, bool, bool)> {
4959 if self.runtime_control.emergency_deny.load(Ordering::SeqCst) {
4960 return Ok((
4961 ToolResult::error("Tool execution cancelled by runtime control"),
4962 false,
4963 true,
4964 false,
4965 ));
4966 }
4967 ctx.deadline = Some(
4972 chrono::Utc::now()
4973 .checked_add_signed(timeout.deadline_delta)
4974 .ok_or_else(|| {
4975 AgentError::Config(
4976 "effective tool timeout_ms exceeds the current UTC deadline range"
4977 .to_string(),
4978 )
4979 })?,
4980 );
4981 let invoked = Arc::new(AtomicBool::new(false));
4985 let invoked_by_future = Arc::clone(&invoked);
4986 let actor_context = current_turn_actor_context();
4987 let future = async move {
4988 invoked_by_future.store(true, Ordering::SeqCst);
4989 if let Some(actor_context) = actor_context {
4990 scope_actor_context(actor_context, tool.execute(args, ctx)).await
4991 } else {
4992 tool.execute(args, ctx).await
4993 }
4994 };
4995 tokio::pin!(future);
4996 let timer = tokio::time::sleep(timeout.timer);
4997 tokio::pin!(timer);
4998 let mut cancel_tick = tokio::time::interval(std::time::Duration::from_millis(50));
4999
5000 loop {
5001 tokio::select! {
5002 result = &mut future => return Ok((result, false, false, true)),
5003 _ = &mut timer => {
5004 return Ok((
5005 ToolResult::error("Tool execution timed out"),
5006 true,
5007 false,
5008 invoked.load(Ordering::SeqCst),
5009 ));
5010 }
5011 _ = cancel_tick.tick() => {
5012 if self.runtime_control.emergency_deny.load(Ordering::SeqCst) {
5013 return Ok((
5014 ToolResult::error("Tool execution cancelled by runtime control"),
5015 false,
5016 true,
5017 invoked.load(Ordering::SeqCst),
5018 ));
5019 }
5020 }
5021 }
5022 }
5023 }
5024
5025 fn truncate_tool_output(output: String, max_chars: Option<usize>) -> (String, bool) {
5027 let Some(max_chars) = max_chars else {
5028 return (output, false);
5029 };
5030 let mut chars = output.chars();
5031 let truncated: String = chars.by_ref().take(max_chars).collect();
5032 if chars.next().is_some() {
5033 (truncated, true)
5034 } else {
5035 (output, false)
5036 }
5037 }
5038
5039 async fn acquire_tool_resource_locks(&self, keys: &[String]) -> Option<ToolResourceGuards> {
5041 let locks = {
5042 let mut table = self.resource_locks.write();
5043 table.retain(|_, lock| lock.strong_count() > 0);
5044 keys.iter()
5045 .map(|key| {
5046 if let Some(lock) = table.get(key).and_then(Weak::upgrade) {
5047 lock
5048 } else {
5049 let lock = Arc::new(tokio::sync::Mutex::new(()));
5050 table.insert(key.clone(), Arc::downgrade(&lock));
5051 lock
5052 }
5053 })
5054 .collect::<Vec<_>>()
5055 };
5056 let mut resource_guards = ToolResourceGuards {
5057 guards: Vec::with_capacity(locks.len()),
5058 locks: Arc::clone(&self.resource_locks),
5059 };
5060 let mut locks = locks.into_iter();
5061 while let Some(lock) = locks.next() {
5062 let mut lock = Box::pin(lock.lock_owned());
5063 loop {
5064 tokio::select! {
5065 guard = &mut lock => {
5066 resource_guards.guards.push(guard);
5067 break;
5068 }
5069 _ = tokio::time::sleep(std::time::Duration::from_millis(10)) => {
5070 if self.runtime_control.emergency_deny.load(Ordering::SeqCst) {
5071 drop(lock);
5072 drop(locks);
5073 drop(resource_guards);
5074 return None;
5075 }
5076 }
5077 }
5078 }
5079 }
5080 Some(resource_guards)
5081 }
5082
5083 async fn run_tool_with_retries(
5087 &self,
5088 canonical_id: &str,
5089 tool: Arc<dyn ai_agents_core::Tool>,
5090 args: Value,
5091 ctx: ToolExecutionContext,
5092 timeout: ValidatedToolTimeout,
5093 max_retries: u32,
5094 ) -> Result<(ToolResult, bool, bool, bool)> {
5095 let max_retries = if ctx.classification.safely_retryable {
5096 max_retries
5097 } else {
5098 0
5099 };
5100 let mut attempts = 0;
5101 let mut invoked = false;
5102 loop {
5103 let (result, timed_out, cancelled, attempt_invoked) = self
5104 .execute_resolved_tool_once(tool.clone(), args.clone(), ctx.clone(), timeout)
5105 .await?;
5106 invoked |= attempt_invoked;
5107 if result.success || timed_out || cancelled || attempts >= max_retries {
5108 return Ok((result, timed_out, cancelled, invoked));
5109 }
5110 attempts += 1;
5111 warn!(tool = %canonical_id, attempt = attempts, error = %result.output, "Retrying failed tool call");
5112 }
5113 }
5114
5115 fn host_tool_unavailability(&self, canonical_id: &str) -> Option<(&'static str, &'static str)> {
5117 match canonical_id {
5118 "command" if !self.tools.command_runner_available() => Some((
5119 "Command runner is unavailable",
5120 "command runner is unavailable",
5121 )),
5122 "diagnostics" if !self.tools.diagnostics_available() => Some((
5123 "Diagnostics provider is unavailable",
5124 "diagnostics provider is unavailable",
5125 )),
5126 "web_search" if !self.tools.web_search_available() => Some((
5127 "Web search provider is unavailable",
5128 "web search provider is unavailable",
5129 )),
5130 _ => None,
5131 }
5132 }
5133
5134 fn execute_tool_record(
5136 &self,
5137 request: ToolExecutionRequest,
5138 ) -> Pin<Box<dyn Future<Output = Result<ToolExecutionRecord>> + Send + '_>> {
5139 Box::pin(self.execute_tool_record_inner(request, ToolFallbackState::default()))
5140 }
5141
5142 async fn execute_tool_record_inner(
5146 &self,
5147 request: ToolExecutionRequest,
5148 fallback_state: ToolFallbackState,
5149 ) -> Result<ToolExecutionRecord> {
5150 let started_at = chrono::Utc::now();
5151 let start = Instant::now();
5152 info!(tool = %request.requested_name, args = %request.arguments, "Executing tool");
5153
5154 if self.runtime_control.emergency_deny.load(Ordering::SeqCst) {
5155 let record = self.record_from_parts(
5156 &request,
5157 request.requested_name.clone(),
5158 request.arguments.clone(),
5159 started_at,
5160 start,
5161 false,
5162 false,
5163 "Tool execution is disabled by runtime control".to_string(),
5164 HashMap::new(),
5165 ToolPolicyDecisionRecord::deny("runtime emergency deny is enabled"),
5166 None,
5167 false,
5168 false,
5169 );
5170 self.finish_tool_record(&record).await;
5171 return Ok(record);
5172 }
5173
5174 let Some(resolved) = self.tools.resolve(&request.requested_name) else {
5175 let record = self.record_from_parts(
5176 &request,
5177 request.requested_name.clone(),
5178 request.arguments.clone(),
5179 started_at,
5180 start,
5181 false,
5182 false,
5183 format!("Tool '{}' is unavailable", request.requested_name),
5184 HashMap::new(),
5185 ToolPolicyDecisionRecord::unavailable(format!(
5186 "Tool '{}' is not registered",
5187 request.requested_name
5188 )),
5189 None,
5190 false,
5191 false,
5192 );
5193 self.finish_tool_record(&record).await;
5194 return Ok(record);
5195 };
5196
5197 let canonical_id = resolved.identity.canonical_id.clone();
5198
5199 let initial_scope_snapshot = self.get_available_tool_ids_snapshot().await?;
5200 if !initial_scope_snapshot
5201 .tool_ids
5202 .iter()
5203 .any(|id| id == &canonical_id)
5204 {
5205 let record = self.record_from_parts(
5206 &request,
5207 canonical_id.clone(),
5208 request.arguments.clone(),
5209 started_at,
5210 start,
5211 false,
5212 false,
5213 format!(
5214 "Tool '{}' is not available in the current scope",
5215 canonical_id
5216 ),
5217 HashMap::new(),
5218 ToolPolicyDecisionRecord::deny(format!(
5219 "Tool '{}' is not granted by the current top-level and state tool scope",
5220 canonical_id
5221 )),
5222 None,
5223 false,
5224 false,
5225 );
5226 self.finish_tool_record(&record).await;
5227 return Ok(record);
5228 }
5229
5230 let approval_control_snapshot = self.runtime_safety_snapshot();
5231 let security_engine = approval_control_snapshot.tool_security.clone();
5232 if let Some(reason) = fallback_state.rejection_reason(&canonical_id) {
5233 let mut metadata = HashMap::new();
5234 metadata.insert(
5235 "fallback_chain".to_string(),
5236 serde_json::to_value(&fallback_state.visited_canonical_ids).unwrap_or(Value::Null),
5237 );
5238 let record = self.record_from_parts(
5239 &request,
5240 canonical_id,
5241 request.arguments.clone(),
5242 started_at,
5243 start,
5244 false,
5245 false,
5246 format!("Denied: {reason}"),
5247 metadata,
5248 ToolPolicyDecisionRecord::deny(reason),
5249 None,
5250 false,
5251 false,
5252 );
5253 self.finish_tool_record(&record).await;
5254 return Ok(record);
5255 }
5256 let admitted_canonical_id = canonical_id.clone();
5257 let fallback_state = fallback_state.with_current(canonical_id.clone());
5258 let bindings = resolved.tool.policy_bindings();
5259 let mut executed_arguments = security_engine.prepare_tool_arguments_with_bindings(
5260 &canonical_id,
5261 &request.arguments,
5262 &bindings,
5263 );
5264 let mut metadata = HashMap::new();
5265 let safety = resolved.tool.safety_metadata();
5266 let classification = resolved.tool.classify_call(&executed_arguments);
5267 let initial_recovery_timeout_ms = self.recovery_manager.get_tool_timeout(&canonical_id);
5268 let (limits, _) = Self::effective_tool_limits(
5269 &security_engine,
5270 &canonical_id,
5271 &safety,
5272 &classification,
5273 initial_recovery_timeout_ms,
5274 )?;
5275 self.hooks
5276 .on_tool_start(&canonical_id, &executed_arguments)
5277 .await;
5278 metadata.insert(
5279 "classification".to_string(),
5280 serde_json::to_value(&classification).unwrap_or(Value::Null),
5281 );
5282 metadata.insert(
5283 "effective_limits".to_string(),
5284 serde_json::to_value(&limits).unwrap_or(Value::Null),
5285 );
5286 let policy_snapshot = security_engine.policy_snapshot(&canonical_id);
5287 if !policy_snapshot.is_null() {
5288 metadata.insert("policy_snapshot".to_string(), policy_snapshot.clone());
5289 }
5290
5291 let mut approval_record = Some(ToolApprovalRecord {
5292 status: ToolApprovalStatus::NotRequired,
5293 reason: None,
5294 modified_arguments: None,
5295 });
5296
5297 let mut security_result = security_engine
5298 .validate_tool_execution_with_bindings(&canonical_id, &executed_arguments, &bindings)
5299 .await?;
5300 if (security_result.is_allowed()
5305 || matches!(
5306 &security_result,
5307 SecurityCheckResult::RequireConfirmation { .. }
5308 ))
5309 && let Some((output, reason)) = self.host_tool_unavailability(&canonical_id)
5310 {
5311 let record = self.record_from_parts(
5312 &request,
5313 canonical_id,
5314 executed_arguments,
5315 started_at,
5316 start,
5317 false,
5318 false,
5319 output.to_string(),
5320 metadata,
5321 ToolPolicyDecisionRecord::unavailable(reason),
5322 Some(ToolApprovalRecord {
5323 status: ToolApprovalStatus::Unavailable,
5324 reason: Some(reason.to_string()),
5325 modified_arguments: None,
5326 }),
5327 false,
5328 false,
5329 );
5330 self.finish_tool_record(&record).await;
5331 return Ok(record);
5332 }
5333 match &security_result {
5334 SecurityCheckResult::Allow => {}
5335 SecurityCheckResult::Warn { message } => {
5336 warn!(tool = %canonical_id, message = %message, "Tool security warning");
5337 }
5338 SecurityCheckResult::Block { reason } => {
5339 let record = self.record_from_parts(
5340 &request,
5341 canonical_id,
5342 executed_arguments,
5343 started_at,
5344 start,
5345 false,
5346 false,
5347 format!("Denied: {}", reason),
5348 metadata,
5349 ToolPolicyDecisionRecord::deny(reason.clone()),
5350 approval_record,
5351 false,
5352 false,
5353 );
5354 self.finish_tool_record(&record).await;
5355 return Ok(record);
5356 }
5357 SecurityCheckResult::Unavailable { reason } => {
5358 let record = self.record_from_parts(
5359 &request,
5360 canonical_id,
5361 executed_arguments,
5362 started_at,
5363 start,
5364 false,
5365 false,
5366 format!("Unavailable: {}", reason),
5367 metadata,
5368 ToolPolicyDecisionRecord::unavailable(reason.clone()),
5369 approval_record,
5370 false,
5371 false,
5372 );
5373 self.finish_tool_record(&record).await;
5374 return Ok(record);
5375 }
5376 SecurityCheckResult::RequireConfirmation { message } => {
5377 if self.hitl_engine.is_none() {
5378 approval_record = Some(ToolApprovalRecord {
5379 status: ToolApprovalStatus::Unavailable,
5380 reason: Some("No HITL engine configured".to_string()),
5381 modified_arguments: None,
5382 });
5383 let record = self.record_from_parts(
5384 &request,
5385 canonical_id,
5386 executed_arguments,
5387 started_at,
5388 start,
5389 false,
5390 false,
5391 format!("Approval unavailable: {}", message),
5392 metadata,
5393 ToolPolicyDecisionRecord::approval(message.clone()),
5394 approval_record,
5395 false,
5396 false,
5397 );
5398 self.finish_tool_record(&record).await;
5399 return Ok(record);
5400 }
5401
5402 let check_result = HITLCheckResult::required(
5403 ApprovalTrigger::tool(&canonical_id, executed_arguments.clone()),
5404 HashMap::new(),
5405 message.clone(),
5406 None,
5407 );
5408 match self.request_hitl_approval(check_result).await? {
5409 ApprovalResult::Approved => {
5410 merge_approved_record(&mut approval_record);
5411 }
5412 ApprovalResult::Modified { changes } => {
5413 if let Some(obj) = executed_arguments.as_object_mut() {
5414 for (key, value) in changes {
5415 obj.insert(key, value);
5416 }
5417 }
5418 security_result = security_engine
5419 .validate_tool_execution_with_bindings(
5420 &canonical_id,
5421 &executed_arguments,
5422 &bindings,
5423 )
5424 .await?;
5425 if !matches!(
5426 security_result,
5427 SecurityCheckResult::Allow
5428 | SecurityCheckResult::Warn { .. }
5429 | SecurityCheckResult::RequireConfirmation { .. }
5430 ) {
5431 let reason = security_result
5432 .reason()
5433 .unwrap_or("modified arguments failed policy")
5434 .to_string();
5435 let record = self.record_from_parts(
5436 &request,
5437 canonical_id,
5438 executed_arguments.clone(),
5439 started_at,
5440 start,
5441 false,
5442 false,
5443 reason.clone(),
5444 metadata,
5445 ToolPolicyDecisionRecord::deny(reason),
5446 Some(ToolApprovalRecord {
5447 status: ToolApprovalStatus::Modified,
5448 reason: None,
5449 modified_arguments: Some(executed_arguments),
5450 }),
5451 false,
5452 false,
5453 );
5454 self.finish_tool_record(&record).await;
5455 return Ok(record);
5456 }
5457 approval_record = Some(ToolApprovalRecord {
5458 status: ToolApprovalStatus::Modified,
5459 reason: None,
5460 modified_arguments: Some(executed_arguments.clone()),
5461 });
5462 }
5463 ApprovalResult::Rejected { reason } => {
5464 let reason = reason.unwrap_or_else(|| "rejected".to_string());
5465 approval_record = Some(ToolApprovalRecord {
5466 status: ToolApprovalStatus::Rejected,
5467 reason: Some(reason.clone()),
5468 modified_arguments: None,
5469 });
5470 let record = self.record_from_parts(
5471 &request,
5472 canonical_id,
5473 executed_arguments,
5474 started_at,
5475 start,
5476 false,
5477 false,
5478 format!("Approval rejected: {}", reason),
5479 metadata,
5480 ToolPolicyDecisionRecord::approval(reason),
5481 approval_record,
5482 false,
5483 false,
5484 );
5485 self.finish_tool_record(&record).await;
5486 return Ok(record);
5487 }
5488 ApprovalResult::Timeout => {
5489 approval_record = Some(ToolApprovalRecord {
5490 status: ToolApprovalStatus::Timeout,
5491 reason: Some("approval timeout".to_string()),
5492 modified_arguments: None,
5493 });
5494 let record = self.record_from_parts(
5495 &request,
5496 canonical_id,
5497 executed_arguments,
5498 started_at,
5499 start,
5500 false,
5501 false,
5502 "Approval timed out".to_string(),
5503 metadata,
5504 ToolPolicyDecisionRecord::approval("approval timeout"),
5505 approval_record,
5506 false,
5507 false,
5508 );
5509 self.finish_tool_record(&record).await;
5510 return Ok(record);
5511 }
5512 }
5513 }
5514 }
5515
5516 if approval_record
5517 .as_ref()
5518 .is_some_and(|record| matches!(record.status, ToolApprovalStatus::NotRequired))
5519 && let Some(message) =
5520 security_engine.classification_approval_message(&canonical_id, &classification)
5521 {
5522 if self.hitl_engine.is_none() {
5523 approval_record = Some(ToolApprovalRecord {
5524 status: ToolApprovalStatus::Unavailable,
5525 reason: Some("No HITL engine configured".to_string()),
5526 modified_arguments: None,
5527 });
5528 let record = self.record_from_parts(
5529 &request,
5530 canonical_id,
5531 executed_arguments,
5532 started_at,
5533 start,
5534 false,
5535 false,
5536 format!("Approval unavailable: {}", message),
5537 metadata,
5538 ToolPolicyDecisionRecord::approval(message),
5539 approval_record,
5540 false,
5541 false,
5542 );
5543 self.finish_tool_record(&record).await;
5544 return Ok(record);
5545 }
5546 let check_result = HITLCheckResult::required(
5547 ApprovalTrigger::tool(&canonical_id, executed_arguments.clone()),
5548 HashMap::new(),
5549 message.clone(),
5550 None,
5551 );
5552 match self.request_hitl_approval(check_result).await? {
5553 ApprovalResult::Approved => {
5554 merge_approved_record(&mut approval_record);
5555 }
5556 ApprovalResult::Modified { changes } => {
5557 if let Some(obj) = executed_arguments.as_object_mut() {
5558 for (key, value) in changes {
5559 obj.insert(key, value);
5560 }
5561 }
5562 let modified_security = security_engine
5563 .validate_tool_execution_with_bindings(
5564 &canonical_id,
5565 &executed_arguments,
5566 &bindings,
5567 )
5568 .await?;
5569 if !matches!(
5570 modified_security,
5571 SecurityCheckResult::Allow | SecurityCheckResult::Warn { .. }
5572 ) {
5573 let reason = modified_security
5574 .reason()
5575 .unwrap_or("modified arguments failed policy")
5576 .to_string();
5577 let record = self.record_from_parts(
5578 &request,
5579 canonical_id,
5580 executed_arguments.clone(),
5581 started_at,
5582 start,
5583 false,
5584 false,
5585 reason.clone(),
5586 metadata,
5587 ToolPolicyDecisionRecord::deny(reason),
5588 Some(ToolApprovalRecord {
5589 status: ToolApprovalStatus::Modified,
5590 reason: None,
5591 modified_arguments: Some(executed_arguments),
5592 }),
5593 false,
5594 false,
5595 );
5596 self.finish_tool_record(&record).await;
5597 return Ok(record);
5598 }
5599 approval_record = Some(ToolApprovalRecord {
5600 status: ToolApprovalStatus::Modified,
5601 reason: None,
5602 modified_arguments: Some(executed_arguments.clone()),
5603 });
5604 }
5605 ApprovalResult::Rejected { reason } => {
5606 let reason = reason.unwrap_or_else(|| "rejected".to_string());
5607 let record = self.record_from_parts(
5608 &request,
5609 canonical_id,
5610 executed_arguments,
5611 started_at,
5612 start,
5613 false,
5614 false,
5615 format!("Approval rejected: {}", reason),
5616 metadata,
5617 ToolPolicyDecisionRecord::approval(reason.clone()),
5618 Some(ToolApprovalRecord {
5619 status: ToolApprovalStatus::Rejected,
5620 reason: Some(reason),
5621 modified_arguments: None,
5622 }),
5623 false,
5624 false,
5625 );
5626 self.finish_tool_record(&record).await;
5627 return Ok(record);
5628 }
5629 ApprovalResult::Timeout => {
5630 let record = self.record_from_parts(
5631 &request,
5632 canonical_id,
5633 executed_arguments,
5634 started_at,
5635 start,
5636 false,
5637 false,
5638 "Approval timed out".to_string(),
5639 metadata,
5640 ToolPolicyDecisionRecord::approval("approval timeout"),
5641 Some(ToolApprovalRecord {
5642 status: ToolApprovalStatus::Timeout,
5643 reason: Some("approval timeout".to_string()),
5644 modified_arguments: None,
5645 }),
5646 false,
5647 false,
5648 );
5649 self.finish_tool_record(&record).await;
5650 return Ok(record);
5651 }
5652 }
5653 }
5654
5655 let hitl_lang_ctx = self.build_hitl_language_context();
5656 if let Some(ref hitl_engine) = self.hitl_engine {
5657 let check_result = self
5658 .observe_purpose(
5659 ObservationPurpose::HitlLocalization,
5660 hitl_engine.check_tool_with_localization(
5661 &canonical_id,
5662 &executed_arguments,
5663 &hitl_lang_ctx,
5664 self.approval_handler.as_ref(),
5665 Some(&self.llm_registry),
5666 ),
5667 )
5668 .await?;
5669 if check_result.is_required() {
5670 match self.request_hitl_approval(check_result).await? {
5671 ApprovalResult::Approved => {
5672 merge_approved_record(&mut approval_record);
5673 }
5674 ApprovalResult::Modified { changes } => {
5675 if let Some(obj) = executed_arguments.as_object_mut() {
5676 for (key, value) in changes {
5677 obj.insert(key, value);
5678 }
5679 }
5680 let modified_security = security_engine
5681 .validate_tool_execution_with_bindings(
5682 &canonical_id,
5683 &executed_arguments,
5684 &bindings,
5685 )
5686 .await?;
5687 if !matches!(
5688 modified_security,
5689 SecurityCheckResult::Allow | SecurityCheckResult::Warn { .. }
5690 ) {
5691 let reason = modified_security
5692 .reason()
5693 .unwrap_or("modified arguments failed policy")
5694 .to_string();
5695 let record = self.record_from_parts(
5696 &request,
5697 canonical_id,
5698 executed_arguments.clone(),
5699 started_at,
5700 start,
5701 false,
5702 false,
5703 reason.clone(),
5704 metadata,
5705 ToolPolicyDecisionRecord::deny(reason),
5706 Some(ToolApprovalRecord {
5707 status: ToolApprovalStatus::Modified,
5708 reason: None,
5709 modified_arguments: Some(executed_arguments),
5710 }),
5711 false,
5712 false,
5713 );
5714 self.finish_tool_record(&record).await;
5715 return Ok(record);
5716 }
5717 approval_record = Some(ToolApprovalRecord {
5718 status: ToolApprovalStatus::Modified,
5719 reason: None,
5720 modified_arguments: Some(executed_arguments.clone()),
5721 });
5722 }
5723 ApprovalResult::Rejected { reason } => {
5724 let reason = reason.unwrap_or_else(|| "rejected".to_string());
5725 let record = self.record_from_parts(
5726 &request,
5727 canonical_id,
5728 executed_arguments,
5729 started_at,
5730 start,
5731 false,
5732 false,
5733 format!("Approval rejected: {}", reason),
5734 metadata,
5735 ToolPolicyDecisionRecord::approval(reason.clone()),
5736 Some(ToolApprovalRecord {
5737 status: ToolApprovalStatus::Rejected,
5738 reason: Some(reason),
5739 modified_arguments: None,
5740 }),
5741 false,
5742 false,
5743 );
5744 self.finish_tool_record(&record).await;
5745 return Ok(record);
5746 }
5747 ApprovalResult::Timeout => {
5748 let record = self.record_from_parts(
5749 &request,
5750 canonical_id,
5751 executed_arguments,
5752 started_at,
5753 start,
5754 false,
5755 false,
5756 "Approval timed out".to_string(),
5757 metadata,
5758 ToolPolicyDecisionRecord::approval("approval timeout"),
5759 Some(ToolApprovalRecord {
5760 status: ToolApprovalStatus::Timeout,
5761 reason: Some("approval timeout".to_string()),
5762 modified_arguments: None,
5763 }),
5764 false,
5765 false,
5766 );
5767 self.finish_tool_record(&record).await;
5768 return Ok(record);
5769 }
5770 }
5771 }
5772
5773 let condition_check = self
5774 .observe_purpose(
5775 ObservationPurpose::HitlLocalization,
5776 hitl_engine.check_conditions_with_localization(
5777 &executed_arguments,
5778 &hitl_lang_ctx,
5779 self.approval_handler.as_ref(),
5780 Some(&self.llm_registry),
5781 ),
5782 )
5783 .await?;
5784 if condition_check.is_required() {
5785 match self.request_hitl_approval(condition_check).await? {
5786 ApprovalResult::Approved => {
5787 merge_approved_record(&mut approval_record);
5788 }
5789 ApprovalResult::Modified { changes } => {
5790 if let Some(obj) = executed_arguments.as_object_mut() {
5791 for (key, value) in changes {
5792 obj.insert(key, value);
5793 }
5794 }
5795 let modified_security = security_engine
5796 .validate_tool_execution_with_bindings(
5797 &canonical_id,
5798 &executed_arguments,
5799 &bindings,
5800 )
5801 .await?;
5802 if !matches!(
5803 modified_security,
5804 SecurityCheckResult::Allow | SecurityCheckResult::Warn { .. }
5805 ) {
5806 let reason = modified_security
5807 .reason()
5808 .unwrap_or("modified arguments failed policy")
5809 .to_string();
5810 let record = self.record_from_parts(
5811 &request,
5812 canonical_id,
5813 executed_arguments,
5814 started_at,
5815 start,
5816 false,
5817 false,
5818 reason.clone(),
5819 metadata,
5820 ToolPolicyDecisionRecord::deny(reason),
5821 approval_record,
5822 false,
5823 false,
5824 );
5825 self.finish_tool_record(&record).await;
5826 return Ok(record);
5827 }
5828 approval_record = Some(ToolApprovalRecord {
5829 status: ToolApprovalStatus::Modified,
5830 reason: None,
5831 modified_arguments: Some(executed_arguments.clone()),
5832 });
5833 }
5834 ApprovalResult::Rejected { reason } => {
5835 let reason = reason.unwrap_or_else(|| "rejected".to_string());
5836 let record = self.record_from_parts(
5837 &request,
5838 canonical_id,
5839 executed_arguments,
5840 started_at,
5841 start,
5842 false,
5843 false,
5844 format!("Approval rejected: {}", reason),
5845 metadata,
5846 ToolPolicyDecisionRecord::approval(reason.clone()),
5847 Some(ToolApprovalRecord {
5848 status: ToolApprovalStatus::Rejected,
5849 reason: Some(reason),
5850 modified_arguments: None,
5851 }),
5852 false,
5853 false,
5854 );
5855 self.finish_tool_record(&record).await;
5856 return Ok(record);
5857 }
5858 ApprovalResult::Timeout => {
5859 let record = self.record_from_parts(
5860 &request,
5861 canonical_id,
5862 executed_arguments,
5863 started_at,
5864 start,
5865 false,
5866 false,
5867 "Approval timed out".to_string(),
5868 metadata,
5869 ToolPolicyDecisionRecord::approval("approval timeout"),
5870 Some(ToolApprovalRecord {
5871 status: ToolApprovalStatus::Timeout,
5872 reason: Some("approval timeout".to_string()),
5873 modified_arguments: None,
5874 }),
5875 false,
5876 false,
5877 );
5878 self.finish_tool_record(&record).await;
5879 return Ok(record);
5880 }
5881 }
5882 }
5883 }
5884
5885 executed_arguments = security_engine.prepare_tool_arguments_with_bindings(
5890 &canonical_id,
5891 &executed_arguments,
5892 &bindings,
5893 );
5894 if let Some(record) = approval_record.as_mut()
5895 && matches!(record.status, ToolApprovalStatus::Modified)
5896 {
5897 record.modified_arguments = Some(executed_arguments.clone());
5898 }
5899 let binding_security_result = security_engine
5900 .validate_tool_execution_with_bindings(&canonical_id, &executed_arguments, &bindings)
5901 .await?;
5902 let approval_confirmation_required = matches!(
5903 binding_security_result,
5904 SecurityCheckResult::RequireConfirmation { .. }
5905 ) || security_engine
5906 .classification_approval_message(
5907 &canonical_id,
5908 &resolved.tool.classify_call(&executed_arguments),
5909 )
5910 .is_some();
5911 let approval_binding = approval_record.as_ref().and_then(|record| {
5912 matches!(
5913 record.status,
5914 ToolApprovalStatus::Approved | ToolApprovalStatus::Modified
5915 )
5916 .then(|| ToolApprovalBinding {
5917 canonical_id: canonical_id.clone(),
5918 arguments: executed_arguments.clone(),
5919 confirmation_required: approval_confirmation_required,
5920 policy_version: security_engine.policy_version(),
5921 runtime_control_version: approval_control_snapshot.version,
5922 state_generation: initial_scope_snapshot.state_generation,
5923 reviewed_tool: Arc::clone(&resolved.tool),
5924 })
5925 });
5926
5927 let control_snapshot = self.runtime_safety_snapshot();
5932 let resolved = self.tools.resolve(&request.requested_name);
5933 let registry_version = self.tools.version();
5934 let mut versions = ToolDecisionVersions {
5935 policy: control_snapshot.tool_security.policy_version(),
5936 registry: registry_version,
5937 runtime_control: control_snapshot.version,
5938 state: None,
5939 };
5940 metadata.insert(
5941 "runtime_scope_snapshot".to_string(),
5942 serde_json::to_value(&control_snapshot.tool_scope_override).unwrap_or(Value::Null),
5943 );
5944 let resolved = match resolved {
5945 Some(resolved) => resolved,
5946 None => {
5947 let reason = format!(
5948 "Tool '{}' became unavailable after approval",
5949 request.requested_name
5950 );
5951 let record = self.record_from_parts_at(
5952 &request,
5953 request.requested_name.clone(),
5954 executed_arguments,
5955 started_at,
5956 start,
5957 false,
5958 false,
5959 reason.clone(),
5960 metadata,
5961 ToolPolicyDecisionRecord::unavailable(reason),
5962 approval_record,
5963 false,
5964 false,
5965 versions,
5966 );
5967 self.finish_tool_record(&record).await;
5968 return Ok(record);
5969 }
5970 };
5971
5972 let canonical_id = resolved.identity.canonical_id.clone();
5973 if let Some(reason) =
5974 fallback_state.final_rejection_reason(&admitted_canonical_id, &canonical_id)
5975 {
5976 metadata.insert(
5980 "fallback_chain".to_string(),
5981 serde_json::to_value(&fallback_state.visited_canonical_ids).unwrap_or(Value::Null),
5982 );
5983 metadata.insert(
5984 "final_resolved_canonical_id".to_string(),
5985 Value::String(canonical_id),
5986 );
5987 let record = self.record_from_parts_at(
5988 &request,
5989 admitted_canonical_id,
5990 executed_arguments,
5991 started_at,
5992 start,
5993 false,
5994 false,
5995 format!("Denied: {reason}"),
5996 metadata,
5997 ToolPolicyDecisionRecord::deny(reason),
5998 approval_record,
5999 false,
6000 false,
6001 versions,
6002 );
6003 self.finish_tool_record(&record).await;
6004 return Ok(record);
6005 }
6006 let bindings = resolved.tool.policy_bindings();
6007 let final_arguments = control_snapshot
6008 .tool_security
6009 .prepare_tool_arguments_with_bindings(&canonical_id, &executed_arguments, &bindings);
6010 if let Some(record) = approval_record.as_mut()
6011 && matches!(record.status, ToolApprovalStatus::Modified)
6012 {
6013 record.modified_arguments = Some(final_arguments.clone());
6014 }
6015 let classification = resolved.tool.classify_call(&final_arguments);
6016 let safety = resolved.tool.safety_metadata();
6017 let security_engine = control_snapshot.tool_security;
6018 let tool_config = self.recovery_manager.get_tool_config(&canonical_id).clone();
6019 let recovery_timeout_ms = self.recovery_manager.get_tool_timeout(&canonical_id);
6020 metadata.insert(
6021 "classification".to_string(),
6022 serde_json::to_value(&classification).unwrap_or(Value::Null),
6023 );
6024 let (limits, timeout) = match Self::effective_tool_limits(
6028 &security_engine,
6029 &canonical_id,
6030 &safety,
6031 &classification,
6032 recovery_timeout_ms,
6033 ) {
6034 Ok(effective) => effective,
6035 Err(error) => {
6036 let reason = error.to_string();
6037 metadata.insert(
6038 "configuration_error".to_string(),
6039 Value::String(reason.clone()),
6040 );
6041 let record = self.record_from_parts_at(
6042 &request,
6043 canonical_id,
6044 final_arguments,
6045 started_at,
6046 start,
6047 false,
6048 false,
6049 format!("Denied: {reason}"),
6050 metadata,
6051 ToolPolicyDecisionRecord::deny(reason),
6052 approval_record,
6053 false,
6054 false,
6055 versions,
6056 );
6057 self.finish_tool_record(&record).await;
6058 return Ok(record);
6059 }
6060 };
6061 let policy_snapshot = security_engine.policy_snapshot(&canonical_id);
6062 let resource_lock_keys =
6063 tool_resource_lock_keys(&canonical_id, &final_arguments, &bindings, &classification);
6064 metadata.insert(
6065 "effective_limits".to_string(),
6066 serde_json::to_value(&limits).unwrap_or(Value::Null),
6067 );
6068 metadata.insert(
6069 "resource_lock_keys".to_string(),
6070 serde_json::to_value(&resource_lock_keys).unwrap_or(Value::Null),
6071 );
6072 if policy_snapshot.is_null() {
6073 metadata.remove("policy_snapshot");
6074 } else {
6075 metadata.insert("policy_snapshot".to_string(), policy_snapshot.clone());
6076 }
6077
6078 let final_denial = |canonical_id: String,
6079 output: String,
6080 policy: ToolPolicyDecisionRecord,
6081 metadata: HashMap<String, Value>,
6082 decision_versions: ToolDecisionVersions| {
6083 self.record_from_parts_at(
6084 &request,
6085 canonical_id,
6086 final_arguments.clone(),
6087 started_at,
6088 start,
6089 false,
6090 false,
6091 output,
6092 metadata,
6093 policy,
6094 approval_record.clone(),
6095 false,
6096 false,
6097 decision_versions,
6098 )
6099 };
6100
6101 if control_snapshot.emergency_deny {
6102 let reason = "Tool execution is disabled by runtime control".to_string();
6103 let record = final_denial(
6104 canonical_id,
6105 reason.clone(),
6106 ToolPolicyDecisionRecord::deny(reason),
6107 metadata,
6108 versions,
6109 );
6110 self.finish_tool_record(&record).await;
6111 return Ok(record);
6112 }
6113
6114 let available_snapshot = self
6119 .get_available_tool_ids_snapshot_for_scope(
6120 control_snapshot.tool_scope_override.as_deref(),
6121 )
6122 .await?;
6123 versions.state = available_snapshot.state_generation;
6124 metadata.insert(
6125 "available_tool_ids_snapshot".to_string(),
6126 serde_json::to_value(&available_snapshot.tool_ids).unwrap_or(Value::Null),
6127 );
6128 metadata.insert(
6129 "state_generation_snapshot".to_string(),
6130 serde_json::to_value(available_snapshot.state_generation).unwrap_or(Value::Null),
6131 );
6132 if !available_snapshot
6133 .tool_ids
6134 .iter()
6135 .any(|tool_id| tool_id == &canonical_id)
6136 {
6137 let reason = format!(
6138 "Tool '{}' is not available in the final runtime scope",
6139 canonical_id
6140 );
6141 let record = final_denial(
6142 canonical_id,
6143 reason.clone(),
6144 ToolPolicyDecisionRecord::deny(reason),
6145 metadata,
6146 versions,
6147 );
6148 self.finish_tool_record(&record).await;
6149 return Ok(record);
6150 }
6151
6152 let final_security_result = security_engine
6157 .validate_tool_execution_with_bindings(&canonical_id, &final_arguments, &bindings)
6158 .await?;
6159 match &final_security_result {
6160 SecurityCheckResult::Block { reason } => {
6161 let record = final_denial(
6162 canonical_id,
6163 format!("Denied: {}", reason),
6164 ToolPolicyDecisionRecord::deny(reason.clone()),
6165 metadata,
6166 versions,
6167 );
6168 self.finish_tool_record(&record).await;
6169 return Ok(record);
6170 }
6171 SecurityCheckResult::Unavailable { reason } => {
6172 let record = final_denial(
6173 canonical_id,
6174 format!("Unavailable: {}", reason),
6175 ToolPolicyDecisionRecord::unavailable(reason.clone()),
6176 metadata,
6177 versions,
6178 );
6179 self.finish_tool_record(&record).await;
6180 return Ok(record);
6181 }
6182 SecurityCheckResult::Warn { message } => {
6183 warn!(tool = %canonical_id, message = %message, "Tool security warning after approval");
6184 }
6185 SecurityCheckResult::Allow | SecurityCheckResult::RequireConfirmation { .. } => {}
6186 }
6187 let final_confirmation_required = matches!(
6188 final_security_result,
6189 SecurityCheckResult::RequireConfirmation { .. }
6190 ) || security_engine
6191 .classification_approval_message(&canonical_id, &classification)
6192 .is_some();
6193 let stale_approval = approval_binding.as_ref().is_some_and(|binding| {
6194 binding.is_stale(
6195 &canonical_id,
6196 &final_arguments,
6197 final_confirmation_required,
6198 versions,
6199 &resolved.tool,
6200 )
6201 });
6202 if stale_approval {
6203 let reason = "Approval became stale before final admission".to_string();
6204 let record = final_denial(
6205 canonical_id,
6206 reason.clone(),
6207 ToolPolicyDecisionRecord::deny(reason),
6208 metadata,
6209 versions,
6210 );
6211 self.finish_tool_record(&record).await;
6212 return Ok(record);
6213 }
6214 if final_confirmation_required && approval_binding.is_none() {
6215 let reason = "Final policy requires fresh approval".to_string();
6216 let record = final_denial(
6217 canonical_id,
6218 reason.clone(),
6219 ToolPolicyDecisionRecord::approval(reason),
6220 metadata,
6221 versions,
6222 );
6223 self.finish_tool_record(&record).await;
6224 return Ok(record);
6225 }
6226
6227 if let Some((_, reason)) = self.host_tool_unavailability(&canonical_id) {
6228 let record = final_denial(
6229 canonical_id,
6230 reason.to_string(),
6231 ToolPolicyDecisionRecord::unavailable(reason),
6232 metadata,
6233 versions,
6234 );
6235 self.finish_tool_record(&record).await;
6236 return Ok(record);
6237 }
6238
6239 let Some(resource_guards) = self.acquire_tool_resource_locks(&resource_lock_keys).await
6244 else {
6245 let reason = "Tool execution cancelled while waiting for resource locks".to_string();
6249 let mut record = final_denial(
6250 canonical_id,
6251 reason.clone(),
6252 ToolPolicyDecisionRecord::deny(reason),
6253 metadata,
6254 versions,
6255 );
6256 record.cancelled = true;
6257 record.cancellation_reason = Some("runtime control cancellation".to_string());
6258 self.finish_tool_record(&record).await;
6259 return Ok(record);
6260 };
6261
6262 let admission = self.admit_tool_execution(
6267 versions.runtime_control,
6268 versions.policy,
6269 versions.state,
6270 &canonical_id,
6271 );
6272 if !matches!(admission, SecurityCheckResult::Allow) {
6273 let latest_control = self.runtime_safety_snapshot();
6274 let reason = admission
6275 .reason()
6276 .unwrap_or("tool admission was denied")
6277 .to_string();
6278 let policy = if admission.is_unavailable() {
6279 ToolPolicyDecisionRecord::unavailable(reason.clone())
6280 } else {
6281 ToolPolicyDecisionRecord::deny(reason.clone())
6282 };
6283 let record = self.record_from_parts_at(
6284 &request,
6285 canonical_id,
6286 final_arguments,
6287 started_at,
6288 start,
6289 false,
6290 false,
6291 reason,
6292 metadata,
6293 policy,
6294 approval_record,
6295 false,
6296 false,
6297 ToolDecisionVersions {
6298 policy: latest_control.tool_security.policy_version(),
6299 registry: versions.registry,
6300 runtime_control: latest_control.version,
6301 state: self
6302 .state_machine
6303 .as_ref()
6304 .map(|state_machine| state_machine.generation()),
6305 },
6306 );
6307 self.finish_tool_record_after_resource_guards(resource_guards, &record)
6308 .await;
6309 return Ok(record);
6310 }
6311 let executed_arguments = final_arguments;
6312
6313 let turn_actor = current_turn_actor_context();
6314 let actor = ToolActorContext {
6315 actor_id: turn_actor
6316 .as_ref()
6317 .and_then(|context| context.effective_actor_id().map(str::to_string))
6318 .or_else(|| self.actor_id()),
6319 origin_actor_id: turn_actor
6320 .as_ref()
6321 .and_then(|context| context.origin_actor_id.clone()),
6322 sender_agent_id: turn_actor
6323 .as_ref()
6324 .and_then(|context| context.sender_agent_id.clone()),
6325 };
6326 let tool_context = ToolExecutionContext {
6327 requested_name: request.requested_name.clone(),
6328 canonical_id: canonical_id.clone(),
6329 display_name: resolved.identity.display_name.clone(),
6330 provider_id: resolved.identity.provider_id.clone(),
6331 registry_version: versions.registry,
6332 policy_version: versions.policy,
6333 runtime_control_version: versions.runtime_control,
6334 call_id: request.call_id.clone(),
6335 source: request.source.clone(),
6336 actor,
6337 cancellation: ToolCancellationToken::new(
6338 Arc::clone(&self.runtime_control.emergency_deny),
6339 Some("runtime control cancellation".to_string()),
6340 ),
6341 started_at,
6342 deadline: None,
6343 permission: ToolPolicyDecisionRecord::allow(),
6344 approval: approval_record.clone(),
6345 classification: classification.clone(),
6346 safety,
6347 limits: limits.clone(),
6348 policy_snapshot,
6349 custom_config: security_engine.custom_config(&canonical_id),
6350 };
6351 let (mut result, timed_out, cancelled, invoked) = self
6352 .run_tool_with_retries(
6353 &canonical_id,
6354 resolved.tool.clone(),
6355 executed_arguments.clone(),
6356 tool_context,
6357 timeout,
6358 tool_config.max_retries,
6359 )
6360 .await?;
6361
6362 let fallback_tool = if !result.success && !cancelled {
6366 match &tool_config.on_failure {
6367 ToolFailureAction::Skip => {
6368 result = ToolResult::ok(format!(
6369 "{{\"skipped\": true, \"reason\": \"Tool '{}' was skipped after failure\"}}",
6370 canonical_id
6371 ));
6372 None
6373 }
6374 ToolFailureAction::Fallback { fallback_tool } => Some(fallback_tool.clone()),
6375 ToolFailureAction::ReportError => None,
6376 }
6377 } else {
6378 None
6379 };
6380
6381 let output_cap = limits.max_output_chars;
6382 let (output, output_truncated) =
6383 Self::truncate_tool_output(result.output.clone(), output_cap);
6384 if let Some(result_metadata) = result.metadata {
6385 metadata.extend(result_metadata);
6386 }
6387 let mut record = self.record_from_parts_at(
6388 &request,
6389 canonical_id,
6390 executed_arguments,
6391 started_at,
6392 start,
6393 invoked,
6394 result.success,
6395 output,
6396 metadata,
6397 ToolPolicyDecisionRecord::allow(),
6398 approval_record,
6399 timed_out,
6400 output_truncated,
6401 versions,
6402 );
6403 record.cancelled = cancelled;
6404 if cancelled {
6405 record.cancellation_reason = Some("runtime control cancellation".to_string());
6406 }
6407 if let Some(fallback_tool) = fallback_tool {
6408 let fallback_arguments = record.executed_arguments.clone();
6409 let original_tool = record.canonical_id.clone();
6410 self.finish_tool_record_after_resource_guards(resource_guards, &record)
6414 .await;
6415 let fallback_request = ToolExecutionRequest::new(
6416 request.call_id.clone(),
6417 fallback_tool,
6418 fallback_arguments,
6419 ToolCallSource::Fallback { original_tool },
6420 );
6421 return Box::pin(self.execute_tool_record_inner(fallback_request, fallback_state))
6422 .await;
6423 }
6424 self.finish_tool_record_after_resource_guards(resource_guards, &record)
6425 .await;
6426 Ok(record)
6427 }
6428
6429 #[instrument(skip(self, tool_call), fields(tool = %tool_call.name))]
6430 async fn execute_tool_smart(&self, tool_call: &ToolCall) -> Result<String> {
6431 let record = self
6432 .execute_tool_record(ToolExecutionRequest::new(
6433 tool_call.id.clone(),
6434 tool_call.name.clone(),
6435 tool_call.arguments.clone(),
6436 ToolCallSource::Model,
6437 ))
6438 .await?;
6439 if record.success {
6440 Ok(record.model_output_string())
6441 } else if matches!(record.policy.outcome, PermissionOutcome::RequiresApproval) {
6442 Err(AgentError::HITLRejected(record.model_output_string()))
6443 } else {
6444 Err(AgentError::Tool(record.model_output_string()))
6445 }
6446 }
6447
6448 async fn select_skill_candidate(&self, input: &str) -> Result<Option<SkillCandidate>> {
6454 let Some(ref router) = self.skill_router else {
6455 return Ok(None);
6456 };
6457 let available_skills = self.get_available_skills();
6458 if available_skills.is_empty() {
6459 return Ok(None);
6460 }
6461 let skill_ids: Vec<&str> = available_skills.iter().map(|s| s.id.as_str()).collect();
6462 let Some(skill_id) = self
6463 .observe_purpose(
6464 ObservationPurpose::SkillRouting,
6465 router.select_skill_filtered(input, &skill_ids),
6466 )
6467 .await?
6468 else {
6469 return Ok(None);
6470 };
6471 let skill = router
6472 .get_skill(&skill_id)
6473 .cloned()
6474 .ok_or_else(|| AgentError::Skill(format!("Skill not found: {}", skill_id)))?;
6475 info!(skill_id = %skill_id, "Skill selected");
6476 Ok(Some(SkillCandidate::new(skill_id, skill)))
6477 }
6478
6479 async fn commit_skill_candidate_route_result(
6484 &self,
6485 candidate: SkillCandidate,
6486 input: &str,
6487 ) -> Result<SkillRouteResult> {
6488 let skill_id = candidate.skill_id;
6489 let skill = candidate.skill;
6490 let expected_state_generation = self
6491 .state_machine
6492 .as_ref()
6493 .map(|state_machine| state_machine.generation());
6494 let expected_disambiguation_epoch = self.disambiguation_epoch.load(Ordering::SeqCst);
6495 if let Some(ref skill_disambig) = skill.disambiguation
6496 && skill_disambig.enabled.unwrap_or(false)
6497 && let Some(ref disambiguator) = self.disambiguation_manager
6498 {
6499 let context = self.build_disambiguation_context().await?;
6500 let state_override = self
6501 .state_machine
6502 .as_ref()
6503 .and_then(|sm| sm.current_definition())
6504 .and_then(|def| def.disambiguation.clone());
6505
6506 let disambiguation_result = self
6507 .observe_purpose(
6508 ObservationPurpose::DisambiguationDetection,
6509 disambiguator.process_input_with_override(
6510 input,
6511 &context,
6512 state_override.as_ref(),
6513 Some(skill_disambig),
6514 ),
6515 )
6516 .await?;
6517 let current_state_generation = self
6518 .state_machine
6519 .as_ref()
6520 .map(|state_machine| state_machine.generation());
6521 if current_state_generation != expected_state_generation
6522 || self.disambiguation_epoch.load(Ordering::SeqCst) != expected_disambiguation_epoch
6523 {
6524 disambiguator.clear_pending().await;
6525 *self.pending_skill_id.write() = None;
6526 return Err(AgentError::Other(
6527 "State or reset ownership changed during skill disambiguation".to_string(),
6528 ));
6529 }
6530 match disambiguation_result {
6531 DisambiguationResult::Clear => {
6532 debug!(skill_id = %skill_id, "Skill disambiguation: clear");
6533 }
6534 DisambiguationResult::NeedsClarification {
6535 question,
6536 detection,
6537 } => {
6538 let admission = self
6539 .admit_disambiguation_redispatch(
6540 expected_disambiguation_epoch,
6541 expected_state_generation,
6542 )
6543 .await?;
6544 let awaiting_confirmation = disambiguator.has_pending_confirmation().await;
6545 info!(
6546 skill_id = %skill_id,
6547 ambiguity_type = ?detection.ambiguity_type,
6548 confidence = detection.confidence,
6549 "Skill requires clarification before execution"
6550 );
6551 *self.pending_skill_id.write() = Some(skill_id.clone());
6552 let response = AgentResponse::new(&question.question).with_metadata(
6553 "disambiguation",
6554 serde_json::json!({
6555 "status": if awaiting_confirmation { "awaiting_confirmation" } else { "awaiting_clarification" },
6556 "skill_id": skill_id,
6557 "options": question.options,
6558 "clarifying": question.clarifying,
6559 "detection": {
6560 "type": detection.ambiguity_type,
6561 "confidence": detection.confidence,
6562 "what_is_unclear": detection.what_is_unclear,
6563 }
6564 }),
6565 );
6566 drop(admission);
6567 return Ok(SkillRouteResult::NeedsClarification {
6568 response,
6569 ownership: Some(DisambiguationOwnership {
6570 epoch: expected_disambiguation_epoch,
6571 state_generation: expected_state_generation,
6572 }),
6573 });
6574 }
6575 DisambiguationResult::Clarified { enriched_input, .. } => {
6576 info!(skill_id = %skill_id, enriched = %enriched_input, "Skill disambiguation clarified");
6577 let admission = self
6578 .admit_disambiguation_redispatch(
6579 expected_disambiguation_epoch,
6580 expected_state_generation,
6581 )
6582 .await?;
6583 drop(admission);
6584 let content = self.execute_skill(&skill, &enriched_input).await?;
6585 return Ok(SkillRouteResult::Response { skill_id, content });
6586 }
6587 DisambiguationResult::ProceedWithBestGuess { enriched_input } => {
6588 info!(skill_id = %skill_id, "Skill disambiguation best guess");
6589 let admission = self
6590 .admit_disambiguation_redispatch(
6591 expected_disambiguation_epoch,
6592 expected_state_generation,
6593 )
6594 .await?;
6595 drop(admission);
6596 let content = self.execute_skill(&skill, &enriched_input).await?;
6597 return Ok(SkillRouteResult::Response { skill_id, content });
6598 }
6599 DisambiguationResult::GiveUp { reason } => {
6600 warn!(skill_id = %skill_id, reason = %reason, "Skill disambiguation gave up");
6601 let apology = self
6602 .generate_localized_apology(
6603 "Generate a brief, polite apology saying you couldn't understand the request. Be concise.",
6604 &reason,
6605 )
6606 .await
6607 .unwrap_or_else(|_| {
6608 format!("I'm sorry, I couldn't understand your request: {}", reason)
6609 });
6610 return Ok(SkillRouteResult::NeedsClarification {
6611 response: AgentResponse::new(&apology),
6612 ownership: None,
6613 });
6614 }
6615 DisambiguationResult::Escalate { reason } => {
6616 info!(skill_id = %skill_id, reason = %reason, "Skill disambiguation escalating");
6617 let apology = self
6618 .generate_localized_apology(
6619 "Explain briefly that you're transferring the user to a human agent for help.",
6620 &reason,
6621 )
6622 .await
6623 .unwrap_or_else(|_| {
6624 format!("I need human assistance to help with your request: {}", reason)
6625 });
6626 return Ok(SkillRouteResult::NeedsClarification {
6627 response: AgentResponse::new(&apology),
6628 ownership: None,
6629 });
6630 }
6631 DisambiguationResult::Abandoned { .. } => {
6632 debug!(skill_id = %skill_id, "Skill disambiguation abandoned");
6633 return Ok(SkillRouteResult::NoMatch);
6634 }
6635 }
6636 }
6637 let admission = self
6638 .admit_disambiguation_redispatch(
6639 expected_disambiguation_epoch,
6640 expected_state_generation,
6641 )
6642 .await?;
6643 drop(admission);
6644 let content = self.execute_skill(&skill, input).await?;
6645 Ok(SkillRouteResult::Response { skill_id, content })
6646 }
6647
6648 async fn try_skill_route(&self, input: &str) -> Result<SkillRouteResult> {
6650 if let Some(candidate) = self.select_skill_candidate(input).await? {
6651 self.commit_skill_candidate_route_result(candidate, input)
6652 .await
6653 } else {
6654 Ok(SkillRouteResult::NoMatch)
6655 }
6656 }
6657
6658 fn skill_clarification_needs_memory_record(response: &AgentResponse) -> bool {
6661 response
6662 .metadata
6663 .as_ref()
6664 .and_then(|m| m.get("disambiguation"))
6665 .and_then(|d| d.get("status"))
6666 .and_then(|s| s.as_str())
6667 == Some("awaiting_clarification")
6668 }
6669
6670 async fn commit_winning_skill_candidate(
6677 &self,
6678 candidate: SkillCandidate,
6679 processed_input: &str,
6680 input_context: &HashMap<String, Value>,
6681 ) -> Result<Option<AgentResponse>> {
6682 self.commit_root_user_message(processed_input).await?;
6683 match self
6684 .commit_skill_candidate_route_result(candidate, processed_input)
6685 .await?
6686 {
6687 SkillRouteResult::Response { skill_id, content } => self
6688 .handle_skill_response(processed_input, &skill_id, content, input_context)
6689 .await
6690 .map(Some),
6691 SkillRouteResult::NeedsClarification {
6692 response,
6693 ownership,
6694 } => {
6695 let admission = self
6696 .admit_optional_disambiguation_ownership(ownership)
6697 .await?;
6698 if Self::skill_clarification_needs_memory_record(&response) {
6699 self.memory
6700 .add_message(ChatMessage::assistant(&response.content))
6701 .await?;
6702 }
6703 drop(admission);
6704 self.finish_turn_if_root(&response).await?;
6705 Ok(Some(response))
6706 }
6707 SkillRouteResult::NoMatch => Ok(None),
6708 }
6709 }
6710
6711 async fn execute_skill(&self, skill: &SkillDefinition, input: &str) -> Result<String> {
6713 if let Some(ref executor) = self.skill_executor {
6714 let skill_reasoning = self.get_skill_reasoning_config(skill);
6715 let skill_reflection = self.get_skill_reflection_config(skill);
6716
6717 debug!(
6718 skill_id = %skill.id,
6719 reasoning_mode = ?skill_reasoning.mode,
6720 reflection_enabled = ?skill_reflection.enabled,
6721 "Skill reasoning/reflection config"
6722 );
6723
6724 let response = self
6725 .observe_purpose(
6726 ObservationPurpose::SkillPrompt,
6727 executor.execute_with_invoker(skill, input, serde_json::json!({}), self),
6728 )
6729 .await?;
6730
6731 if skill_reflection.requires_evaluation() && skill_reflection.is_enabled() {
6732 let should_reflect = self
6733 .should_reflect_with_config(input, &response, &skill_reflection)
6734 .await?;
6735 if should_reflect {
6736 let evaluated = self
6737 .evaluate_and_retry_with_config(input, response, &skill_reflection)
6738 .await?;
6739 return Ok(evaluated);
6740 }
6741 }
6742
6743 return Ok(response);
6744 }
6745 Err(AgentError::Skill(
6746 "No skill executor configured".to_string(),
6747 ))
6748 }
6749
6750 async fn execute_skill_by_id(&self, skill_id: &str, input: &str) -> Result<String> {
6753 let skill = self
6754 .skill_router
6755 .as_ref()
6756 .and_then(|r| r.get_skill(skill_id).cloned())
6757 .ok_or_else(|| AgentError::Skill(format!("Skill not found: {}", skill_id)))?;
6758 self.execute_skill(&skill, input).await
6759 }
6760
6761 async fn should_reflect_with_config(
6763 &self,
6764 input: &str,
6765 response: &str,
6766 config: &ReflectionConfig,
6767 ) -> Result<bool> {
6768 if !config.requires_evaluation() {
6769 return Ok(false);
6770 }
6771
6772 if config.is_enabled() {
6773 return Ok(true);
6774 }
6775
6776 let evaluator_llm = config
6777 .evaluator_llm
6778 .as_ref()
6779 .and_then(|alias| self.llm_registry.get(alias).ok())
6780 .or_else(|| self.llm_registry.router().ok())
6781 .or_else(|| self.llm_registry.default().ok());
6782
6783 let Some(llm) = evaluator_llm else {
6784 return Ok(false);
6785 };
6786
6787 let response_preview: String = response.chars().take(500).collect();
6788 let prompt = format!(
6789 r#"Should this response be evaluated for quality? Consider if it's a complex or important response.
6790
6791User query: "{}"
6792Response: "{}"
6793
6794Answer YES or NO only."#,
6795 input, response_preview
6796 );
6797
6798 let messages = vec![ChatMessage::user(&prompt)];
6799 let result = self
6800 .observe_purpose(
6801 ObservationPurpose::ReflectionDecision,
6802 llm.complete(&messages, None),
6803 )
6804 .await;
6805
6806 match result {
6807 Ok(resp) => Ok(resp.content.trim().to_uppercase().contains("YES")),
6808 Err(_) => Ok(false),
6809 }
6810 }
6811
6812 async fn evaluate_and_retry_with_config(
6813 &self,
6814 input: &str,
6815 mut response: String,
6816 config: &ReflectionConfig,
6817 ) -> Result<String> {
6818 let llm = self.get_state_llm()?;
6819 let mut attempts = 0u32;
6820 let max_retries = config.max_retries;
6821
6822 loop {
6823 let evaluation = self
6824 .evaluate_response_with_config(input, &response, config)
6825 .await?;
6826
6827 if evaluation.passed || attempts >= max_retries {
6828 info!(
6829 passed = evaluation.passed,
6830 confidence = evaluation.confidence,
6831 attempts = attempts + 1,
6832 "Skill reflection evaluation complete"
6833 );
6834 return Ok(response);
6835 }
6836
6837 debug!(
6838 attempt = attempts + 1,
6839 failed_criteria = evaluation.failed_criteria().count(),
6840 "Skill response did not meet criteria, retrying"
6841 );
6842
6843 let feedback: Vec<String> = evaluation
6844 .failed_criteria()
6845 .map(|c| format!("- {}", c.criterion))
6846 .collect();
6847
6848 let retry_prompt = format!(
6849 "Your previous response did not meet these criteria:\n{}\n\nPlease provide an improved response to: {}",
6850 feedback.join("\n"),
6851 input
6852 );
6853
6854 let messages = vec![ChatMessage::user(&retry_prompt)];
6855 let retry_response = self
6856 .observe_purpose(
6857 ObservationPurpose::ReflectionEvaluation,
6858 llm.complete(&messages, None),
6859 )
6860 .await
6861 .map_err(|e| AgentError::LLM(e.to_string()))?;
6862
6863 response = retry_response.content.trim().to_string();
6864 attempts += 1;
6865 }
6866 }
6867
6868 async fn evaluate_response_with_config(
6869 &self,
6870 input: &str,
6871 response: &str,
6872 config: &ReflectionConfig,
6873 ) -> Result<EvaluationResult> {
6874 let evaluator_llm = config
6875 .evaluator_llm
6876 .as_ref()
6877 .and_then(|alias| self.llm_registry.get(alias).ok())
6878 .or_else(|| self.llm_registry.router().ok())
6879 .or_else(|| self.llm_registry.default().ok())
6880 .ok_or_else(|| AgentError::Config("No LLM available for evaluation".into()))?;
6881
6882 let criteria = &config.criteria;
6883 let criteria_list = criteria
6884 .iter()
6885 .enumerate()
6886 .map(|(i, c)| format!("{}. {}", i + 1, c))
6887 .collect::<Vec<_>>()
6888 .join("\n");
6889
6890 let prompt = format!(
6891 r#"Evaluate this response against the criteria.
6892
6893User query: "{}"
6894
6895Response to evaluate: "{}"
6896
6897Criteria:
6898{}
6899
6900For each criterion, respond with:
6901- criterion number
6902- PASS or FAIL
6903- brief reason
6904
6905Then provide overall confidence (0.0 to 1.0) and whether it passes overall.
6906
6907Format:
69081. PASS/FAIL - reason
69092. PASS/FAIL - reason
6910...
6911CONFIDENCE: 0.X
6912OVERALL: PASS/FAIL"#,
6913 input, response, criteria_list
6914 );
6915
6916 let messages = vec![ChatMessage::user(&prompt)];
6917 let eval_response = self
6918 .observe_purpose(
6919 ObservationPurpose::ReflectionEvaluation,
6920 evaluator_llm.complete(&messages, None),
6921 )
6922 .await
6923 .map_err(|e| AgentError::LLM(format!("Evaluation failed: {}", e)))?;
6924
6925 let content = eval_response.content.to_uppercase();
6926 let llm_pass = content.contains("OVERALL: PASS");
6927
6928 let confidence = content
6929 .lines()
6930 .find(|l| l.contains("CONFIDENCE:"))
6931 .and_then(|l| {
6932 l.split(':')
6933 .nth(1)
6934 .and_then(|v| v.trim().parse::<f32>().ok())
6935 })
6936 .unwrap_or(if llm_pass { 0.8 } else { 0.4 });
6937
6938 let overall_pass = llm_pass && confidence >= config.pass_threshold;
6941
6942 let mut criteria_results = Vec::new();
6943 for (i, criterion) in criteria.iter().enumerate() {
6944 let line_marker = format!("{}.", i + 1);
6945 let passed = eval_response
6946 .content
6947 .lines()
6948 .find(|l| l.contains(&line_marker))
6949 .map(|l| l.to_uppercase().contains("PASS"))
6950 .unwrap_or(overall_pass);
6951
6952 if passed {
6953 criteria_results.push(CriterionResult::pass(criterion));
6954 } else {
6955 criteria_results.push(CriterionResult::fail(criterion, "Did not meet criterion"));
6956 }
6957 }
6958
6959 Ok(EvaluationResult::new(overall_pass, confidence).with_criteria(criteria_results))
6960 }
6961
6962 async fn process_input(&self, input: &str) -> Result<ProcessData> {
6964 if let Some(processor) = self.get_state_process_processor() {
6965 let purpose = observation_purpose_for_process(processor.input_purpose_hint());
6966 return self
6967 .observe_purpose(purpose, processor.process_input(input))
6968 .await;
6969 }
6970 if let Some(ref processor) = self.process_processor {
6971 let purpose = observation_purpose_for_process(processor.input_purpose_hint());
6972 self.observe_purpose(purpose, processor.process_input(input))
6973 .await
6974 } else {
6975 Ok(ProcessData::new(input))
6976 }
6977 }
6978
6979 async fn process_output(
6981 &self,
6982 output: &str,
6983 input_context: &std::collections::HashMap<String, serde_json::Value>,
6984 ) -> Result<ProcessData> {
6985 if let Some(processor) = self.get_state_process_processor() {
6986 let purpose = observation_purpose_for_process(processor.output_purpose_hint());
6987 return self
6988 .observe_purpose(purpose, processor.process_output(output, input_context))
6989 .await;
6990 }
6991 if let Some(ref processor) = self.process_processor {
6992 let purpose = observation_purpose_for_process(processor.output_purpose_hint());
6993 self.observe_purpose(purpose, processor.process_output(output, input_context))
6994 .await
6995 } else {
6996 Ok(ProcessData::new(output))
6997 }
6998 }
6999
7000 fn get_state_process_processor(&self) -> Option<ProcessProcessor> {
7002 let sm = self.state_machine.as_ref()?;
7003 let def = sm.current_definition()?;
7004 let config = def.process.as_ref()?;
7005 let mut processor = ProcessProcessor::new(config.clone());
7006 if let Some(ref registry) = Some(self.llm_registry.clone()) {
7007 processor = processor.with_llm_registry(registry.clone());
7008 }
7009 processor = processor.with_stage_observer(Arc::new(ObservabilityProcessStageObserver));
7010 Some(processor)
7011 }
7012
7013 async fn check_turn_timeout(&self) -> Result<()> {
7015 let Some(ref sm) = self.state_machine else {
7016 return Ok(());
7017 };
7018 let Some(timeout_state) = sm.check_timeout() else {
7019 return Ok(());
7020 };
7021 let claim_admission = self.disambiguation_admission.write().await;
7022 if sm.check_timeout().as_deref() != Some(timeout_state.as_str()) {
7023 return Ok(());
7024 }
7025 let Some(reservation) = self.reserve_state_transition() else {
7026 return Ok(());
7027 };
7028 let from_state = sm.current();
7029 let expected_state_generation = sm.generation();
7030 let expected_disambiguation_epoch = self.disambiguation_epoch.load(Ordering::SeqCst);
7031 let history_before = sm.history();
7032 drop(claim_admission);
7033
7034 self.execute_state_exit_actions(&from_state).await;
7035
7036 let admission = self.disambiguation_admission.write().await;
7037 if sm.current() != from_state
7038 || sm.generation() != expected_state_generation
7039 || self.disambiguation_epoch.load(Ordering::SeqCst) != expected_disambiguation_epoch
7040 || sm.check_timeout().as_deref() != Some(timeout_state.as_str())
7041 {
7042 return Ok(());
7043 }
7044 sm.transition_to(&timeout_state, "max_turns exceeded")?;
7045 self.invalidate_pending_confirmation("state_timeout").await;
7046 let entered = sm.current();
7047 let is_reentry = Self::state_was_previously_entered(&entered, &from_state, &history_before);
7048 drop(admission);
7049
7050 self.execute_state_enter_actions(&entered, is_reentry).await;
7051 drop(reservation);
7052 info!(to = %entered, "Timeout transition");
7053 Ok(())
7054 }
7055
7056 fn increment_turn(&self) {
7057 if let Some(ref sm) = self.state_machine {
7058 sm.increment_turn();
7059 }
7060 }
7061
7062 fn transitions_available_for_commit(&self) -> Option<(Vec<Transition>, String)> {
7063 let sm = self.state_machine.as_ref()?;
7064 let current = sm.current();
7065 let transitions: Vec<_> = sm
7066 .auto_transitions()
7067 .into_iter()
7068 .filter(|t| match t.cooldown_turns {
7069 Some(cd) if cd > 0 => {
7070 let resolved = sm.config().resolve_full_path(¤t, &t.to);
7071 !sm.is_on_cooldown(&resolved, cd)
7072 }
7073 _ => true,
7074 })
7075 .collect();
7076 Some((transitions, current))
7077 }
7078
7079 fn transition_reason(transition: &Transition) -> String {
7080 if transition.when.is_empty() {
7081 "guard condition met".to_string()
7082 } else {
7083 transition.when.clone()
7084 }
7085 }
7086
7087 fn build_transition_context(
7089 &self,
7090 user_message: &str,
7091 response: &str,
7092 current_state: &str,
7093 staged: Option<&HashMap<String, Value>>,
7094 ) -> TransitionContext {
7095 let context_map = staged
7096 .map(|writes| self.build_context_with_staged(writes))
7097 .unwrap_or_else(|| self.build_context_with_overlays());
7098 TransitionContext::new(user_message, response, current_state).with_context(context_map)
7099 }
7100
7101 async fn select_transition_candidate(
7103 &self,
7104 user_message: &str,
7105 response: &str,
7106 ) -> Result<Option<TransitionCandidate>> {
7107 let Some((transitions, current_state)) = self.transitions_available_for_commit() else {
7108 return Ok(None);
7109 };
7110 let transitions: Vec<Transition> = transitions
7111 .into_iter()
7112 .filter(|transition| matches!(transition.timing, TransitionTiming::PostResponse))
7113 .collect();
7114 if transitions.is_empty() {
7115 return Ok(None);
7116 }
7117 let Some(evaluator) = self.transition_evaluator.as_ref() else {
7118 return Ok(None);
7119 };
7120 let context = self.build_transition_context(user_message, response, ¤t_state, None);
7121 let selected = self
7122 .observe_purpose(
7123 ObservationPurpose::StateTransitionEvaluation,
7124 evaluator.select_transition(&transitions, &context),
7125 )
7126 .await?;
7127 Ok(selected.map(|index| {
7128 let transition = transitions[index].clone();
7129 TransitionCandidate::new(
7130 current_state,
7131 transition.clone(),
7132 Self::transition_reason(&transition),
7133 )
7134 }))
7135 }
7136
7137 fn select_deterministic_transition_candidate(
7139 &self,
7140 user_message: &str,
7141 current_state: &str,
7142 transitions: &[Transition],
7143 staged: &HashMap<String, Value>,
7144 ) -> Option<TransitionCandidate> {
7145 let context = self.build_transition_context(user_message, "", current_state, Some(staged));
7146
7147 for transition in transitions {
7148 if let Some(guard) = transition.guard.as_ref()
7149 && evaluate_guard(guard, &context)
7150 {
7151 return Some(TransitionCandidate::new(
7152 current_state,
7153 transition.clone(),
7154 Self::transition_reason(transition),
7155 ));
7156 }
7157 }
7158
7159 let resolved_intent = context
7160 .context
7161 .get("resolved_intent")
7162 .and_then(Value::as_str)
7163 .filter(|value| !value.is_empty());
7164 if let Some(resolved_intent) = resolved_intent {
7165 for transition in transitions {
7166 if transition.intent.as_deref() == Some(resolved_intent) {
7167 return Some(TransitionCandidate::new(
7168 current_state,
7169 transition.clone(),
7170 Self::transition_reason(transition),
7171 ));
7172 }
7173 }
7174 }
7175
7176 None
7177 }
7178
7179 async fn commit_transition_candidate(&self, candidate: &TransitionCandidate) -> Result<bool> {
7181 self.commit_transition_target(&candidate.from_state, candidate.target(), &candidate.reason)
7182 .await
7183 }
7184
7185 async fn approve_transition_target(&self, from_state: &str, target: &str) -> Result<bool> {
7187 let approved = self.check_state_hitl(Some(from_state), target).await?;
7188 if !approved {
7189 info!(to = %target, "State transition rejected by HITL");
7190 }
7191 Ok(approved)
7192 }
7193
7194 async fn apply_transition_target(
7196 &self,
7197 from_state: &str,
7198 target: &str,
7199 reason: &str,
7200 staged: Option<&HashMap<String, Value>>,
7201 ) -> Result<bool> {
7202 let Some(ref sm) = self.state_machine else {
7203 return Ok(false);
7204 };
7205 let claim_admission = self.disambiguation_admission.write().await;
7206 if sm.current() != from_state {
7207 return Ok(false);
7208 }
7209 let Some(reservation) = self.reserve_state_transition() else {
7210 return Ok(false);
7211 };
7212 let expected_state_generation = sm.generation();
7213 let expected_disambiguation_epoch = self.disambiguation_epoch.load(Ordering::SeqCst);
7214 let history_before = sm.history();
7215 drop(claim_admission);
7216
7217 self.execute_state_exit_actions(from_state).await;
7218
7219 let admission = self.disambiguation_admission.write().await;
7220 if sm.current() != from_state
7221 || sm.generation() != expected_state_generation
7222 || self.disambiguation_epoch.load(Ordering::SeqCst) != expected_disambiguation_epoch
7223 {
7224 return Ok(false);
7225 }
7226 sm.transition_to(target, reason)?;
7227 self.invalidate_pending_confirmation("state_transition")
7228 .await;
7229 sm.reset_no_transition();
7230 if let Some(staged) = staged {
7231 self.commit_staged_context_writes(staged);
7232 }
7233 let entered = sm.current();
7234 let is_reentry = Self::state_was_previously_entered(&entered, from_state, &history_before);
7235 drop(admission);
7236
7237 self.execute_state_enter_actions(&entered, is_reentry).await;
7238 drop(reservation);
7239 self.hooks
7240 .on_state_transition(Some(from_state), &entered, reason)
7241 .await;
7242 info!(from = %from_state, to = %entered, "State transition");
7243 Ok(true)
7244 }
7245
7246 async fn commit_transition_target(
7248 &self,
7249 from_state: &str,
7250 target: &str,
7251 reason: &str,
7252 ) -> Result<bool> {
7253 if !self.approve_transition_target(from_state, target).await? {
7254 return Ok(false);
7255 }
7256 self.apply_transition_target(from_state, target, reason, None)
7257 .await
7258 }
7259
7260 async fn apply_pre_response_transition_candidate(
7262 &self,
7263 candidate: &TransitionCandidate,
7264 staged: &HashMap<String, Value>,
7265 processed_input: &str,
7266 ) -> Result<bool> {
7267 self.commit_root_user_message(processed_input).await?;
7268 self.apply_transition_target(
7269 &candidate.from_state,
7270 candidate.target(),
7271 &candidate.reason,
7272 Some(staged),
7273 )
7274 .await
7275 }
7276
7277 async fn commit_pre_response_transition_candidate(
7279 &self,
7280 candidate: &TransitionCandidate,
7281 staged: &HashMap<String, Value>,
7282 processed_input: &str,
7283 ) -> Result<bool> {
7284 if !self
7285 .approve_transition_target(&candidate.from_state, candidate.target())
7286 .await?
7287 {
7288 return Ok(false);
7289 }
7290 self.apply_pre_response_transition_candidate(candidate, staged, processed_input)
7291 .await
7292 }
7293
7294 async fn handle_transition_miss(&self, current_state: &str) -> Result<bool> {
7296 let Some(ref sm) = self.state_machine else {
7297 return Ok(false);
7298 };
7299 sm.increment_no_transition();
7300 let Some(fallback) = sm.check_fallback() else {
7301 return Ok(false);
7302 };
7303 self.commit_transition_target(current_state, &fallback, "fallback after no transitions")
7304 .await
7305 }
7306
7307 async fn evaluate_transitions(&self, user_message: &str, response: &str) -> Result<bool> {
7309 let Some((transitions, current_state)) = self.transitions_available_for_commit() else {
7310 return Ok(false);
7311 };
7312 if transitions.is_empty() {
7313 return Ok(false);
7314 }
7315 if let Some(candidate) = self
7316 .select_transition_candidate(user_message, response)
7317 .await?
7318 {
7319 return self.commit_transition_candidate(&candidate).await;
7320 }
7321 self.handle_transition_miss(¤t_state).await
7322 }
7323
7324 async fn try_pre_response_transition(
7326 &self,
7327 processed_input: &str,
7328 ) -> Result<Option<AgentResponse>> {
7329 let optimization = &self.runtime_config.optimization;
7330 if !optimization.enabled || !optimization.pre_response_deterministic_transitions {
7331 return Ok(None);
7332 }
7333 let Some((transitions, current_state)) = self.transitions_available_for_commit() else {
7334 return Ok(None);
7335 };
7336 let eligible: Vec<Transition> = transitions
7337 .into_iter()
7338 .filter(|transition| !transition.requires_response)
7339 .filter(|transition| matches!(transition.timing, TransitionTiming::PreResponse))
7340 .collect();
7341 if eligible.is_empty() {
7342 return Ok(None);
7343 }
7344
7345 let empty_staged = HashMap::new();
7346 let mut extracted_staged: Option<HashMap<String, Value>> = None;
7347 let mut selected: Option<(TransitionCandidate, HashMap<String, Value>)> = None;
7348
7349 for transition in &eligible {
7350 let use_extractors = optimization.pre_response_extractors || transition.run_extractors;
7351 let staged_for_eval = if use_extractors {
7352 if extracted_staged.is_none() {
7353 extracted_staged =
7354 Some(self.run_context_extractors_staged(processed_input).await);
7355 }
7356 extracted_staged.as_ref().unwrap_or(&empty_staged)
7357 } else {
7358 &empty_staged
7359 };
7360
7361 if let Some(candidate) = self.select_deterministic_transition_candidate(
7362 processed_input,
7363 ¤t_state,
7364 std::slice::from_ref(transition),
7365 staged_for_eval,
7366 ) {
7367 let staged_for_commit = if use_extractors {
7368 staged_for_eval.clone()
7369 } else {
7370 HashMap::new()
7371 };
7372 selected = Some((candidate, staged_for_commit));
7373 break;
7374 }
7375 }
7376
7377 let Some((candidate, staged)) = selected else {
7378 return Ok(None);
7379 };
7380
7381 if !self
7382 .commit_pre_response_transition_candidate(&candidate, &staged, processed_input)
7383 .await?
7384 {
7385 return Ok(None);
7386 }
7387 self.redispatch_current_state(processed_input)
7388 .await
7389 .map(Some)
7390 }
7391
7392 async fn try_speculative_branches(
7397 &self,
7398 processed_input: &str,
7399 input_context: &HashMap<String, Value>,
7400 ) -> Result<Option<AgentResponse>> {
7401 let optimization = &self.runtime_config.optimization;
7402 if !optimization.enabled {
7403 return Ok(None);
7404 }
7405
7406 let effective_reasoning_mode = self.get_effective_reasoning_config().mode.clone();
7407 if !matches!(
7408 effective_reasoning_mode,
7409 ReasoningMode::None | ReasoningMode::Auto
7410 ) {
7411 return Ok(None);
7412 }
7413
7414 let mut transition_enabled =
7415 optimization.speculative_state_transitions && self.has_parallel_transition_candidates();
7416 let mut skill_enabled = optimization.speculative_skill_routing
7417 && self.skill_router.is_some()
7418 && self.pending_skill_id.read().is_none();
7419 let mut reasoning_enabled = optimization.speculative_reasoning_auto
7420 && matches!(effective_reasoning_mode, ReasoningMode::Auto);
7421
7422 if matches!(effective_reasoning_mode, ReasoningMode::Auto)
7423 && (!reasoning_enabled || optimization.max_speculative_llm_calls_per_turn < 2)
7424 {
7425 return Ok(None);
7426 }
7427
7428 if !transition_enabled && !skill_enabled && !reasoning_enabled {
7429 return Ok(None);
7430 }
7431
7432 let mut optional_slots = optimization.max_parallel_runtime_tasks.saturating_sub(1);
7433 let mut speculative_call_slots = optimization
7434 .max_speculative_llm_calls_per_turn
7435 .saturating_sub(1);
7436 if reasoning_enabled {
7437 if optional_slots == 0 || speculative_call_slots == 0 {
7438 return Ok(None);
7439 }
7440 optional_slots -= 1;
7441 speculative_call_slots -= 1;
7442 }
7443 if transition_enabled {
7444 if optional_slots == 0 {
7445 transition_enabled = false;
7446 } else {
7447 optional_slots -= 1;
7448 }
7449 }
7450 if skill_enabled && (optional_slots == 0 || speculative_call_slots == 0) {
7451 skill_enabled = false;
7452 }
7453
7454 if !transition_enabled && !skill_enabled && !reasoning_enabled {
7455 return Ok(None);
7456 }
7457
7458 let main_kind = if transition_enabled {
7459 RuntimeOptimizationKind::ParallelStateTransition
7460 } else if skill_enabled {
7461 RuntimeOptimizationKind::SpeculativeSkillRouting
7462 } else {
7463 RuntimeOptimizationKind::SpeculativeReasoningAuto
7464 };
7465 if !self.reserve_active_speculative_llm_call(main_kind) {
7466 return Ok(None);
7467 }
7468
7469 let mut branch_set = ScheduledBranchSet::new(optimization.max_parallel_runtime_tasks)?;
7470 let main_branch = RuntimeBranch::new(
7471 RuntimeTaskPurpose::MainResponse,
7472 main_kind,
7473 RuntimeTaskPriority::Normal,
7474 RuntimeCommitBehavior::FinalResponse,
7475 );
7476 let transition_branch = RuntimeBranch::new(
7477 RuntimeTaskPurpose::StateTransition,
7478 RuntimeOptimizationKind::ParallelStateTransition,
7479 RuntimeTaskPriority::Critical,
7480 RuntimeCommitBehavior::TransitionDecision,
7481 );
7482 let skill_branch = RuntimeBranch::new(
7483 RuntimeTaskPurpose::SkillRouting,
7484 RuntimeOptimizationKind::SpeculativeSkillRouting,
7485 RuntimeTaskPriority::High,
7486 RuntimeCommitBehavior::SkillSelection,
7487 );
7488 let reasoning_branch = RuntimeBranch::new(
7489 RuntimeTaskPurpose::ReasoningJudge,
7490 RuntimeOptimizationKind::SpeculativeReasoningAuto,
7491 RuntimeTaskPriority::Normal,
7492 RuntimeCommitBehavior::ReasoningDecision,
7493 );
7494 let main_id = main_branch.branch_id();
7495 let transition_id = transition_branch.branch_id();
7496 let skill_id = skill_branch.branch_id();
7497 let reasoning_id = reasoning_branch.branch_id();
7498
7499 let main_id_for_future = main_id.clone();
7500 if !branch_set.schedule(
7501 main_branch,
7502 Box::pin(async move {
7503 match crate::optimization::observability::with_branch_observation(
7504 &main_id_for_future,
7505 main_kind,
7506 RuntimeCommitBehavior::FinalResponse,
7507 self.generate_main_response_draft(processed_input, &ReasoningMode::None),
7508 )
7509 .await
7510 {
7511 Ok(draft) => RuntimeBranchResult::MainDraft(draft),
7512 Err(error) => RuntimeBranchResult::Failed(error),
7513 }
7514 }),
7515 ) {
7516 return Ok(None);
7517 }
7518
7519 if transition_enabled {
7520 let transition_id_for_future = transition_id.clone();
7521 if !branch_set.schedule(
7522 transition_branch,
7523 Box::pin(async move {
7524 match crate::optimization::observability::with_branch_observation(
7525 &transition_id_for_future,
7526 RuntimeOptimizationKind::ParallelStateTransition,
7527 RuntimeCommitBehavior::TransitionDecision,
7528 self.select_parallel_transition_candidate(processed_input),
7529 )
7530 .await
7531 {
7532 Ok(ParallelTransitionSelection::Candidate(candidate)) => {
7533 RuntimeBranchResult::Transition(Some(candidate))
7534 }
7535 Ok(ParallelTransitionSelection::NoMatch) => {
7536 RuntimeBranchResult::Transition(None)
7537 }
7538 Ok(ParallelTransitionSelection::ReservationExhausted) => {
7539 RuntimeBranchResult::Cancelled
7540 }
7541 Err(error) => RuntimeBranchResult::Failed(error),
7542 }
7543 }),
7544 ) {
7545 transition_enabled = false;
7546 }
7547 }
7548
7549 if skill_enabled {
7550 let skill_id_for_future = skill_id.clone();
7551 if !branch_set.schedule(
7552 skill_branch,
7553 Box::pin(async move {
7554 if !self.reserve_active_speculative_llm_call(
7555 RuntimeOptimizationKind::SpeculativeSkillRouting,
7556 ) {
7557 return RuntimeBranchResult::Cancelled;
7558 }
7559 match crate::optimization::observability::with_branch_observation(
7560 &skill_id_for_future,
7561 RuntimeOptimizationKind::SpeculativeSkillRouting,
7562 RuntimeCommitBehavior::SkillSelection,
7563 self.select_skill_candidate(processed_input),
7564 )
7565 .await
7566 {
7567 Ok(candidate) => RuntimeBranchResult::Skill(candidate),
7568 Err(error) => RuntimeBranchResult::Failed(error),
7569 }
7570 }),
7571 ) {
7572 skill_enabled = false;
7573 }
7574 }
7575
7576 if reasoning_enabled {
7577 let reasoning_id_for_future = reasoning_id.clone();
7578 if !branch_set.schedule(
7579 reasoning_branch,
7580 Box::pin(async move {
7581 if !self.reserve_active_speculative_llm_call(
7582 RuntimeOptimizationKind::SpeculativeReasoningAuto,
7583 ) {
7584 return RuntimeBranchResult::Cancelled;
7585 }
7586 match crate::optimization::observability::with_branch_observation(
7587 &reasoning_id_for_future,
7588 RuntimeOptimizationKind::SpeculativeReasoningAuto,
7589 RuntimeCommitBehavior::ReasoningDecision,
7590 self.determine_reasoning_mode_strict(processed_input),
7591 )
7592 .await
7593 {
7594 Ok(mode) => RuntimeBranchResult::Reasoning(mode),
7595 Err(error) => RuntimeBranchResult::Failed(error),
7596 }
7597 }),
7598 ) {
7599 reasoning_enabled = false;
7600 }
7601 }
7602
7603 if matches!(effective_reasoning_mode, ReasoningMode::Auto) && !reasoning_enabled {
7604 self.finalize_pending_branches(branch_set.cancel_pending());
7605 return Ok(None);
7606 }
7607
7608 if !transition_enabled && !skill_enabled && !reasoning_enabled {
7609 self.finalize_pending_branches(branch_set.cancel_pending());
7610 return Ok(None);
7611 }
7612
7613 let mut main_pending = true;
7614 let mut skill_pending = skill_enabled;
7615 let mut reasoning_pending = reasoning_enabled;
7616 let mut transition_finalized = !transition_enabled;
7617 let mut skill_finalized = !skill_enabled && self.skill_router.is_none();
7620 let mut reasoning_finalized = !reasoning_enabled;
7621 let mut main_result: Option<Result<MainResponseDraft>> = None;
7622 let mut transition_candidate: Option<TransitionCandidate> = None;
7623 let mut skill_candidate: Option<SkillCandidate> = None;
7624 let mut reasoning_decision: Option<ReasoningMode> = None;
7625 let mut transition_fallback_required = false;
7626 let mut skill_fallback_required = false;
7627 let mut reasoning_fallback_required = false;
7628
7629 loop {
7630 if let Some(candidate) = transition_candidate.take() {
7631 if self
7632 .approve_transition_target(&candidate.from_state, candidate.target())
7633 .await?
7634 {
7635 self.finalize_pending_branches(branch_set.cancel_pending());
7637 if !main_pending {
7638 self.finalize_branch_loss(
7639 &main_id,
7640 main_kind,
7641 RuntimeCommitBehavior::FinalResponse,
7642 false,
7643 main_result.as_ref().map(|result| result.is_err()),
7644 );
7645 }
7646 if skill_enabled && !skill_pending {
7647 self.finalize_branch_loss(
7648 &skill_id,
7649 RuntimeOptimizationKind::SpeculativeSkillRouting,
7650 RuntimeCommitBehavior::SkillSelection,
7651 false,
7652 Some(false),
7653 );
7654 }
7655 if reasoning_enabled && !reasoning_pending {
7656 self.finalize_branch_loss(
7657 &reasoning_id,
7658 RuntimeOptimizationKind::SpeculativeReasoningAuto,
7659 RuntimeCommitBehavior::ReasoningDecision,
7660 false,
7661 Some(false),
7662 );
7663 }
7664 if !self
7665 .apply_pre_response_transition_candidate(
7666 &candidate,
7667 &HashMap::new(),
7668 processed_input,
7669 )
7670 .await?
7671 {
7672 self.finalize_optional_branch(
7673 &transition_id,
7674 RuntimeOptimizationKind::ParallelStateTransition,
7675 RuntimeCommitBehavior::TransitionDecision,
7676 "discarded",
7677 false,
7678 );
7679 return Ok(None);
7680 }
7681 self.finalize_optional_branch(
7682 &transition_id,
7683 RuntimeOptimizationKind::ParallelStateTransition,
7684 RuntimeCommitBehavior::TransitionDecision,
7685 "committed",
7686 true,
7687 );
7688 return self
7689 .redispatch_current_state(processed_input)
7690 .await
7691 .map(Some);
7692 }
7693 self.finalize_optional_branch(
7694 &transition_id,
7695 RuntimeOptimizationKind::ParallelStateTransition,
7696 RuntimeCommitBehavior::TransitionDecision,
7697 "discarded",
7698 false,
7699 );
7700 transition_finalized = true;
7701 }
7702
7703 if transition_finalized
7712 && !skill_finalized
7713 && !skill_enabled
7714 && self.skill_router.is_some()
7715 {
7716 match self.select_skill_candidate(processed_input).await {
7717 Ok(Some(candidate)) => skill_candidate = Some(candidate),
7718 Ok(None) => {}
7719 Err(error) => {
7720 self.finalize_pending_branches(branch_set.cancel_pending());
7722 return Err(error);
7723 }
7724 }
7725 skill_finalized = true;
7726 }
7727
7728 if transition_finalized && skill_candidate.is_some() {
7729 let candidate = skill_candidate.take().unwrap();
7730 if skill_enabled {
7732 self.finalize_optional_branch(
7733 &skill_id,
7734 RuntimeOptimizationKind::SpeculativeSkillRouting,
7735 RuntimeCommitBehavior::SkillSelection,
7736 "committed",
7737 true,
7738 );
7739 }
7740 if !main_pending {
7741 self.finalize_branch_loss(
7742 &main_id,
7743 main_kind,
7744 RuntimeCommitBehavior::FinalResponse,
7745 false,
7746 main_result.as_ref().map(|result| result.is_err()),
7747 );
7748 }
7749 if reasoning_enabled && !reasoning_pending {
7750 self.finalize_branch_loss(
7751 &reasoning_id,
7752 RuntimeOptimizationKind::SpeculativeReasoningAuto,
7753 RuntimeCommitBehavior::ReasoningDecision,
7754 false,
7755 Some(false),
7756 );
7757 }
7758 self.finalize_pending_branches(branch_set.cancel_pending());
7759 return self
7760 .commit_winning_skill_candidate(candidate, processed_input, input_context)
7761 .await;
7762 }
7763
7764 if transition_finalized
7765 && skill_finalized
7766 && let Some(reasoning_mode) = reasoning_decision.take()
7767 {
7768 if !matches!(reasoning_mode, ReasoningMode::None) {
7769 self.finalize_optional_branch(
7770 &reasoning_id,
7771 RuntimeOptimizationKind::SpeculativeReasoningAuto,
7772 RuntimeCommitBehavior::ReasoningDecision,
7773 "committed",
7774 true,
7775 );
7776 if !main_pending {
7777 self.finalize_branch_loss(
7778 &main_id,
7779 main_kind,
7780 RuntimeCommitBehavior::FinalResponse,
7781 false,
7782 main_result.as_ref().map(|result| result.is_err()),
7783 );
7784 }
7785 self.finalize_pending_branches(branch_set.cancel_pending());
7786 self.commit_root_user_message(processed_input).await?;
7787 return if matches!(reasoning_mode, ReasoningMode::PlanAndExecute) {
7788 self.handle_plan_and_execute(processed_input, input_context, true)
7789 .await
7790 .map(Some)
7791 } else {
7792 self.run_committed_response_loop_with_reasoning(
7793 processed_input,
7794 input_context,
7795 reasoning_mode,
7796 true,
7797 )
7798 .await
7799 .map(Some)
7800 };
7801 }
7802 self.finalize_optional_branch(
7803 &reasoning_id,
7804 RuntimeOptimizationKind::SpeculativeReasoningAuto,
7805 RuntimeCommitBehavior::ReasoningDecision,
7806 "committed",
7807 true,
7808 );
7809 reasoning_finalized = true;
7810 }
7811
7812 if transition_finalized && skill_finalized && reasoning_finalized {
7813 if transition_fallback_required
7814 || skill_fallback_required
7815 || reasoning_fallback_required
7816 {
7817 if !main_pending {
7818 self.finalize_branch_loss(
7819 &main_id,
7820 main_kind,
7821 RuntimeCommitBehavior::FinalResponse,
7822 false,
7823 main_result.as_ref().map(|result| result.is_err()),
7824 );
7825 }
7826 self.finalize_pending_branches(branch_set.cancel_pending());
7827 return Ok(None);
7828 }
7829
7830 if let Some(result) = main_result.take() {
7831 let draft = match result {
7832 Ok(draft) => draft,
7833 Err(error) => {
7834 self.finalize_optional_branch(
7835 &main_id,
7836 main_kind,
7837 RuntimeCommitBehavior::FinalResponse,
7838 "failed",
7839 false,
7840 );
7841 self.finalize_pending_branches(branch_set.cancel_pending());
7842 return Err(error);
7843 }
7844 };
7845 self.finalize_optional_branch(
7846 &main_id,
7847 main_kind,
7848 RuntimeCommitBehavior::FinalResponse,
7849 "committed",
7850 true,
7851 );
7852 self.finalize_pending_branches(branch_set.cancel_pending());
7853 return self
7854 .commit_main_response_draft(
7855 processed_input,
7856 input_context,
7857 draft,
7858 ReasoningMode::None,
7859 reasoning_enabled,
7860 )
7861 .await
7862 .map(Some);
7863 }
7864 }
7865
7866 if branch_set.is_empty() {
7867 return Ok(None);
7868 }
7869
7870 let Some(outcome) = branch_set.next_completed().await else {
7871 return Ok(None);
7872 };
7873 let branch_id = outcome.branch.branch_id();
7874 match outcome.result {
7875 RuntimeBranchResult::MainDraft(draft) => {
7876 main_pending = false;
7877 main_result = Some(Ok(draft));
7878 }
7879 RuntimeBranchResult::Transition(candidate) => {
7880 if let Some(candidate) = candidate {
7881 transition_candidate = Some(candidate);
7882 } else {
7883 self.finalize_optional_branch(
7884 &transition_id,
7885 RuntimeOptimizationKind::ParallelStateTransition,
7886 RuntimeCommitBehavior::TransitionDecision,
7887 "discarded",
7888 false,
7889 );
7890 transition_finalized = true;
7891 }
7892 }
7893 RuntimeBranchResult::Skill(candidate) => {
7894 skill_pending = false;
7895 if let Some(candidate) = candidate {
7896 skill_candidate = Some(candidate);
7897 } else {
7898 self.finalize_optional_branch(
7899 &skill_id,
7900 RuntimeOptimizationKind::SpeculativeSkillRouting,
7901 RuntimeCommitBehavior::SkillSelection,
7902 "discarded",
7903 false,
7904 );
7905 skill_finalized = true;
7906 }
7907 }
7908 RuntimeBranchResult::Reasoning(mode) => {
7909 reasoning_pending = false;
7910 reasoning_decision = Some(mode);
7911 }
7912 RuntimeBranchResult::Failed(error) => {
7913 if branch_id == main_id {
7914 main_pending = false;
7915 main_result = Some(Err(error));
7916 } else if branch_id == transition_id {
7917 self.finalize_optional_branch(
7918 &transition_id,
7919 RuntimeOptimizationKind::ParallelStateTransition,
7920 RuntimeCommitBehavior::TransitionDecision,
7921 "failed",
7922 false,
7923 );
7924 transition_finalized = true;
7925 } else if branch_id == skill_id {
7926 skill_pending = false;
7927 self.finalize_optional_branch(
7928 &skill_id,
7929 RuntimeOptimizationKind::SpeculativeSkillRouting,
7930 RuntimeCommitBehavior::SkillSelection,
7931 "failed",
7932 false,
7933 );
7934 skill_finalized = true;
7935 } else if branch_id == reasoning_id {
7936 reasoning_pending = false;
7937 self.finalize_optional_branch(
7938 &reasoning_id,
7939 RuntimeOptimizationKind::SpeculativeReasoningAuto,
7940 RuntimeCommitBehavior::ReasoningDecision,
7941 "failed",
7942 false,
7943 );
7944 reasoning_finalized = true;
7945 }
7946 }
7947 RuntimeBranchResult::Cancelled => {
7948 self.finalize_optional_branch(
7949 &branch_id,
7950 outcome.branch.optimization,
7951 outcome.branch.commit_behavior,
7952 "cancelled",
7953 false,
7954 );
7955 if branch_id == main_id {
7956 main_pending = false;
7957 main_result =
7958 Some(Err(AgentError::Other("main branch cancelled".to_string())));
7959 } else if branch_id == transition_id {
7960 transition_finalized = true;
7961 transition_fallback_required = true;
7962 } else if branch_id == skill_id {
7963 skill_pending = false;
7964 skill_finalized = true;
7965 skill_fallback_required = true;
7966 } else if branch_id == reasoning_id {
7967 reasoning_pending = false;
7968 reasoning_finalized = true;
7969 reasoning_fallback_required = true;
7970 }
7971 }
7972 }
7973 }
7974 }
7975
7976 fn finalize_pending_branches(&self, branches: Vec<RuntimeBranch>) {
7977 for branch in branches {
7978 self.finalize_optional_branch(
7979 &branch.branch_id(),
7980 branch.optimization,
7981 branch.commit_behavior,
7982 "cancelled",
7983 false,
7984 );
7985 }
7986 }
7987
7988 fn finalize_branch_loss(
7993 &self,
7994 branch_id: &str,
7995 optimization: RuntimeOptimizationKind,
7996 commit_behavior: RuntimeCommitBehavior,
7997 pending: bool,
7998 completed_failed: Option<bool>,
7999 ) {
8000 let status = if pending {
8001 "cancelled"
8002 } else if completed_failed.unwrap_or(false) {
8003 "failed"
8004 } else {
8005 "discarded"
8006 };
8007 self.finalize_optional_branch(branch_id, optimization, commit_behavior, status, false);
8008 }
8009
8010 fn finalize_optional_branch(
8015 &self,
8016 branch_id: &str,
8017 optimization: RuntimeOptimizationKind,
8018 commit_behavior: RuntimeCommitBehavior,
8019 status: &str,
8020 winner: bool,
8021 ) {
8022 crate::optimization::observability::finalize_branch(
8023 self.observability_manager.as_ref(),
8024 branch_id,
8025 status,
8026 winner,
8027 optimization,
8028 commit_behavior,
8029 );
8030 }
8031
8032 fn has_parallel_transition_candidates(&self) -> bool {
8037 self.transitions_available_for_commit()
8038 .map(|(transitions, _)| {
8039 transitions
8040 .iter()
8041 .any(|transition| matches!(transition.timing, TransitionTiming::Parallel))
8042 })
8043 .unwrap_or(false)
8044 }
8045
8046 async fn select_parallel_transition_candidate(
8051 &self,
8052 processed_input: &str,
8053 ) -> Result<ParallelTransitionSelection> {
8054 let Some((transitions, current_state)) = self.transitions_available_for_commit() else {
8055 return Ok(ParallelTransitionSelection::NoMatch);
8056 };
8057 let parallel: Vec<Transition> = transitions
8058 .into_iter()
8059 .filter(|transition| matches!(transition.timing, TransitionTiming::Parallel))
8060 .filter(|transition| !transition.requires_response)
8061 .collect();
8062 if parallel.is_empty() {
8063 return Ok(ParallelTransitionSelection::NoMatch);
8064 }
8065 let empty_staged = HashMap::new();
8066 if let Some(candidate) = self.select_deterministic_transition_candidate(
8067 processed_input,
8068 ¤t_state,
8069 ¶llel,
8070 &empty_staged,
8071 ) {
8072 return Ok(ParallelTransitionSelection::Candidate(candidate));
8073 }
8074 let when_transitions: Vec<(usize, &Transition)> = parallel
8075 .iter()
8076 .enumerate()
8077 .filter(|(_, transition)| !transition.when.trim().is_empty())
8078 .collect();
8079 if when_transitions.is_empty() {
8080 return Ok(ParallelTransitionSelection::NoMatch);
8081 }
8082 let llm = self
8083 .llm_registry
8084 .router()
8085 .or_else(|_| self.llm_registry.default())
8086 .map_err(|e| AgentError::Config(e.to_string()))?;
8087 let conditions = when_transitions
8088 .iter()
8089 .enumerate()
8090 .map(|(display_idx, (_, transition))| {
8091 format!("{}. {}", display_idx + 1, transition.when)
8092 })
8093 .collect::<Vec<_>>()
8094 .join("\n");
8095 if !self
8096 .reserve_active_speculative_llm_call(RuntimeOptimizationKind::ParallelStateTransition)
8097 {
8098 return Ok(ParallelTransitionSelection::ReservationExhausted);
8099 }
8100 let context_preview = self.branch_context_preview();
8101 let prompt = format!(
8102 "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-{}).",
8103 current_state,
8104 processed_input,
8105 context_preview,
8106 conditions,
8107 when_transitions.len()
8108 );
8109 let response = self
8110 .observe_purpose(
8111 ObservationPurpose::StateTransitionEvaluation,
8112 llm.complete(&[ChatMessage::user(prompt)], None),
8113 )
8114 .await
8115 .map_err(|e| AgentError::LLM(e.to_string()))?;
8116 let choice = response.content.trim().parse::<usize>().unwrap_or(0);
8117 if choice == 0 || choice > when_transitions.len() {
8118 return Ok(ParallelTransitionSelection::NoMatch);
8119 }
8120 let transition = when_transitions[choice - 1].1.clone();
8121 Ok(ParallelTransitionSelection::Candidate(
8122 TransitionCandidate::new(
8123 current_state,
8124 transition.clone(),
8125 Self::transition_reason(&transition),
8126 ),
8127 ))
8128 }
8129
8130 async fn redispatch_current_state(&self, processed_input: &str) -> Result<AgentResponse> {
8132 const MAX_REDISPATCH_DEPTH: u32 = 3;
8133 let current_depth = *self.redispatch_depth.read();
8134 if current_depth >= MAX_REDISPATCH_DEPTH {
8135 warn!(depth = current_depth, "Re-dispatch depth limit reached");
8136 let response = AgentResponse::new("");
8137 self.finish_turn_if_root(&response).await?;
8138 return Ok(response);
8139 }
8140 *self.redispatch_depth.write() += 1;
8141 if let Some(context) = self.active_turn_context.write().as_mut() {
8142 context.enter_redispatch();
8143 }
8144 let result = Box::pin(self.run_loop_internal(processed_input)).await;
8145 *self.redispatch_depth.write() -= 1;
8146 if let Some(context) = self.active_turn_context.write().as_mut() {
8147 context.exit_redispatch();
8148 }
8149 let response = result?;
8150 self.finish_turn_if_root(&response).await?;
8151 Ok(response)
8152 }
8153
8154 async fn finish_turn_if_root(&self, response: &AgentResponse) -> Result<()> {
8156 if *self.redispatch_depth.read() == 0 {
8157 self.post_turn_session_lifecycle().await?;
8158 if let Some(context) = self.active_turn_context.write().as_mut() {
8159 context.mark_post_turn_lifecycle_completed();
8160 }
8161 self.hooks.on_response(response).await;
8162 self.end_root_turn();
8163 }
8164 Ok(())
8165 }
8166
8167 async fn execute_state_exit_actions(&self, state_path: &str) {
8169 if let Some(ref sm) = self.state_machine
8170 && let Some(def) = sm.get_definition(state_path)
8171 && !def.on_exit.is_empty()
8172 {
8173 debug!(state = %state_path, count = def.on_exit.len(), "Executing on_exit actions");
8174 self.execute_state_actions(&def.on_exit).await;
8175 }
8176 }
8177
8178 fn state_was_previously_entered(
8180 state_path: &str,
8181 from_state: &str,
8182 history_before: &[StateTransitionEvent],
8183 ) -> bool {
8184 state_path == from_state
8185 || history_before
8186 .iter()
8187 .any(|event| event.from == state_path || event.to == state_path)
8188 }
8189
8190 async fn execute_state_enter_actions(&self, state_path: &str, is_reentry: bool) {
8192 if let Some(ref sm) = self.state_machine
8193 && let Some(def) = sm.get_definition(state_path)
8194 {
8195 if is_reentry && !def.on_reenter.is_empty() {
8196 debug!(state = %state_path, count = def.on_reenter.len(), "Executing on_reenter actions");
8197 self.execute_state_actions(&def.on_reenter).await;
8198 } else if !def.on_enter.is_empty() {
8199 debug!(state = %state_path, count = def.on_enter.len(), "Executing on_enter actions");
8200 self.execute_state_actions(&def.on_enter).await;
8201 }
8202 }
8203 }
8204
8205 async fn execute_state_actions(&self, actions: &[StateAction]) {
8207 for (action_index, action) in actions.iter().enumerate() {
8208 match action {
8209 StateAction::Tool { tool, args } => {
8210 let raw_args = args.clone().unwrap_or(Value::Object(Default::default()));
8211 let args_value = self.render_action_args(&raw_args);
8212 let state = self.state_machine.as_ref().map(|sm| sm.current());
8213 let request = ToolExecutionRequest::new(
8214 uuid::Uuid::new_v4().to_string(),
8215 tool.clone(),
8216 args_value,
8217 ToolCallSource::StateAction {
8218 state,
8219 action_index,
8220 },
8221 );
8222 match self.execute_tool_record(request).await {
8223 Ok(record) if record.success => {
8224 debug!(tool = %record.canonical_id, "State action: tool executed");
8225 let _ = self.context_manager.set(
8226 "last_tool_result",
8227 serde_json::Value::String(record.model_output_string()),
8228 );
8229 let _ = self.context_manager.set(
8230 "last_tool_record",
8231 serde_json::to_value(record).unwrap_or(Value::Null),
8232 );
8233 }
8234 Ok(record) => {
8235 warn!(tool = %record.canonical_id, error = %record.output, "State action: tool failed");
8236 }
8237 Err(e) => {
8238 warn!(tool = %tool, error = %e, "State action: tool failed")
8239 }
8240 }
8241 }
8242 StateAction::Skill { skill } => {
8243 if let Some(ref executor) = self.skill_executor {
8244 if let Some(def) = self.skills.iter().find(|s| s.id == *skill) {
8245 match executor
8246 .execute_with_invoker(def, "", serde_json::json!({}), self)
8247 .await
8248 {
8249 Ok(_) => debug!(skill = %skill, "State action: skill executed"),
8250 Err(e) => {
8251 warn!(skill = %skill, error = %e, "State action: skill failed")
8252 }
8253 }
8254 } else {
8255 warn!(skill = %skill, "State action: skill not found");
8256 }
8257 }
8258 }
8259 StateAction::SetContext { set_context } => {
8260 for (key, value) in set_context {
8261 if let Err(e) = self.context_manager.set(key, value.clone()) {
8262 warn!(key = %key, error = %e, "State action: set_context failed");
8263 } else {
8264 debug!(key = %key, "State action: context set");
8265 }
8266 }
8267 }
8268 StateAction::Prompt {
8269 prompt,
8270 llm,
8271 store_as,
8272 } => {
8273 let llm_result = if let Some(alias) = llm {
8274 self.llm_registry.get(alias)
8275 } else {
8276 self.llm_registry.default()
8277 };
8278 match llm_result {
8279 Ok(llm_provider) => {
8280 let context = self.build_context_with_overlays();
8282 let rendered_prompt = self
8283 .template_renderer
8284 .render(prompt, &context)
8285 .unwrap_or_else(|_| prompt.clone());
8286 let recent =
8287 self.memory.get_messages(Some(5)).await.unwrap_or_default();
8288 let mut messages: Vec<ChatMessage> = recent;
8289 messages.push(ChatMessage::user(&rendered_prompt));
8290 match self
8291 .observe_purpose(
8292 ObservationPurpose::StateAction,
8293 llm_provider.complete(&messages, None),
8294 )
8295 .await
8296 {
8297 Ok(response) => {
8298 if let Some(key) = store_as {
8299 let _ = self
8300 .context_manager
8301 .set(key, Value::String(response.content));
8302 debug!(key = %key, "State action: prompt result stored");
8303 }
8304 }
8305 Err(e) => {
8306 warn!(error = %e, "State action: prompt LLM call failed");
8307 }
8308 }
8309 }
8310 Err(e) => {
8311 warn!(error = %e, "State action: LLM not found for prompt");
8312 }
8313 }
8314 }
8315 }
8316 }
8317 }
8318
8319 async fn run_context_extractors_staged(&self, user_message: &str) -> HashMap<String, Value> {
8320 let extractors = match &self.state_machine {
8321 Some(sm) => match sm.current_definition() {
8322 Some(def) if !def.extract.is_empty() => def.extract.clone(),
8323 _ => return HashMap::new(),
8324 },
8325 None => return HashMap::new(),
8326 };
8327
8328 let mut staged = HashMap::new();
8329 for extractor in &extractors {
8330 let prompt = if let Some(ref custom) = extractor.llm_extract {
8331 format!(
8332 "User message:\n\"{}\"\n\nInstruction:\n{}",
8333 user_message, custom
8334 )
8335 } else if let Some(ref desc) = extractor.description {
8336 format!(
8337 "From the following message, extract: {}\n\n\
8338 Message: \"{}\"\n\n\
8339 If the information is present, return ONLY the extracted value.\n\
8340 If NOT present, return exactly: __NONE__",
8341 desc, user_message
8342 )
8343 } else {
8344 continue;
8345 };
8346
8347 let llm = match self
8348 .llm_registry
8349 .get(&extractor.llm)
8350 .or_else(|_| self.llm_registry.get("router"))
8351 .or_else(|_| self.llm_registry.get("default"))
8352 {
8353 Ok(llm) => llm,
8354 Err(e) => {
8355 warn!(key = %extractor.key, error = %e, "Extractor LLM not found");
8356 continue;
8357 }
8358 };
8359
8360 let messages = vec![ChatMessage::user(&prompt)];
8361 match self
8362 .observe_purpose(
8363 ObservationPurpose::ContextExtraction,
8364 llm.complete(&messages, None),
8365 )
8366 .await
8367 {
8368 Ok(response) => {
8369 let value = response.content.trim().to_string();
8370 if value != "__NONE__" && !value.is_empty() {
8371 staged.insert(
8372 extractor.key.clone(),
8373 serde_json::Value::String(value.clone()),
8374 );
8375 debug!(key = %extractor.key, value = %value, "Context extracted");
8376 } else if extractor.required {
8377 warn!(key = %extractor.key, "Required extraction returned no value");
8378 }
8379 }
8380 Err(e) => {
8381 warn!(key = %extractor.key, error = %e, "Context extraction LLM call failed");
8382 }
8383 }
8384 }
8385 staged
8386 }
8387
8388 fn commit_staged_context_writes(&self, staged: &HashMap<String, Value>) {
8389 for (key, value) in staged {
8390 if let Err(error) = self.context_manager.update(key, value.clone()) {
8391 warn!(key = %key, error = %error, "staged context write failed");
8392 }
8393 }
8394 }
8395
8396 async fn run_context_extractors(&self, user_message: &str) {
8398 let staged = self.run_context_extractors_staged(user_message).await;
8399 self.commit_staged_context_writes(&staged);
8400 }
8401
8402 async fn check_memory_compression(&self) -> Result<()> {
8403 if self.memory.needs_compression() {
8404 let result = self.memory.compress(None).await?;
8405 if let CompressResult::Compressed {
8406 messages_summarized,
8407 new_summary_length,
8408 tokens_saved,
8409 } = result
8410 {
8411 let event = MemoryCompressEvent::new(
8412 messages_summarized,
8413 tokens_saved,
8414 new_summary_length as u32,
8415 );
8416 self.hooks.on_memory_compress(&event).await;
8417 debug!(
8418 messages = messages_summarized,
8419 tokens_saved = tokens_saved,
8420 "Memory compressed"
8421 );
8422 }
8423 }
8424
8425 self.handle_memory_overflow().await?;
8427 self.check_memory_budget().await;
8428
8429 Ok(())
8430 }
8431
8432 async fn check_memory_budget(&self) {
8433 let Some(ref budget) = self.memory_token_budget else {
8434 return;
8435 };
8436
8437 let context = match self.memory.get_context().await {
8438 Ok(ctx) => ctx,
8439 Err(_) => return,
8440 };
8441
8442 let used_tokens = context.estimated_tokens();
8444 if budget.is_over_warn_threshold(used_tokens) {
8445 let event = MemoryBudgetEvent::new("memory", used_tokens, budget.total);
8446 self.hooks.on_memory_budget_warning(&event).await;
8447 debug!(
8448 used = used_tokens,
8449 total = budget.total,
8450 percent = event.usage_percent,
8451 "Memory budget warning"
8452 );
8453 }
8454
8455 if let Some(ref summary) = context.summary {
8457 let summary_tokens = ai_agents_memory::estimate_tokens(summary);
8458 let summary_budget = budget.allocation.summary;
8459 if summary_budget > 0 {
8460 let warn_threshold =
8461 (summary_budget as f64 * budget.warn_at_percent as f64 / 100.0) as u32;
8462 if summary_tokens >= warn_threshold {
8463 let event = MemoryBudgetEvent::new("summary", summary_tokens, summary_budget);
8464 self.hooks.on_memory_budget_warning(&event).await;
8465 }
8466 }
8467 }
8468
8469 let recent_tokens: u32 = context
8471 .messages
8472 .iter()
8473 .map(ai_agents_memory::estimate_message_tokens)
8474 .sum();
8475 let recent_budget = budget.allocation.recent_messages;
8476 if recent_budget > 0 {
8477 let warn_threshold =
8478 (recent_budget as f64 * budget.warn_at_percent as f64 / 100.0) as u32;
8479 if recent_tokens >= warn_threshold {
8480 let event = MemoryBudgetEvent::new("recent_messages", recent_tokens, recent_budget);
8481 self.hooks.on_memory_budget_warning(&event).await;
8482 }
8483 }
8484
8485 let relationship_budget = budget.allocation.relationships;
8486 if relationship_budget > 0 {
8487 let relationship_tokens = self
8488 .relationship_memory_text()
8489 .map(|text| ai_agents_memory::estimate_tokens(&text))
8490 .unwrap_or(0);
8491 let warn_threshold =
8492 (relationship_budget as f64 * budget.warn_at_percent as f64 / 100.0) as u32;
8493 if relationship_tokens >= warn_threshold {
8494 let event = MemoryBudgetEvent::new(
8495 "relationships",
8496 relationship_tokens,
8497 relationship_budget,
8498 );
8499 self.hooks.on_memory_budget_warning(&event).await;
8500 }
8501 }
8502 }
8503
8504 async fn handle_memory_overflow(&self) -> Result<()> {
8505 let Some(ref budget) = self.memory_token_budget else {
8506 return Ok(());
8507 };
8508
8509 let context = self.memory.get_context().await?;
8510 let used_tokens = context.estimated_tokens();
8511
8512 if used_tokens <= budget.total {
8513 return Ok(());
8514 }
8515
8516 match budget.overflow_strategy {
8517 OverflowStrategy::TruncateOldest => {
8518 let tokens_to_free = used_tokens - budget.total;
8519 let messages_to_evict = self.calculate_eviction_count(tokens_to_free);
8520 if messages_to_evict > 0 {
8521 self.evict_messages(messages_to_evict, EvictionReason::TokenBudgetExceeded)
8522 .await?;
8523 }
8524 }
8525 OverflowStrategy::SummarizeMore => {
8526 let max_attempts = context.total_messages.max(1);
8527 for _ in 0..max_attempts {
8528 match self.memory.compress(None).await? {
8529 CompressResult::Compressed {
8530 messages_summarized,
8531 ..
8532 } if messages_summarized > 0 => {
8533 let context = self.memory.get_context().await?;
8534 if context.estimated_tokens() <= budget.total {
8535 return Ok(());
8536 }
8537 }
8538 _ => break,
8539 }
8540 }
8541 let context = self.memory.get_context().await?;
8542 let used_tokens = context.estimated_tokens();
8543 if used_tokens > budget.total {
8544 return Err(AgentError::MemoryBudgetExceeded {
8545 used: used_tokens,
8546 budget: budget.total,
8547 });
8548 }
8549 }
8550 OverflowStrategy::Error => {
8551 return Err(AgentError::MemoryBudgetExceeded {
8552 used: used_tokens,
8553 budget: budget.total,
8554 });
8555 }
8556 }
8557 Ok(())
8558 }
8559
8560 fn calculate_eviction_count(&self, tokens_to_free: u32) -> usize {
8561 ((tokens_to_free as f64 / 50.0).ceil() as usize).max(1)
8563 }
8564
8565 async fn evict_messages(&self, count: usize, reason: EvictionReason) -> Result<()> {
8566 let evicted = self.memory.evict_oldest(count).await?;
8567 if !evicted.is_empty() {
8568 let event = MemoryEvictEvent {
8569 reason,
8570 messages_evicted: evicted.len(),
8571 importance_scores: vec![],
8572 };
8573 self.hooks.on_memory_evict(&event).await;
8574 debug!(count = evicted.len(), "Messages evicted from memory");
8575 }
8576 Ok(())
8577 }
8578
8579 #[instrument(skip(self, input), fields(agent = %self.info.name))]
8580 async fn determine_reasoning_mode(&self, input: &str) -> Result<ReasoningMode> {
8581 match self.determine_reasoning_mode_strict(input).await {
8582 Ok(mode) => Ok(mode),
8583 Err(_) => Ok(ReasoningMode::None),
8584 }
8585 }
8586
8587 async fn determine_reasoning_mode_strict(&self, input: &str) -> Result<ReasoningMode> {
8588 let effective_config = self.get_effective_reasoning_config();
8589
8590 if !matches!(effective_config.mode, ReasoningMode::Auto) {
8591 return Ok(effective_config.mode.clone());
8592 }
8593
8594 let judge_llm = effective_config
8595 .judge_llm
8596 .as_ref()
8597 .and_then(|alias| self.llm_registry.get(alias).ok())
8598 .or_else(|| self.llm_registry.router().ok())
8599 .or_else(|| self.llm_registry.default().ok());
8600
8601 let Some(llm) = judge_llm else {
8602 return Ok(ReasoningMode::None);
8603 };
8604
8605 let prompt = format!(
8606 r#"Analyze this user request and determine the appropriate reasoning mode.
8607
8608User request: "{}"
8609
8610Choose ONE of these modes:
8611- none: Simple queries, greetings, direct answers (fastest)
8612- cot: Complex analysis, multi-step reasoning, math problems
8613- react: Tasks requiring multiple tool calls with observation
8614- plan_and_execute: Complex multi-step tasks requiring coordination
8615
8616Respond with ONLY the mode name (none, cot, react, or plan_and_execute)."#,
8617 input
8618 );
8619
8620 let messages = vec![ChatMessage::user(&prompt)];
8621 let response = self
8622 .observe_purpose(
8623 ObservationPurpose::ReflectionDecision,
8624 llm.complete(&messages, None),
8625 )
8626 .await
8627 .map_err(|e| AgentError::LLM(e.to_string()))?;
8628
8629 let mode_str = response.content.trim().to_lowercase();
8630 Ok(match mode_str.as_str() {
8631 "cot" => ReasoningMode::CoT,
8632 "react" => ReasoningMode::React,
8633 "plan_and_execute" => ReasoningMode::PlanAndExecute,
8634 _ => ReasoningMode::None,
8635 })
8636 }
8637
8638 fn build_cot_system_prompt(&self, base_prompt: &str) -> String {
8639 format!(
8640 "{}\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>",
8641 base_prompt
8642 )
8643 }
8644
8645 fn build_react_system_prompt(&self, base_prompt: &str) -> String {
8646 format!(
8647 "{}\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>",
8648 base_prompt
8649 )
8650 }
8651
8652 async fn generate_plan(&self, input: &str) -> Result<Plan> {
8654 let effective = self.get_effective_reasoning_config();
8655 let planning_config = effective.get_planning();
8656
8657 let planner_llm = planning_config
8658 .and_then(|c| c.planner_llm.as_ref())
8659 .and_then(|alias| self.llm_registry.get(alias).ok())
8660 .or_else(|| self.llm_registry.router().ok())
8661 .or_else(|| self.llm_registry.default().ok())
8662 .ok_or_else(|| AgentError::Config("No LLM available for planning".into()))?;
8663
8664 let mut available_tool_ids = self.get_available_tool_ids().await?;
8665 let mut available_skills: Vec<String> = self.skills.iter().map(|s| s.id.clone()).collect();
8666
8667 if let Some(config) = planning_config {
8669 if !config.available.tools.is_all() {
8670 available_tool_ids.retain(|t| config.available.tools.allows(t));
8671 }
8672 if !config.available.skills.is_all() {
8673 available_skills.retain(|s| config.available.skills.allows(s));
8674 }
8675 }
8676
8677 let tool_descriptions: Vec<String> = available_tool_ids
8680 .iter()
8681 .filter_map(|id| {
8682 self.tools.get(id).map(|tool| {
8683 let schema = tool.input_schema();
8684 let args_desc = schema
8685 .get("properties")
8686 .and_then(|p| serde_json::to_string(p).ok())
8687 .unwrap_or_else(|| "{}".to_string());
8688 format!(
8689 "- {} ({}): {}\n Arguments: {}",
8690 id,
8691 tool.name(),
8692 tool.description(),
8693 args_desc
8694 )
8695 })
8696 })
8697 .collect();
8698
8699 let tools_section = if tool_descriptions.is_empty() {
8700 "Available tools: none".to_string()
8701 } else {
8702 format!("Available tools:\n{}", tool_descriptions.join("\n"))
8703 };
8704
8705 let skills_section = if available_skills.is_empty() {
8706 "Available skills: none".to_string()
8707 } else {
8708 format!("Available skills: {}", available_skills.join(", "))
8709 };
8710
8711 let prompt = format!(
8712 r#"Create a step-by-step plan to accomplish this goal.
8713
8714Goal: "{}"
8715
8716{}
8717
8718{}
8719
8720Create a plan with clear steps. For each step, specify:
8721- description: What this step accomplishes
8722- action_type: "tool", "skill", "think", or "respond"
8723- action_target: The tool/skill id (if applicable)
8724- args: The arguments object matching the tool's schema (if action_type is "tool")
8725- dependencies: List of step IDs this depends on (empty if none)
8726
8727Respond in JSON format:
8728{{
8729 "steps": [
8730 {{"id": "step1", "description": "...", "action_type": "tool", "action_target": "tool_id", "args": {{"required_field": "value"}}, "dependencies": []}},
8731 {{"id": "step2", "description": "...", "action_type": "think", "action_target": "...", "dependencies": ["step1"]}}
8732 ]
8733}}"#,
8734 input, tools_section, skills_section,
8735 );
8736
8737 let messages = vec![ChatMessage::user(&prompt)];
8738 let response = self
8739 .observe_purpose(
8740 ObservationPurpose::PlanGeneration,
8741 planner_llm.complete(&messages, None),
8742 )
8743 .await
8744 .map_err(|e| AgentError::LLM(format!("Planning failed: {}", e)))?;
8745
8746 let mut plan = Plan::new(input);
8747
8748 if let Some(json_start) = response.content.find('{')
8749 && let Some(json_end) = response.content.rfind('}')
8750 {
8751 let json_str = &response.content[json_start..=json_end];
8752 if let Ok(parsed) = serde_json::from_str::<serde_json::Value>(json_str)
8753 && let Some(steps) = parsed.get("steps").and_then(|s| s.as_array())
8754 {
8755 for step_value in steps {
8756 let id = step_value
8757 .get("id")
8758 .and_then(|v| v.as_str())
8759 .unwrap_or("step");
8760 let desc = step_value
8761 .get("description")
8762 .and_then(|v| v.as_str())
8763 .unwrap_or("");
8764 let action_type = step_value
8765 .get("action_type")
8766 .and_then(|v| v.as_str())
8767 .unwrap_or("think");
8768 let action_target = step_value
8769 .get("action_target")
8770 .and_then(|v| v.as_str())
8771 .unwrap_or("");
8772 let args = step_value
8773 .get("args")
8774 .cloned()
8775 .unwrap_or(serde_json::json!({}));
8776 let deps: Vec<String> = step_value
8777 .get("dependencies")
8778 .and_then(|v| v.as_array())
8779 .map(|arr| {
8780 arr.iter()
8781 .filter_map(|v| v.as_str().map(String::from))
8782 .collect()
8783 })
8784 .unwrap_or_default();
8785
8786 let action = match action_type {
8787 "tool" => PlanAction::tool(action_target, args),
8788 "skill" => PlanAction::skill(action_target),
8789 "respond" => PlanAction::respond(action_target),
8790 _ => PlanAction::think(desc),
8791 };
8792
8793 let step = PlanStep::new(desc, action)
8794 .with_id(id)
8795 .with_dependencies(deps);
8796 plan.add_step(step);
8797 }
8798 }
8799 }
8800
8801 if plan.steps.is_empty() {
8802 plan.add_step(PlanStep::new(
8803 "Process the request",
8804 PlanAction::think(input),
8805 ));
8806 plan.add_step(PlanStep::new(
8807 "Provide response",
8808 PlanAction::respond("Answer based on analysis"),
8809 ));
8810 }
8811
8812 Ok(plan)
8813 }
8814
8815 async fn execute_plan(&self, plan: &mut Plan) -> Result<String> {
8816 let llm = self.get_state_llm()?;
8817 let mut results: HashMap<String, serde_json::Value> = HashMap::new();
8818 let effective = self.get_effective_reasoning_config();
8819 let max_steps = effective.get_planning().map(|c| c.max_steps).unwrap_or(10);
8820
8821 plan.status = PlanStatus::InProgress;
8822
8823 for step_idx in 0..plan.steps.len().min(max_steps as usize) {
8824 let step = &plan.steps[step_idx];
8825
8826 let deps_satisfied = step.dependencies.iter().all(|dep| {
8827 plan.steps
8828 .iter()
8829 .find(|s| &s.id == dep)
8830 .map(|s| s.status.is_completed())
8831 .unwrap_or(false)
8832 });
8833
8834 if !deps_satisfied {
8835 continue;
8836 }
8837
8838 plan.steps[step_idx].mark_running();
8839
8840 let result = match &plan.steps[step_idx].action {
8841 PlanAction::Tool { tool, args } => {
8842 let has_dep_results = plan.steps[step_idx]
8848 .dependencies
8849 .iter()
8850 .any(|dep| results.contains_key(dep));
8851
8852 let final_args = if has_dep_results {
8853 let dep_context: String = plan.steps[step_idx]
8854 .dependencies
8855 .iter()
8856 .filter_map(|dep| results.get(dep).map(|r| format!("{}: {}", dep, r)))
8857 .collect::<Vec<_>>()
8858 .join("\n");
8859
8860 let tool_schema = self
8861 .tools
8862 .get(tool)
8863 .map(|t| {
8864 let schema = t.input_schema();
8865 let props = schema
8866 .get("properties")
8867 .and_then(|p| serde_json::to_string(p).ok())
8868 .unwrap_or_else(|| "{}".to_string());
8869 format!(
8870 "{}: {}\nArguments schema: {}",
8871 t.id(),
8872 t.description(),
8873 props
8874 )
8875 })
8876 .unwrap_or_default();
8877
8878 let step_desc = &plan.steps[step_idx].description;
8879 let arg_prompt = format!(
8880 "Generate the JSON arguments for a tool call.\n\n\
8881 Tool: {}\n\n\
8882 Task: {}\n\n\
8883 Previous step results:\n{}\n\n\
8884 Planner's draft arguments: {}\n\n\
8885 Produce ONLY a valid JSON object with the correct argument values.\n\
8886 Use actual values from the previous step results, not template references.",
8887 tool_schema,
8888 step_desc,
8889 dep_context,
8890 serde_json::to_string(args).unwrap_or_default()
8891 );
8892 let messages = vec![ChatMessage::user(&arg_prompt)];
8893 match self
8894 .observe_purpose(
8895 ObservationPurpose::PlanStep,
8896 llm.complete(&messages, None),
8897 )
8898 .await
8899 {
8900 Ok(resp) => {
8901 let content = resp.content.trim();
8902 let json_start = content.find('{');
8904 let json_end = content.rfind('}');
8905 if let (Some(start), Some(end)) = (json_start, json_end) {
8906 serde_json::from_str(&content[start..=end])
8907 .unwrap_or_else(|_| args.clone())
8908 } else {
8909 args.clone()
8910 }
8911 }
8912 Err(_) => args.clone(),
8913 }
8914 } else {
8915 args.clone()
8916 };
8917
8918 let request = ToolExecutionRequest::new(
8919 uuid::Uuid::new_v4().to_string(),
8920 tool.clone(),
8921 final_args,
8922 ToolCallSource::Plan {
8923 step_index: step_idx,
8924 },
8925 );
8926 match self.execute_tool_record(request).await {
8927 Ok(record) if record.success => {
8928 serde_json::json!({ "output": record.model_output_string() })
8929 }
8930 Ok(record) => {
8931 plan.steps[step_idx].mark_failed(record.model_output_string());
8932 continue;
8933 }
8934 Err(e) => {
8935 plan.steps[step_idx].mark_failed(e.to_string());
8936 continue;
8937 }
8938 }
8939 }
8940 PlanAction::Skill { skill } => {
8941 if let Some(skill_def) = self.skills.iter().find(|s| &s.id == skill) {
8942 if let Some(ref executor) = self.skill_executor {
8943 match executor
8944 .execute_with_invoker(skill_def, "", serde_json::json!({}), self)
8945 .await
8946 {
8947 Ok(output) => serde_json::json!({ "output": output }),
8948 Err(e) => {
8949 plan.steps[step_idx].mark_failed(e.to_string());
8950 continue;
8951 }
8952 }
8953 } else {
8954 serde_json::json!({ "output": "Skill executor not available" })
8955 }
8956 } else {
8957 plan.steps[step_idx].mark_failed("Skill not found");
8958 continue;
8959 }
8960 }
8961 PlanAction::Think { prompt } => {
8962 let context: String = results
8963 .iter()
8964 .map(|(k, v)| format!("{}: {}", k, v))
8965 .collect::<Vec<_>>()
8966 .join("\n");
8967
8968 let think_prompt = format!("Context:\n{}\n\nTask: {}", context, prompt);
8969 let messages = vec![ChatMessage::user(&think_prompt)];
8970
8971 match self
8972 .observe_purpose(
8973 ObservationPurpose::PlanStep,
8974 llm.complete(&messages, None),
8975 )
8976 .await
8977 {
8978 Ok(resp) => serde_json::json!({ "output": resp.content }),
8979 Err(e) => {
8980 plan.steps[step_idx].mark_failed(e.to_string());
8981 continue;
8982 }
8983 }
8984 }
8985 PlanAction::Respond { template } => {
8986 let context: String = results
8987 .iter()
8988 .map(|(k, v)| format!("{}: {}", k, v))
8989 .collect::<Vec<_>>()
8990 .join("\n");
8991
8992 let respond_prompt = format!(
8993 "Based on this context:\n{}\n\nGenerate a response following this template/instruction: {}",
8994 context, template
8995 );
8996 let messages = vec![ChatMessage::user(&respond_prompt)];
8997
8998 match self
8999 .observe_purpose(
9000 ObservationPurpose::PlanStep,
9001 llm.complete(&messages, None),
9002 )
9003 .await
9004 {
9005 Ok(resp) => serde_json::json!({ "output": resp.content }),
9006 Err(e) => {
9007 plan.steps[step_idx].mark_failed(e.to_string());
9008 continue;
9009 }
9010 }
9011 }
9012 };
9013
9014 results.insert(plan.steps[step_idx].id.clone(), result.clone());
9015 plan.steps[step_idx].mark_completed(Some(result));
9016 }
9017
9018 let has_failures = plan.steps.iter().any(|s| s.status.is_failed());
9020 if has_failures {
9021 let failed_ids: Vec<String> = plan
9022 .steps
9023 .iter()
9024 .filter(|s| s.status.is_failed())
9025 .map(|s| s.id.clone())
9026 .collect();
9027 plan.status = PlanStatus::Failed {
9028 error: format!("Steps failed: {}", failed_ids.join(", ")),
9029 };
9030 } else {
9031 plan.status = PlanStatus::Completed;
9032 }
9033
9034 let all_outputs: Vec<String> = plan
9036 .steps
9037 .iter()
9038 .filter(|s| s.status.is_completed())
9039 .filter_map(|s| {
9040 s.result
9041 .as_ref()
9042 .and_then(|r| r.get("output"))
9043 .and_then(|o| o.as_str())
9044 .map(|o| format!("{}: {}", s.description, o))
9045 })
9046 .collect();
9047
9048 if all_outputs.is_empty() {
9049 return Ok("Plan execution completed but produced no results.".to_string());
9050 }
9051
9052 if all_outputs.len() == 1 {
9053 return Ok(all_outputs.into_iter().next().unwrap());
9054 }
9055
9056 let context = all_outputs.join("\n\n");
9058 let prompt = format!(
9059 "You completed a multi-step plan for: \"{}\"\n\nStep results:\n{}\n\nProvide a coherent final response that synthesizes these results.",
9060 plan.goal, context
9061 );
9062 let messages = vec![ChatMessage::user(&prompt)];
9063 match self
9064 .observe_purpose(ObservationPurpose::PlanStep, llm.complete(&messages, None))
9065 .await
9066 {
9067 Ok(resp) => Ok(resp.content.trim().to_string()),
9068 Err(_) => Ok(context),
9069 }
9070 }
9071
9072 fn extract_thinking(&self, content: &str) -> (Option<String>, String) {
9073 if let Some(start) = content.find("<thinking>")
9074 && let Some(end) = content.find("</thinking>")
9075 {
9076 let thinking = content[start + 10..end].trim().to_string();
9077 let answer = content[end + 11..].trim().to_string();
9078 return (Some(thinking), answer);
9079 }
9080 (None, content.to_string())
9081 }
9082
9083 fn format_response_with_thinking(&self, thinking: Option<&str>, answer: &str) -> String {
9084 match self.get_effective_reasoning_config().output {
9085 ReasoningOutput::Hidden => answer.to_string(),
9086 ReasoningOutput::Visible => {
9087 if let Some(t) = thinking {
9088 format!("Thinking:\n{}\n\nAnswer:\n{}", t, answer)
9089 } else {
9090 answer.to_string()
9091 }
9092 }
9093 ReasoningOutput::Tagged => {
9094 if let Some(t) = thinking {
9095 format!("<thinking>{}</thinking>\n{}", t, answer)
9096 } else {
9097 answer.to_string()
9098 }
9099 }
9100 }
9101 }
9102
9103 fn disambiguation_question_response(
9106 question: &ClarificationQuestion,
9107 detection: &AmbiguityDetectionResult,
9108 awaiting_confirmation: bool,
9109 ) -> AgentResponse {
9110 let status = if awaiting_confirmation {
9111 "awaiting_confirmation"
9112 } else {
9113 "awaiting_clarification"
9114 };
9115 AgentResponse::new(&question.question).with_metadata(
9116 "disambiguation",
9117 serde_json::json!({
9118 "status": status,
9119 "options": question.options,
9120 "clarifying": question.clarifying,
9121 "detection": {
9122 "type": detection.ambiguity_type,
9123 "confidence": detection.confidence,
9124 "what_is_unclear": detection.what_is_unclear,
9125 }
9126 }),
9127 )
9128 }
9129
9130 async fn resolve_disambiguation(&self, input: &str) -> Result<DisambiguationDispatch> {
9142 let Some(ref disambiguator) = self.disambiguation_manager else {
9143 return Ok(DisambiguationDispatch::Proceed(input.to_string()));
9144 };
9145 let disambiguation_context = self.build_disambiguation_context().await?;
9146
9147 let state_override = self
9149 .state_machine
9150 .as_ref()
9151 .and_then(|sm| sm.current_definition())
9152 .and_then(|def| def.disambiguation.clone());
9153
9154 let state_generation = self
9155 .state_machine
9156 .as_ref()
9157 .map(|state_machine| state_machine.generation());
9158 let disambiguation_epoch = self.disambiguation_epoch.load(Ordering::SeqCst);
9159 let mut disambiguation_result = self
9160 .observe_purpose(
9161 ObservationPurpose::DisambiguationDetection,
9162 disambiguator.process_input_with_override(
9163 input,
9164 &disambiguation_context,
9165 state_override.as_ref(),
9166 None,
9167 ),
9168 )
9169 .await?;
9170 let current_state_generation = self
9171 .state_machine
9172 .as_ref()
9173 .map(|state_machine| state_machine.generation());
9174 if current_state_generation != state_generation
9175 || self.disambiguation_epoch.load(Ordering::SeqCst) != disambiguation_epoch
9176 {
9177 disambiguator.clear_pending().await;
9178 *self.pending_skill_id.write() = None;
9179 disambiguation_result = DisambiguationResult::Abandoned { new_input: None };
9180 info!(
9181 confirmation_event = "invalidated",
9182 invalidation_reason = "state_generation_changed",
9183 "Disambiguation result invalidated before redispatch"
9184 );
9185 }
9186 match disambiguation_result {
9187 DisambiguationResult::Clear => {
9188 debug!("Input is clear, proceeding normally");
9189 Ok(DisambiguationDispatch::Proceed(input.to_string()))
9190 }
9191 DisambiguationResult::NeedsClarification {
9192 question,
9193 detection,
9194 } => {
9195 let admission = match self
9196 .admit_disambiguation_redispatch(disambiguation_epoch, state_generation)
9197 .await
9198 {
9199 Ok(admission) => admission,
9200 Err(error) => {
9201 *self.pending_skill_id.write() = None;
9202 return Err(error);
9203 }
9204 };
9205 let awaiting_confirmation = disambiguator.has_pending_confirmation().await;
9206 info!(
9207 ambiguity_type = ?detection.ambiguity_type,
9208 confidence = detection.confidence,
9209 "Input requires clarification"
9210 );
9211
9212 self.commit_root_user_message(input).await?;
9215 self.memory
9216 .add_message(ChatMessage::assistant(&question.question))
9217 .await?;
9218
9219 let response = Self::disambiguation_question_response(
9220 &question,
9221 &detection,
9222 awaiting_confirmation,
9223 );
9224 drop(admission);
9225 self.finish_turn_if_root(&response).await?;
9226 Ok(DisambiguationDispatch::Terminal(response))
9227 }
9228 DisambiguationResult::Clarified {
9229 enriched_input,
9230 resolved,
9231 ..
9232 } => {
9233 let admission = match self
9234 .admit_disambiguation_redispatch(disambiguation_epoch, state_generation)
9235 .await
9236 {
9237 Ok(admission) => admission,
9238 Err(error) => {
9239 *self.pending_skill_id.write() = None;
9240 return Err(error);
9241 }
9242 };
9243 info!(
9244 resolved_count = resolved.len(),
9245 enriched = %enriched_input,
9246 "Input clarified, injecting resolved intent into context"
9247 );
9248
9249 for (key, value) in &resolved {
9252 let context_key = format!("disambiguation.{}", key);
9253 let _ = self.context_manager.set(&context_key, value.clone());
9254 }
9255
9256 if let Some(intent) = resolved.get("intent") {
9257 let _ = self.context_manager.set("resolved_intent", intent.clone());
9258 }
9259
9260 let _ = self
9261 .context_manager
9262 .set("disambiguation.resolved", serde_json::Value::Bool(true));
9263
9264 let skill_id = self.pending_skill_id.read().clone();
9268 drop(admission);
9269 if let Some(skill_id) = skill_id {
9270 info!(skill_id = %skill_id, "Re-checking skill disambiguation on clarified input");
9271 return Ok(DisambiguationDispatch::RecheckSkill {
9272 skill_id,
9273 enriched_input,
9274 disambiguation_epoch,
9275 state_generation,
9276 });
9277 }
9278 Ok(DisambiguationDispatch::Proceed(enriched_input))
9279 }
9280 DisambiguationResult::ProceedWithBestGuess { enriched_input } => {
9281 info!("Proceeding with best guess interpretation");
9282
9283 let skill_id = self.pending_skill_id.read().clone();
9285 if let Some(skill_id) = skill_id {
9286 info!(skill_id = %skill_id, "Re-checking skill disambiguation on best-guess input");
9287 return Ok(DisambiguationDispatch::RecheckSkill {
9288 skill_id,
9289 enriched_input,
9290 disambiguation_epoch,
9291 state_generation,
9292 });
9293 }
9294 Ok(DisambiguationDispatch::Proceed(enriched_input))
9295 }
9296 DisambiguationResult::GiveUp { reason } => {
9297 *self.pending_skill_id.write() = None;
9298 warn!(reason = %reason, "Disambiguation gave up");
9299 let apology = self
9300 .generate_localized_apology(
9301 "Generate a brief, polite apology saying you couldn't understand the request. Be concise.",
9302 &reason,
9303 )
9304 .await
9305 .unwrap_or_else(|_| {
9306 format!("I'm sorry, I couldn't understand your request: {}", reason)
9307 });
9308 let response = AgentResponse::new(&apology);
9309 self.finish_turn_if_root(&response).await?;
9310 Ok(DisambiguationDispatch::Terminal(response))
9311 }
9312 DisambiguationResult::Escalate { reason } => {
9313 *self.pending_skill_id.write() = None;
9314 info!(reason = %reason, "Escalating to human");
9315 if let Some(ref hitl) = self.hitl_engine {
9316 let trigger =
9317 ApprovalTrigger::condition("disambiguation_escalation", reason.clone());
9318 let mut context_map = HashMap::new();
9319 context_map.insert("original_input".to_string(), serde_json::json!(input));
9320 context_map.insert("reason".to_string(), serde_json::json!(&reason));
9321 let check_result = HITLCheckResult::required(
9322 trigger,
9323 context_map,
9324 format!("User request needs human assistance: {}", reason),
9325 Some(hitl.config().default_timeout_seconds),
9326 );
9327 let result = self.request_hitl_approval(check_result).await?;
9328 if matches!(
9329 result,
9330 ApprovalResult::Approved | ApprovalResult::Modified { .. }
9331 ) {
9332 return Ok(DisambiguationDispatch::Proceed(input.to_string()));
9334 }
9335 }
9336 let apology = self
9337 .generate_localized_apology(
9338 "Explain briefly that you're transferring the user to a human agent for help.",
9339 &reason,
9340 )
9341 .await
9342 .unwrap_or_else(|_| {
9343 format!("I need human assistance to help with your request: {}", reason)
9344 });
9345 let response = AgentResponse::new(&apology);
9346 self.finish_turn_if_root(&response).await?;
9347 Ok(DisambiguationDispatch::Terminal(response))
9348 }
9349 DisambiguationResult::Abandoned { new_input } => {
9350 *self.pending_skill_id.write() = None;
9351
9352 info!(
9353 has_new_input = new_input.is_some(),
9354 "Clarification abandoned by user"
9355 );
9356
9357 self.commit_root_user_message(input).await?;
9358
9359 match new_input {
9360 Some(fresh_input) => {
9361 Ok(DisambiguationDispatch::Proceed(fresh_input))
9364 }
9365 None => {
9366 let ack = self
9368 .generate_localized_apology(
9369 "The user changed their mind about their previous request. \
9370 Generate a brief, friendly acknowledgment (e.g. 'OK, no problem. What else can I help with?'). \
9371 Do NOT apologize excessively. Be concise.",
9372 "User abandoned clarification",
9373 )
9374 .await
9375 .unwrap_or_else(|_| {
9376 "OK, no problem. What else can I help with?".to_string()
9377 });
9378
9379 self.memory
9380 .add_message(ChatMessage::assistant(&ack))
9381 .await?;
9382
9383 let response = AgentResponse::new(&ack);
9384 self.finish_turn_if_root(&response).await?;
9385 Ok(DisambiguationDispatch::Terminal(response))
9386 }
9387 }
9388 }
9389 }
9390 }
9391
9392 async fn prepare_turn_context(&self) -> Result<()> {
9396 if !self.context_initialized.load(Ordering::SeqCst) {
9397 self.context_manager.initialize().await?;
9398 self.context_initialized.store(true, Ordering::SeqCst);
9399 debug!("Context manager initialized (defaults, env, builtins)");
9400 }
9401
9402 self.check_turn_timeout().await?;
9403 self.context_manager.refresh_per_turn().await?;
9404 self.context_manager.validate()
9405 }
9406
9407 async fn run_loop(&self, input: &str) -> Result<AgentResponse> {
9410 self.init_storage().await?;
9414 self.begin_root_turn();
9415 let _root_cleanup = RootTurnCleanup::new(self);
9416 info!(input_len = input.len(), "Starting chat");
9417
9418 self.hooks.on_message_received(input).await;
9419
9420 self.prepare_turn_context().await?;
9421
9422 self.clear_disambiguation_context();
9425
9426 let input_to_run = match self.resolve_disambiguation(input).await? {
9429 DisambiguationDispatch::Terminal(response) => return Ok(response),
9430 DisambiguationDispatch::RecheckSkill {
9431 skill_id,
9432 enriched_input,
9433 disambiguation_epoch,
9434 state_generation,
9435 } => {
9436 return self
9437 .recheck_skill_disambiguation(
9438 &skill_id,
9439 &enriched_input,
9440 disambiguation_epoch,
9441 state_generation,
9442 )
9443 .await;
9444 }
9445 DisambiguationDispatch::Proceed(input) => input,
9446 };
9447
9448 self.run_loop_internal(&input_to_run).await
9449 }
9450
9451 async fn generate_localized_apology(&self, instruction: &str, reason: &str) -> Result<String> {
9453 let llm = self.llm_registry.router().map_err(|e| {
9454 AgentError::LLM(format!(
9455 "Router LLM not available for localized response: {}",
9456 e
9457 ))
9458 })?;
9459
9460 let recent: Vec<String> = self
9461 .memory
9462 .get_messages(Some(3))
9463 .await?
9464 .iter()
9465 .map(|m| m.content.clone())
9466 .collect();
9467
9468 let context_hint = if recent.is_empty() {
9469 String::new()
9470 } else {
9471 format!(
9472 "\nRecent conversation (detect the user's language from this):\n{}\n",
9473 recent.join("\n")
9474 )
9475 };
9476
9477 let prompt = format!(
9478 "{}\nReason: {}\n{}Respond in the same language as the user. Output ONLY the message, nothing else.",
9479 instruction, reason, context_hint
9480 );
9481
9482 let messages = vec![ChatMessage::user(&prompt)];
9483 let response = self
9484 .observe_purpose(
9485 ObservationPurpose::DisambiguationClarification,
9486 llm.complete(&messages, None),
9487 )
9488 .await
9489 .map_err(|e| AgentError::LLM(format!("Localized response generation failed: {}", e)))?;
9490
9491 Ok(response.content.trim().to_string())
9492 }
9493
9494 fn render_action_args(&self, args: &Value) -> Value {
9498 let context = self.build_context_with_overlays();
9499 match args {
9500 Value::Object(map) => {
9501 let mut rendered = serde_json::Map::new();
9502 for (k, v) in map {
9503 match v {
9504 Value::String(s) if s.contains("{{") => {
9505 match self.template_renderer.render(s, &context) {
9506 Ok(rendered_str) => {
9507 rendered.insert(k.clone(), Value::String(rendered_str));
9508 }
9509 Err(_) => {
9510 rendered.insert(k.clone(), v.clone());
9511 }
9512 }
9513 }
9514 _ => {
9515 rendered.insert(k.clone(), v.clone());
9516 }
9517 }
9518 }
9519 Value::Object(rendered)
9520 }
9521 _ => args.clone(),
9522 }
9523 }
9524
9525 fn clear_disambiguation_context(&self) {
9527 let _ = self
9528 .context_manager
9529 .set("resolved_intent", serde_json::Value::Null);
9530
9531 let all = self.context_manager.get_all();
9532 for key in all.keys() {
9533 if key.starts_with("disambiguation.") {
9534 let _ = self.context_manager.set(key, serde_json::Value::Null);
9535 }
9536 }
9537 }
9538
9539 async fn recheck_skill_disambiguation(
9545 &self,
9546 skill_id: &str,
9547 enriched_input: &str,
9548 expected_disambiguation_epoch: u64,
9549 expected_state_generation: Option<u64>,
9550 ) -> Result<AgentResponse> {
9551 let skill = self
9552 .skill_router
9553 .as_ref()
9554 .and_then(|r| r.get_skill(skill_id).cloned());
9555
9556 if let Some(ref skill) = skill
9558 && let Some(ref skill_disambig) = skill.disambiguation
9559 && skill_disambig.enabled.unwrap_or(false)
9560 && let Some(ref disambiguator) = self.disambiguation_manager
9561 {
9562 let context = self.build_disambiguation_context().await?;
9563 let state_override = self
9564 .state_machine
9565 .as_ref()
9566 .and_then(|sm| sm.current_definition())
9567 .and_then(|def| def.disambiguation.clone());
9568
9569 let disambiguation_result = self
9570 .observe_purpose(
9571 ObservationPurpose::DisambiguationDetection,
9572 disambiguator.process_input_with_override(
9573 enriched_input,
9574 &context,
9575 state_override.as_ref(),
9576 Some(skill_disambig),
9577 ),
9578 )
9579 .await?;
9580 let current_state_generation = self
9581 .state_machine
9582 .as_ref()
9583 .map(|state_machine| state_machine.generation());
9584 if current_state_generation != expected_state_generation
9585 || self.disambiguation_epoch.load(Ordering::SeqCst) != expected_disambiguation_epoch
9586 {
9587 disambiguator.clear_pending().await;
9588 *self.pending_skill_id.write() = None;
9589 return Err(AgentError::Other(
9590 "State or reset ownership changed during skill disambiguation recheck"
9591 .to_string(),
9592 ));
9593 }
9594 match disambiguation_result {
9595 DisambiguationResult::Clear => {
9596 debug!(skill_id = %skill_id, "Skill re-check: all fields present");
9597 }
9598 DisambiguationResult::NeedsClarification {
9599 question,
9600 detection,
9601 } => {
9602 let admission = self
9603 .admit_disambiguation_redispatch(
9604 expected_disambiguation_epoch,
9605 expected_state_generation,
9606 )
9607 .await?;
9608 let awaiting_confirmation = disambiguator.has_pending_confirmation().await;
9609 info!(
9610 skill_id = %skill_id,
9611 ambiguity_type = ?detection.ambiguity_type,
9612 what_is_unclear = ?detection.what_is_unclear,
9613 "Skill re-check: still missing fields, asking again"
9614 );
9615 self.memory
9619 .add_message(ChatMessage::user(enriched_input))
9620 .await?;
9621 self.memory
9622 .add_message(ChatMessage::assistant(&question.question))
9623 .await?;
9624
9625 let response = AgentResponse::new(&question.question).with_metadata(
9626 "disambiguation",
9627 serde_json::json!({
9628 "status": if awaiting_confirmation { "awaiting_confirmation" } else { "awaiting_clarification" },
9629 "skill_id": skill_id,
9630 "options": question.options,
9631 "clarifying": question.clarifying,
9632 "detection": {
9633 "type": detection.ambiguity_type,
9634 "confidence": detection.confidence,
9635 "what_is_unclear": detection.what_is_unclear,
9636 }
9637 }),
9638 );
9639 drop(admission);
9640 self.finish_turn_if_root(&response).await?;
9641 return Ok(response);
9642 }
9643 DisambiguationResult::Clarified {
9644 enriched_input: re_enriched,
9645 ..
9646 } => {
9647 debug!(skill_id = %skill_id, "Skill re-check: clarified immediately, executing");
9648 let admission = self
9649 .admit_disambiguation_redispatch(
9650 expected_disambiguation_epoch,
9651 expected_state_generation,
9652 )
9653 .await?;
9654 *self.pending_skill_id.write() = None;
9655 drop(admission);
9656 let skill_response = self.execute_skill_by_id(skill_id, &re_enriched).await?;
9657 self.memory
9658 .add_message(ChatMessage::user(&re_enriched))
9659 .await?;
9660 return self
9661 .handle_skill_response(
9662 &re_enriched,
9663 skill_id,
9664 skill_response,
9665 &HashMap::new(),
9666 )
9667 .await;
9668 }
9669 DisambiguationResult::ProceedWithBestGuess {
9670 enriched_input: re_enriched,
9671 } => {
9672 debug!(skill_id = %skill_id, "Skill re-check: proceeding with best guess");
9673 let admission = self
9674 .admit_disambiguation_redispatch(
9675 expected_disambiguation_epoch,
9676 expected_state_generation,
9677 )
9678 .await?;
9679 *self.pending_skill_id.write() = None;
9680 drop(admission);
9681 let skill_response = self.execute_skill_by_id(skill_id, &re_enriched).await?;
9682 self.memory
9683 .add_message(ChatMessage::user(&re_enriched))
9684 .await?;
9685 return self
9686 .handle_skill_response(
9687 &re_enriched,
9688 skill_id,
9689 skill_response,
9690 &HashMap::new(),
9691 )
9692 .await;
9693 }
9694 DisambiguationResult::GiveUp { reason } => {
9695 *self.pending_skill_id.write() = None;
9696 let apology = self
9697 .generate_localized_apology(
9698 "Generate a brief, polite apology saying you couldn't understand the request. Be concise.",
9699 &reason,
9700 )
9701 .await
9702 .unwrap_or_else(|_| {
9703 format!("I'm sorry, I couldn't understand your request: {}", reason)
9704 });
9705 let response = AgentResponse::new(&apology);
9706 self.finish_turn_if_root(&response).await?;
9707 return Ok(response);
9708 }
9709 DisambiguationResult::Escalate { reason } => {
9710 *self.pending_skill_id.write() = None;
9711 let apology = self
9712 .generate_localized_apology(
9713 "Explain briefly that you're transferring the user to a human agent for help.",
9714 &reason,
9715 )
9716 .await
9717 .unwrap_or_else(|_| {
9718 format!("I need human assistance to help with your request: {}", reason)
9719 });
9720 let response = AgentResponse::new(&apology);
9721 self.finish_turn_if_root(&response).await?;
9722 return Ok(response);
9723 }
9724 DisambiguationResult::Abandoned { new_input } => {
9725 *self.pending_skill_id.write() = None;
9728 debug!(skill_id = %skill_id, "Skill re-check: abandoned by user");
9729 if let Some(fresh) = new_input {
9730 return self.run_loop_internal(&fresh).await;
9731 }
9732 let ack = self
9733 .generate_localized_apology(
9734 "The user changed their mind about their previous request. \
9735 Generate a brief, friendly acknowledgment (e.g. 'OK, no problem. What else can I help with?'). \
9736 Do NOT apologize excessively. Be concise.",
9737 "User abandoned clarification",
9738 )
9739 .await
9740 .unwrap_or_else(|_| {
9741 "OK, no problem. What else can I help with?".to_string()
9742 });
9743 self.memory
9744 .add_message(ChatMessage::assistant(&ack))
9745 .await?;
9746 let response = AgentResponse::new(&ack);
9747 self.finish_turn_if_root(&response).await?;
9748 return Ok(response);
9749 }
9750 }
9751 }
9752
9753 let admission = self
9755 .admit_disambiguation_redispatch(
9756 expected_disambiguation_epoch,
9757 expected_state_generation,
9758 )
9759 .await?;
9760 *self.pending_skill_id.write() = None;
9761 drop(admission);
9762 let skill_response = self.execute_skill_by_id(skill_id, enriched_input).await?;
9763 self.memory
9764 .add_message(ChatMessage::user(enriched_input))
9765 .await?;
9766 self.handle_skill_response(enriched_input, skill_id, skill_response, &HashMap::new())
9767 .await
9768 }
9769
9770 async fn handle_skill_response(
9773 &self,
9774 processed_input: &str,
9775 skill_id: &str,
9776 skill_response: String,
9777 input_context: &HashMap<String, Value>,
9778 ) -> Result<AgentResponse> {
9779 let output_data = self.process_output(&skill_response, input_context).await?;
9780 let final_response = output_data.content;
9781
9782 self.memory
9783 .add_message(ChatMessage::assistant(&final_response))
9784 .await?;
9785
9786 self.check_memory_compression().await?;
9787
9788 self.increment_turn();
9789 self.evaluate_transitions(processed_input, &final_response)
9790 .await?;
9791
9792 let response = AgentResponse::new(final_response)
9793 .with_metadata("skill_id", serde_json::json!(skill_id));
9794 self.finish_turn_if_root(&response).await?;
9795 Ok(response)
9796 }
9797
9798 async fn handle_plan_and_execute(
9801 &self,
9802 processed_input: &str,
9803 input_context: &HashMap<String, Value>,
9804 auto_detected: bool,
9805 ) -> Result<AgentResponse> {
9806 let effective = self.get_effective_reasoning_config();
9807 let plan_reflection = effective
9808 .get_planning()
9809 .map(|c| c.reflection.clone())
9810 .unwrap_or_default();
9811
9812 let max_attempts = if plan_reflection.enabled {
9813 1 + plan_reflection.max_replans
9814 } else {
9815 1
9816 };
9817
9818 let mut plan = self.generate_plan(processed_input).await?;
9819 info!(
9820 plan_id = %plan.id,
9821 steps = plan.steps.len(),
9822 "Plan generated"
9823 );
9824
9825 let mut plan_result = String::new();
9826
9827 for attempt in 0..max_attempts {
9828 *self.current_plan.write() = Some(plan.clone());
9829 plan_result = self.execute_plan(&mut plan).await?;
9830
9831 info!(
9832 plan_status = ?plan.status,
9833 completed_steps = plan.completed_steps().count(),
9834 attempt = attempt + 1,
9835 "Plan execution completed"
9836 );
9837
9838 if !plan_reflection.enabled {
9839 break;
9840 }
9841
9842 let has_failures = plan.steps.iter().any(|s| s.status.is_failed());
9843 if !has_failures {
9844 break;
9845 }
9846
9847 if attempt + 1 >= max_attempts {
9848 break;
9849 }
9850
9851 match plan_reflection.on_step_failure {
9852 StepFailureAction::Replan => {
9853 info!(attempt = attempt + 1, "Plan had failures, replanning");
9854 plan = self.generate_plan(processed_input).await?;
9855 }
9856 StepFailureAction::Abort => {
9857 warn!("Plan step failed, aborting");
9858 break;
9859 }
9860 StepFailureAction::Skip | StepFailureAction::Continue => {
9861 break;
9862 }
9863 }
9864 }
9865
9866 *self.current_plan.write() = Some(plan);
9867
9868 let output_data = self.process_output(&plan_result, input_context).await?;
9869 let final_content = output_data.content;
9870
9871 self.memory
9872 .add_message(ChatMessage::assistant(&final_content))
9873 .await?;
9874
9875 self.check_memory_compression().await?;
9876 self.increment_turn();
9877 self.evaluate_transitions(processed_input, &final_content)
9878 .await?;
9879
9880 let reasoning_metadata =
9881 ReasoningMetadata::new(ReasoningMode::PlanAndExecute).with_auto_detected(auto_detected);
9882
9883 let response = AgentResponse::new(&final_content).with_metadata(
9884 "reasoning",
9885 serde_json::to_value(&reasoning_metadata).unwrap_or_default(),
9886 );
9887
9888 self.finish_turn_if_root(&response).await?;
9889 Ok(response)
9890 }
9891
9892 fn inject_reasoning_prompt(
9894 &self,
9895 messages: &mut [ChatMessage],
9896 reasoning_mode: &ReasoningMode,
9897 is_first_iteration: bool,
9898 ) {
9899 if !is_first_iteration {
9900 return;
9901 }
9902 match reasoning_mode {
9903 ReasoningMode::CoT => {
9904 if let Some(msg) = messages.first_mut()
9905 && matches!(msg.role, ai_agents_core::Role::System)
9906 {
9907 msg.content = self.build_cot_system_prompt(&msg.content);
9908 debug!("Applied Chain-of-Thought system prompt");
9909 }
9910 }
9911 ReasoningMode::React => {
9912 if let Some(msg) = messages.first_mut()
9913 && matches!(msg.role, ai_agents_core::Role::System)
9914 {
9915 msg.content = self.build_react_system_prompt(&msg.content);
9916 debug!("Applied ReAct system prompt");
9917 }
9918 }
9919 _ => {}
9920 }
9921 }
9922
9923 async fn generate_main_response_draft(
9928 &self,
9929 processed_input: &str,
9930 reasoning_mode: &ReasoningMode,
9931 ) -> Result<MainResponseDraft> {
9932 let llm = self.get_state_llm()?;
9933 let protocol = self.main_tool_protocol(llm.as_ref(), true).await?;
9934 let mut messages = self
9935 .build_messages_internal(false, Some(processed_input), protocol.choice.is_none())
9936 .await?;
9937 self.inject_reasoning_prompt(&mut messages, reasoning_mode, true);
9938 let response = self
9939 .complete_main_llm_with_recovery(llm, &messages, &protocol)
9940 .await?;
9941 let content = response.content.trim().to_string();
9942 let (thinking, answer) = self.extract_thinking(&content);
9943 if let Some(calls) = self.parse_main_tool_calls(&content, &protocol)? {
9944 return Ok(MainResponseDraft::ToolCalls {
9945 raw_content: content,
9946 calls,
9947 thinking,
9948 });
9949 }
9950 Ok(MainResponseDraft::Text {
9951 raw_content: answer,
9952 thinking,
9953 })
9954 }
9955
9956 async fn commit_main_response_draft(
9961 &self,
9962 processed_input: &str,
9963 input_context: &HashMap<String, Value>,
9964 draft: MainResponseDraft,
9965 reasoning_mode: ReasoningMode,
9966 auto_detected: bool,
9967 ) -> Result<AgentResponse> {
9968 self.commit_root_user_message(processed_input).await?;
9969 match draft {
9970 MainResponseDraft::Text {
9971 raw_content,
9972 thinking,
9973 } => {
9974 self.finish_text_response_from_model(CommittedTextResponse {
9975 processed_input,
9976 input_context,
9977 answer: raw_content,
9978 reasoning_mode,
9979 auto_detected,
9980 iterations: 1,
9981 thinking_content: thinking,
9982 all_tool_calls: Vec::new(),
9983 })
9984 .await
9985 }
9986 MainResponseDraft::ToolCalls {
9987 raw_content,
9988 calls,
9989 thinking: _,
9990 } => {
9991 let mut all_tool_calls = Vec::new();
9992 match self
9993 .handle_tool_calls(
9994 processed_input,
9995 &raw_content,
9996 calls,
9997 &mut all_tool_calls,
9998 None,
9999 )
10000 .await?
10001 {
10002 ToolCallOutcome::Rejected(response) => {
10003 self.finish_turn_if_root(&response).await?;
10004 Ok(response)
10005 }
10006 ToolCallOutcome::Continue | ToolCallOutcome::TransitionFired => {
10007 self.continue_after_committed_tool_draft(processed_input)
10008 .await
10009 }
10010 }
10011 }
10012 }
10013 }
10014
10015 async fn continue_after_committed_tool_draft(
10020 &self,
10021 processed_input: &str,
10022 ) -> Result<AgentResponse> {
10023 *self.redispatch_depth.write() += 1;
10024 if let Some(context) = self.active_turn_context.write().as_mut() {
10025 context.enter_redispatch();
10026 }
10027 let result = Box::pin(self.run_loop_internal(processed_input)).await;
10028 *self.redispatch_depth.write() -= 1;
10029 if let Some(context) = self.active_turn_context.write().as_mut() {
10030 context.exit_redispatch();
10031 }
10032 let response = result?;
10033 self.finish_turn_if_root(&response).await?;
10034 Ok(response)
10035 }
10036
10037 async fn finish_text_response_from_model(
10042 &self,
10043 response: CommittedTextResponse<'_>,
10044 ) -> Result<AgentResponse> {
10045 let CommittedTextResponse {
10046 processed_input,
10047 input_context,
10048 answer,
10049 reasoning_mode,
10050 auto_detected,
10051 iterations,
10052 thinking_content,
10053 all_tool_calls,
10054 } = response;
10055 let output_data = self.process_output(&answer, input_context).await?;
10056 let mut final_content = if output_data.metadata.rejected {
10057 output_data
10058 .metadata
10059 .rejection_reason
10060 .unwrap_or_else(|| answer.to_string())
10061 } else {
10062 output_data.content
10063 };
10064 let llm = self.get_state_llm()?;
10065 let reflection_metadata;
10066 (final_content, reflection_metadata) = self
10067 .run_reflection(&*llm, processed_input, final_content)
10068 .await?;
10069 final_content =
10070 self.format_response_with_thinking(thinking_content.as_deref(), &final_content);
10071 let final_content = {
10072 let result = self
10073 .post_loop_processing(processed_input, final_content)
10074 .await?;
10075 self.apply_post_loop_result(processed_input, result)
10076 .await?
10077 .content
10078 };
10079 let response = self.build_agent_response(AgentResponseParts {
10080 content: final_content,
10081 all_tool_calls,
10082 reasoning_mode,
10083 auto_detected,
10084 iterations,
10085 thinking: thinking_content,
10086 reflection_metadata,
10087 });
10088 self.finish_turn_if_root(&response).await?;
10089 Ok(response)
10090 }
10091
10092 async fn run_committed_response_loop_with_reasoning(
10097 &self,
10098 processed_input: &str,
10099 input_context: &HashMap<String, Value>,
10100 reasoning_mode: ReasoningMode,
10101 auto_detected: bool,
10102 ) -> Result<AgentResponse> {
10103 self.commit_root_user_message(processed_input).await?;
10104 let llm = self.get_state_llm()?;
10105 let mut iterations = 0u32;
10106 let mut all_tool_calls = Vec::new();
10107 let mut thinking_content = None;
10108 loop {
10109 let effective_max = if reasoning_mode != ReasoningMode::None {
10110 let rc = self.get_effective_reasoning_config();
10111 self.max_iterations.min(rc.max_iterations)
10112 } else {
10113 self.max_iterations
10114 };
10115 if iterations >= effective_max {
10116 return Err(AgentError::Other(format!(
10117 "Max iterations ({}) exceeded",
10118 effective_max
10119 )));
10120 }
10121 iterations += 1;
10122 *self.iteration_count.write() = iterations;
10123 let protocol = self.main_tool_protocol(llm.as_ref(), false).await?;
10124 let mut messages = self
10125 .build_messages_internal(true, None, protocol.choice.is_none())
10126 .await?;
10127 self.inject_reasoning_prompt(&mut messages, &reasoning_mode, iterations == 1);
10128 self.hooks.on_llm_start(&messages).await;
10129 let llm_start = Instant::now();
10130 let response = self
10131 .complete_main_llm_with_recovery(Arc::clone(&llm), &messages, &protocol)
10132 .await?;
10133 let llm_duration_ms = llm_start.elapsed().as_millis() as u64;
10134 self.hooks.on_llm_complete(&response, llm_duration_ms).await;
10135 let content = response.content.trim();
10136 if let Some(tool_calls) = self.parse_main_tool_calls(content, &protocol)? {
10137 match self
10138 .handle_tool_calls(
10139 processed_input,
10140 content,
10141 tool_calls,
10142 &mut all_tool_calls,
10143 None,
10144 )
10145 .await?
10146 {
10147 ToolCallOutcome::Continue | ToolCallOutcome::TransitionFired => continue,
10148 ToolCallOutcome::Rejected(resp) => {
10149 self.finish_turn_if_root(&resp).await?;
10150 return Ok(resp);
10151 }
10152 }
10153 }
10154 let (extracted_thinking, answer) = self.extract_thinking(content);
10155 if extracted_thinking.is_some() {
10156 thinking_content = extracted_thinking;
10157 }
10158 return self
10159 .finish_text_response_from_model(CommittedTextResponse {
10160 processed_input,
10161 input_context,
10162 answer,
10163 reasoning_mode,
10164 auto_detected,
10165 iterations,
10166 thinking_content,
10167 all_tool_calls,
10168 })
10169 .await;
10170 }
10171 }
10172
10173 async fn handle_tool_calls(
10179 &self,
10180 processed_input: &str,
10181 content: &str,
10182 tool_calls: Vec<ToolCall>,
10183 all_tool_calls: &mut Vec<ToolCall>,
10184 mut events: Option<&mut Vec<StreamChunk>>,
10185 ) -> Result<ToolCallOutcome> {
10186 let include_tool_events = self.streaming.include_tool_events;
10187 let transition_content = native_readable_projection(content)
10191 .map_err(|error| AgentError::LLM(error.to_string()))?;
10192 let transition_fired = self
10193 .evaluate_transitions(processed_input, &transition_content)
10194 .await?;
10195 if transition_fired {
10196 self.memory
10197 .add_message(ChatMessage::assistant(
10198 "(Transitioned to new state — tool call handled by workflow)",
10199 ))
10200 .await?;
10201 if let Some(events) = events.as_deref_mut()
10202 && self.streaming.include_state_events
10203 && let Some(state) = self.current_state()
10204 {
10205 events.push(StreamChunk::state_transition(None, state));
10206 }
10207 return Ok(ToolCallOutcome::TransitionFired);
10208 }
10209
10210 self.memory
10212 .add_message(ChatMessage::assistant(content))
10213 .await?;
10214 self.remember_committed_native_exchange(content).await?;
10215 let native_tool_call = Self::is_native_tool_call_content(content)?;
10216
10217 if let Some(events) = events.as_deref_mut()
10218 && include_tool_events
10219 {
10220 for tool_call in &tool_calls {
10221 events.push(StreamChunk::tool_start(&tool_call.id, &tool_call.name));
10222 }
10223 }
10224 let results = self.execute_tools_parallel(&tool_calls).await;
10225 let mut rejection = None;
10226
10227 for ((_id, result), tool_call) in results.into_iter().zip(tool_calls.iter()) {
10228 match result {
10229 Ok(output) => {
10230 if let Some(events) = events.as_deref_mut()
10231 && include_tool_events
10232 {
10233 events.push(StreamChunk::tool_result(
10234 &tool_call.id,
10235 &tool_call.name,
10236 &output,
10237 true,
10238 ));
10239 }
10240 self.memory
10241 .add_message(Self::tool_result_message(
10242 tool_call,
10243 &output,
10244 native_tool_call,
10245 )?)
10246 .await?;
10247 }
10248 Err(e) => {
10249 if matches!(e, AgentError::HITLRejected(_)) {
10250 if !native_tool_call {
10251 self.memory
10252 .add_message(ChatMessage::assistant(format!(
10253 "The operation was rejected by the approver: {e}"
10254 )))
10255 .await?;
10256 return Ok(ToolCallOutcome::Rejected(AgentResponse {
10257 content: format!("Operation cancelled: {e}"),
10258 metadata: None,
10259 tool_calls: Some(all_tool_calls.clone()),
10260 }));
10261 }
10262 if rejection.is_none() {
10263 rejection = Some(e.to_string());
10264 }
10265 }
10266 if let Some(events) = events.as_deref_mut()
10267 && include_tool_events
10268 {
10269 events.push(StreamChunk::tool_result(
10270 &tool_call.id,
10271 &tool_call.name,
10272 e.to_string(),
10273 false,
10274 ));
10275 }
10276 self.memory
10277 .add_message(Self::tool_result_message(
10278 tool_call,
10279 &format!("Error: {}", e),
10280 native_tool_call,
10281 )?)
10282 .await?;
10283 }
10284 }
10285 all_tool_calls.push(tool_call.clone());
10286 if let Some(events) = events.as_deref_mut()
10287 && include_tool_events
10288 {
10289 events.push(StreamChunk::tool_end(&tool_call.id));
10290 }
10291 }
10292 if let Some(rejection) = rejection {
10293 self.memory
10294 .add_message(ChatMessage::assistant(format!(
10295 "The operation was rejected by the approver: {rejection}"
10296 )))
10297 .await?;
10298 return Ok(ToolCallOutcome::Rejected(AgentResponse {
10299 content: format!("Operation cancelled: {rejection}"),
10300 metadata: None,
10301 tool_calls: Some(all_tool_calls.clone()),
10302 }));
10303 }
10304 Ok(ToolCallOutcome::Continue)
10305 }
10306
10307 async fn run_reflection(
10309 &self,
10310 llm: &dyn LLMProvider,
10311 processed_input: &str,
10312 mut content: String,
10313 ) -> Result<(String, Option<ReflectionMetadata>)> {
10314 let config = self.get_effective_reflection_config();
10315 let should_reflect = self
10316 .should_reflect_with_config(processed_input, &content, &config)
10317 .await?;
10318 if !should_reflect {
10319 return Ok((content, None));
10320 }
10321
10322 info!("Starting response reflection evaluation");
10323 let mut attempts = 0u32;
10324 let max_retries = config.max_retries;
10325 let mut history: Vec<ReflectionAttempt> = Vec::new();
10326
10327 loop {
10328 let evaluation = self
10329 .evaluate_response_with_config(processed_input, &content, &config)
10330 .await?;
10331
10332 if evaluation.passed || attempts >= max_retries {
10333 info!(
10334 passed = evaluation.passed,
10335 confidence = evaluation.confidence,
10336 attempts = attempts + 1,
10337 "Reflection evaluation complete"
10338 );
10339 let reflection_metadata = Some(
10340 ReflectionMetadata::new(evaluation)
10341 .with_attempts(attempts + 1)
10342 .with_history(history),
10343 );
10344 return Ok((content, reflection_metadata));
10345 }
10346
10347 debug!(
10348 attempt = attempts + 1,
10349 failed_criteria = evaluation.failed_criteria().count(),
10350 "Response did not meet criteria, retrying"
10351 );
10352
10353 history.push(
10354 ReflectionAttempt::new(&content, evaluation.clone())
10355 .with_feedback("Response did not meet quality criteria"),
10356 );
10357
10358 let feedback: Vec<String> = evaluation
10359 .failed_criteria()
10360 .map(|c| format!("- {}", c.criterion))
10361 .collect();
10362
10363 let retry_prompt = format!(
10364 "Your previous response did not meet these criteria:\n{}\n\nPlease provide an improved response.",
10365 feedback.join("\n")
10366 );
10367
10368 self.memory
10369 .add_message(ChatMessage::user(&retry_prompt))
10370 .await?;
10371
10372 let retry_messages = self.build_messages().await?;
10373 let retry_response = self
10374 .observe_purpose(
10375 ObservationPurpose::ReflectionEvaluation,
10376 llm.complete(&retry_messages, None),
10377 )
10378 .await
10379 .map_err(|e| AgentError::LLM(e.to_string()))?;
10380
10381 content = retry_response.content.trim().to_string();
10382 attempts += 1;
10383 }
10384 }
10385
10386 async fn post_loop_processing(
10389 &self,
10390 processed_input: &str,
10391 content: String,
10392 ) -> Result<PostLoopResult> {
10393 self.increment_turn();
10398
10399 self.run_context_extractors(processed_input).await;
10401
10402 let transitioned = self.evaluate_transitions(processed_input, &content).await?;
10403
10404 if !transitioned {
10405 self.memory
10406 .add_message(ChatMessage::assistant(&content))
10407 .await?;
10408 self.check_memory_compression().await?;
10409 return Ok(PostLoopResult::NoTransition(content));
10410 }
10411
10412 if !self.should_regenerate_after_transition() {
10414 self.memory
10415 .add_message(ChatMessage::assistant(&content))
10416 .await?;
10417 self.check_memory_compression().await?;
10418 return Ok(PostLoopResult::Transitioned {
10419 content,
10420 regenerated: false,
10421 });
10422 }
10423
10424 if self.needs_redispatch_for_new_state() {
10428 info!("Post-transition NeedsRedispatch: new state requires full dispatch");
10429 return Ok(PostLoopResult::NeedsRedispatch);
10432 }
10433
10434 self.memory
10437 .add_message(ChatMessage::assistant(&content))
10438 .await?;
10439 self.check_memory_compression().await?;
10440
10441 let new_llm = self.get_state_llm()?;
10447 let mut final_content;
10448
10449 for post_iter in 0..self.max_iterations {
10450 let protocol = self.main_tool_protocol(new_llm.as_ref(), false).await?;
10451 let new_messages = self
10452 .build_messages_internal(true, None, protocol.choice.is_none())
10453 .await?;
10454 if post_iter == 0
10455 && let Some(system_msg) = new_messages.first()
10456 && system_msg.role == ai_agents_core::Role::System
10457 {
10458 debug!(
10459 prompt_preview =
10460 &system_msg.content[system_msg.content.len().saturating_sub(200)..],
10461 "Post-transition system prompt (last 200 chars)"
10462 );
10463 }
10464
10465 let new_response = self
10466 .complete_main_llm_with_recovery(Arc::clone(&new_llm), &new_messages, &protocol)
10467 .await?;
10468 final_content = new_response.content.trim().to_string();
10469
10470 if let Some(tool_calls) = self.parse_main_tool_calls(&final_content, &protocol)? {
10473 let native_tool_call = Self::is_native_tool_call_content(&final_content)?;
10474 debug!(
10475 post_iter = post_iter,
10476 tools = tool_calls.len(),
10477 "Post-transition tool call detected, executing"
10478 );
10479
10480 self.memory
10481 .add_message(ChatMessage::assistant(&final_content))
10482 .await?;
10483 self.remember_committed_native_exchange(&final_content)
10484 .await?;
10485
10486 let results = self.execute_tools_parallel(&tool_calls).await;
10487 let mut rejection = None;
10488 for ((_id, result), tool_call) in results.into_iter().zip(tool_calls.iter()) {
10489 match result {
10490 Ok(output) => {
10491 self.memory
10492 .add_message(Self::tool_result_message(
10493 tool_call,
10494 &output,
10495 native_tool_call,
10496 )?)
10497 .await?;
10498 }
10499 Err(e) => {
10500 if native_tool_call
10501 && rejection.is_none()
10502 && matches!(e, AgentError::HITLRejected(_))
10503 {
10504 rejection = Some(e.to_string());
10505 }
10506 self.memory
10507 .add_message(Self::tool_result_message(
10508 tool_call,
10509 &format!("Error: {}", e),
10510 native_tool_call,
10511 )?)
10512 .await?;
10513 }
10514 }
10515 }
10516 if let Some(rejection) = rejection {
10517 self.memory
10518 .add_message(ChatMessage::assistant(format!(
10519 "The operation was rejected by the approver: {rejection}"
10520 )))
10521 .await?;
10522 return Err(AgentError::HITLRejected(rejection));
10523 }
10524 continue;
10526 }
10527
10528 self.memory
10530 .add_message(ChatMessage::assistant(&final_content))
10531 .await?;
10532 return Ok(PostLoopResult::Transitioned {
10533 content: final_content,
10534 regenerated: true,
10535 });
10536 }
10537
10538 final_content = "Post-transition processing completed.".to_string();
10540 self.memory
10541 .add_message(ChatMessage::assistant(&final_content))
10542 .await?;
10543
10544 Ok(PostLoopResult::Transitioned {
10545 content: final_content,
10546 regenerated: true,
10547 })
10548 }
10549
10550 fn should_regenerate_after_transition(&self) -> bool {
10553 if let Some(ref sm) = self.state_machine {
10554 if !sm.config().regenerate_on_transition {
10556 return false;
10557 }
10558 if let Some(def) = sm.current_definition()
10560 && let Some(regen) = def.regenerate_on_enter
10561 {
10562 return regen;
10563 }
10564 }
10565 true
10566 }
10567
10568 fn needs_redispatch_for_new_state(&self) -> bool {
10571 if let Some(ref sm) = self.state_machine
10572 && let Some(def) = sm.current_definition()
10573 {
10574 if def.concurrent.is_some()
10575 || def.group_chat.is_some()
10576 || def.pipeline.is_some()
10577 || def.handoff.is_some()
10578 || def.delegate.is_some()
10579 {
10580 return true;
10581 }
10582 let effective = self.get_effective_reasoning_config();
10584 if !matches!(effective.mode, ReasoningMode::None) {
10585 return true;
10586 }
10587 }
10588 false
10589 }
10590
10591 async fn apply_post_loop_result(
10597 &self,
10598 processed_input: &str,
10599 result: PostLoopResult,
10600 ) -> Result<AppliedPostLoop> {
10601 match result {
10602 PostLoopResult::NoTransition(content) => Ok(AppliedPostLoop {
10603 content,
10604 transitioned: false,
10605 regenerated: false,
10606 }),
10607 PostLoopResult::Transitioned {
10608 content,
10609 regenerated,
10610 } => Ok(AppliedPostLoop {
10611 content,
10612 transitioned: true,
10613 regenerated,
10614 }),
10615 PostLoopResult::NeedsRedispatch => {
10616 const MAX_REDISPATCH_DEPTH: u32 = 3;
10617 let current_depth = *self.redispatch_depth.read();
10618 if current_depth >= MAX_REDISPATCH_DEPTH {
10619 warn!(
10620 depth = current_depth,
10621 "Post-transition re-dispatch depth limit reached, returning empty response"
10622 );
10623 let content = String::new();
10624 self.memory
10625 .add_message(ChatMessage::assistant(&content))
10626 .await?;
10627 return Ok(AppliedPostLoop {
10629 content,
10630 transitioned: true,
10631 regenerated: false,
10632 });
10633 }
10634 *self.redispatch_depth.write() += 1;
10635 if let Some(context) = self.active_turn_context.write().as_mut() {
10636 context.enter_redispatch();
10637 }
10638 info!(
10639 depth = current_depth + 1,
10640 "Re-dispatching for new state after transition"
10641 );
10642 let resp = Box::pin(self.run_loop_internal(processed_input)).await;
10643 *self.redispatch_depth.write() -= 1;
10644 if let Some(context) = self.active_turn_context.write().as_mut() {
10645 context.exit_redispatch();
10646 }
10647 resp.map(|r| AppliedPostLoop {
10648 content: r.content,
10649 transitioned: true,
10650 regenerated: true,
10651 })
10652 }
10653 }
10654 }
10655
10656 fn build_agent_response(&self, parts: AgentResponseParts) -> AgentResponse {
10658 let AgentResponseParts {
10659 content,
10660 all_tool_calls,
10661 reasoning_mode,
10662 auto_detected,
10663 iterations,
10664 thinking,
10665 reflection_metadata,
10666 } = parts;
10667 let reasoning_metadata = ReasoningMetadata::new(reasoning_mode.clone())
10668 .with_thinking(thinking.clone().unwrap_or_default())
10669 .with_iterations(iterations)
10670 .with_auto_detected(auto_detected);
10671
10672 let mut response = AgentResponse::new(&content);
10673 if !all_tool_calls.is_empty() {
10674 response = response.with_tool_calls(all_tool_calls);
10675 }
10676
10677 if let Some(state) = self.current_state() {
10678 response = response.with_metadata("current_state", serde_json::json!(state));
10679 }
10680
10681 response = response.with_metadata(
10682 "reasoning",
10683 serde_json::to_value(&reasoning_metadata).unwrap_or_default(),
10684 );
10685
10686 if let Some(ref refl_meta) = reflection_metadata {
10687 response = response.with_metadata(
10688 "reflection",
10689 serde_json::to_value(refl_meta).unwrap_or_default(),
10690 );
10691 }
10692
10693 response
10694 }
10695
10696 async fn handle_delegated_state(
10698 &self,
10699 input: &str,
10700 delegate_id: &str,
10701 state_def: &ai_agents_state::StateDefinition,
10702 ) -> Result<AgentResponse> {
10703 use std::time::Instant;
10704
10705 let registry = self.spawner_registry.as_ref().ok_or_else(|| {
10706 AgentError::Config(format!(
10707 "State delegates to '{}' but no agent registry is configured. \
10708 Add a spawner section with auto_spawn to your YAML.",
10709 delegate_id
10710 ))
10711 })?;
10712
10713 let state_name = self
10714 .state_machine
10715 .as_ref()
10716 .map(|sm| sm.current())
10717 .unwrap_or_else(|| "unknown".to_string());
10718
10719 self.hooks.on_delegate_start(delegate_id, &state_name).await;
10720 let start = Instant::now();
10721
10722 let delegate = registry.get(delegate_id).ok_or_else(|| {
10723 AgentError::Other(format!(
10724 "State '{}' delegates to '{}' but no agent with that ID exists in the registry.",
10725 state_name, delegate_id
10726 ))
10727 })?;
10728
10729 let context_mode = state_def.delegate_context.clone().unwrap_or_default();
10731 let effective_input = self
10732 .observe_purpose(
10733 ObservationPurpose::OrchestrationRouting,
10734 crate::orchestration::context::prepare_delegate_input(
10735 input,
10736 &context_mode,
10737 &*self.memory,
10738 self.llm_registry.get("router").ok().as_deref(),
10739 ),
10740 )
10741 .await?;
10742
10743 let response = delegate
10744 .chat_with_actor_context(&effective_input, self.outbound_actor_context())
10745 .await?;
10746
10747 let duration_ms = start.elapsed().as_millis() as u64;
10748 self.hooks
10749 .on_delegate_complete(delegate_id, &state_name, duration_ms)
10750 .await;
10751
10752 let ctx_key = format!("delegation.{}.last_response", delegate_id);
10754 let _ = self.context_manager.set(
10755 &ctx_key,
10756 serde_json::Value::String(response.content.clone()),
10757 );
10758
10759 let _ = self.context_manager.set(
10761 "orchestration",
10762 serde_json::json!({
10763 "type": "delegate",
10764 "agent": delegate_id,
10765 "state": state_name,
10766 "response": response.content,
10767 "duration_ms": duration_ms,
10768 }),
10769 );
10770
10771 self.commit_root_user_message(input).await?;
10772
10773 let post_result = self
10776 .post_loop_processing(
10777 input,
10778 format!("[Delegated to {}]: {}", delegate_id, response.content),
10779 )
10780 .await?;
10781 let final_content = self
10782 .apply_post_loop_result(input, post_result)
10783 .await?
10784 .content;
10785
10786 let mut result = AgentResponse::new(final_content);
10787
10788 let metadata = serde_json::json!({
10789 "orchestration": {
10790 "type": "delegate",
10791 "agent": delegate_id,
10792 "state": state_name,
10793 "response": response.content,
10794 "duration_ms": duration_ms,
10795 }
10796 });
10797 result.metadata = Some(
10798 serde_json::from_value::<std::collections::HashMap<String, serde_json::Value>>(
10799 metadata,
10800 )
10801 .unwrap_or_default(),
10802 );
10803
10804 self.finish_turn_if_root(&result).await?;
10805 Ok(result)
10806 }
10807
10808 async fn handle_concurrent_state(
10810 &self,
10811 input: &str,
10812 config: &ai_agents_state::ConcurrentStateConfig,
10813 ) -> Result<AgentResponse> {
10814 use std::time::Instant;
10815
10816 let registry = self.spawner_registry.as_ref().ok_or_else(|| {
10817 AgentError::Config(
10818 "Concurrent state requires an agent registry. Add a spawner section.".into(),
10819 )
10820 })?;
10821
10822 let context_mode = config.context_mode.clone().unwrap_or_default();
10827 let context_input = self
10828 .observe_purpose(
10829 ObservationPurpose::OrchestrationRouting,
10830 crate::orchestration::context::prepare_delegate_input(
10831 input,
10832 &context_mode,
10833 &*self.memory,
10834 self.llm_registry.get("router").ok().as_deref(),
10835 ),
10836 )
10837 .await?;
10838
10839 let effective_input = if let Some(ref tmpl) = config.input {
10840 render_concurrent_template(tmpl, &context_input, &self.build_context_with_overlays())
10841 .unwrap_or_else(|_| context_input.clone())
10842 } else {
10843 context_input
10844 };
10845
10846 let start = Instant::now();
10847
10848 let llm_name = config
10849 .aggregation
10850 .synthesizer_llm
10851 .as_deref()
10852 .unwrap_or("router");
10853 let llm_provider = self.llm_registry.get(llm_name).ok();
10854
10855 let vote_parallelism = if self.runtime_config.optimization.enabled
10856 && self
10857 .runtime_config
10858 .optimization
10859 .parallel_orchestration_vote_extraction
10860 {
10861 Some(self.runtime_config.optimization.max_parallel_runtime_tasks)
10862 } else {
10863 None
10864 };
10865
10866 let result = self
10867 .observe_purpose(
10868 ObservationPurpose::OrchestrationAggregation,
10869 scope_actor_context(
10870 self.outbound_actor_context(),
10871 crate::orchestration::concurrent(
10872 registry,
10873 &effective_input,
10874 &config.agents,
10875 &config.aggregation,
10876 llm_provider.as_deref(),
10877 config.min_required,
10878 config.timeout_ms,
10879 config.on_partial_failure.clone(),
10880 vote_parallelism,
10881 ),
10882 ),
10883 )
10884 .await?;
10885
10886 let duration_ms = start.elapsed().as_millis() as u64;
10887 let agent_ids: Vec<String> = config.agents.iter().map(|a| a.id().to_string()).collect();
10888 let strategy = format!("{:?}", config.aggregation.strategy);
10889 self.hooks
10890 .on_concurrent_complete(&agent_ids, &strategy, duration_ms)
10891 .await;
10892
10893 let _ = self.context_manager.set(
10895 "concurrent.result",
10896 serde_json::Value::String(result.response.content.clone()),
10897 );
10898
10899 let agents_json: Vec<serde_json::Value> = result
10901 .agent_results
10902 .iter()
10903 .map(|ar| {
10904 serde_json::json!({
10905 "id": ar.agent_id,
10906 "response": ar.response.as_ref().map(|r| r.content.as_str()),
10907 "success": ar.success,
10908 "error": ar.error,
10909 "duration_ms": ar.duration_ms,
10910 })
10911 })
10912 .collect();
10913
10914 let _ = self.context_manager.set(
10916 "orchestration",
10917 serde_json::json!({
10918 "type": "concurrent",
10919 "result": result.response.content,
10920 "strategy": strategy,
10921 "agents": agents_json,
10922 "duration_ms": duration_ms,
10923 }),
10924 );
10925
10926 self.commit_root_user_message(input).await?;
10927
10928 let post_result = self
10929 .post_loop_processing(input, result.response.content.clone())
10930 .await?;
10931 let final_content = self
10932 .apply_post_loop_result(input, post_result)
10933 .await?
10934 .content;
10935
10936 let mut response = AgentResponse::new(final_content);
10937 let metadata = serde_json::json!({
10938 "orchestration": {
10939 "type": "concurrent",
10940 "result": result.response.content,
10941 "strategy": strategy,
10942 "agents": agents_json,
10943 "duration_ms": duration_ms,
10944 }
10945 });
10946 response.metadata = Some(
10947 serde_json::from_value::<std::collections::HashMap<String, serde_json::Value>>(
10948 metadata,
10949 )
10950 .unwrap_or_default(),
10951 );
10952
10953 self.finish_turn_if_root(&response).await?;
10954 Ok(response)
10955 }
10956
10957 async fn handle_group_chat_state(
10959 &self,
10960 input: &str,
10961 config: &ai_agents_state::GroupChatStateConfig,
10962 ) -> Result<AgentResponse> {
10963 use std::time::Instant;
10964
10965 let registry = self.spawner_registry.as_ref().ok_or_else(|| {
10966 AgentError::Config(
10967 "Group chat state requires an agent registry. Add a spawner section.".into(),
10968 )
10969 })?;
10970
10971 let start = Instant::now();
10972
10973 let llm_provider = self.llm_registry.get("router").ok();
10974
10975 let context_mode = config.context_mode.clone().unwrap_or_default();
10977 let context_input = self
10978 .observe_purpose(
10979 ObservationPurpose::OrchestrationRouting,
10980 crate::orchestration::context::prepare_delegate_input(
10981 input,
10982 &context_mode,
10983 &*self.memory,
10984 self.llm_registry.get("router").ok().as_deref(),
10985 ),
10986 )
10987 .await?;
10988
10989 let effective_topic = if let Some(ref tmpl) = config.input {
10991 render_concurrent_template(tmpl, &context_input, &self.build_context_with_overlays())
10992 .unwrap_or_else(|_| context_input.clone())
10993 } else {
10994 context_input
10995 };
10996
10997 let result = self
10998 .observe_purpose(
10999 ObservationPurpose::OrchestrationConversation,
11000 scope_actor_context(
11001 self.outbound_actor_context(),
11002 crate::orchestration::group_chat(
11003 registry,
11004 &effective_topic,
11005 config,
11006 llm_provider.as_deref(),
11007 Some(&*self.hooks),
11008 ),
11009 ),
11010 )
11011 .await?;
11012
11013 let duration_ms = start.elapsed().as_millis() as u64;
11014
11015 let _ = self.context_manager.set(
11017 "group_chat.conclusion",
11018 serde_json::Value::String(result.response.content.clone()),
11019 );
11020
11021 let transcript_json: Vec<serde_json::Value> = result
11023 .transcript
11024 .iter()
11025 .map(|t| {
11026 serde_json::json!({
11027 "speaker": t.speaker,
11028 "round": t.round,
11029 "content": t.content,
11030 })
11031 })
11032 .collect();
11033
11034 let _ = self.context_manager.set(
11036 "orchestration",
11037 serde_json::json!({
11038 "type": "group_chat",
11039 "conclusion": result.response.content,
11040 "transcript": transcript_json,
11041 "rounds": result.rounds_completed,
11042 "termination": result.termination_reason,
11043 "duration_ms": duration_ms,
11044 }),
11045 );
11046
11047 self.commit_root_user_message(input).await?;
11048
11049 let post_result = self
11050 .post_loop_processing(input, result.response.content.clone())
11051 .await?;
11052 let final_content = self
11053 .apply_post_loop_result(input, post_result)
11054 .await?
11055 .content;
11056
11057 let mut response = AgentResponse::new(final_content);
11058 let metadata = serde_json::json!({
11059 "orchestration": {
11060 "type": "group_chat",
11061 "conclusion": result.response.content,
11062 "transcript": transcript_json,
11063 "rounds": result.rounds_completed,
11064 "termination": result.termination_reason,
11065 "duration_ms": duration_ms,
11066 }
11067 });
11068 response.metadata = Some(
11069 serde_json::from_value::<std::collections::HashMap<String, serde_json::Value>>(
11070 metadata,
11071 )
11072 .unwrap_or_default(),
11073 );
11074
11075 self.finish_turn_if_root(&response).await?;
11076 Ok(response)
11077 }
11078
11079 async fn handle_pipeline_state(
11081 &self,
11082 input: &str,
11083 config: &ai_agents_state::PipelineStateConfig,
11084 ) -> Result<AgentResponse> {
11085 use std::time::Instant;
11086
11087 let registry = self.spawner_registry.as_ref().ok_or_else(|| {
11088 AgentError::Config(
11089 "Pipeline state requires an agent registry. Add a spawner section.".into(),
11090 )
11091 })?;
11092
11093 let start = Instant::now();
11094
11095 let stages: Vec<crate::orchestration::PipelineStage> = config
11096 .stages
11097 .iter()
11098 .map(|entry| {
11099 let mut stage = crate::orchestration::PipelineStage::id(entry.id());
11100 if let Some(tmpl) = entry.input() {
11101 stage = stage.with_input(tmpl);
11102 }
11103 stage
11104 })
11105 .collect();
11106
11107 let context_mode = config.context_mode.clone().unwrap_or_default();
11109 let context_input = self
11110 .observe_purpose(
11111 ObservationPurpose::OrchestrationRouting,
11112 crate::orchestration::context::prepare_delegate_input(
11113 input,
11114 &context_mode,
11115 &*self.memory,
11116 self.llm_registry.get("router").ok().as_deref(),
11117 ),
11118 )
11119 .await?;
11120
11121 let context_values = self.build_context_with_overlays();
11122 let result = self
11123 .observe_purpose(
11124 ObservationPurpose::OrchestrationRouting,
11125 scope_actor_context(
11126 self.outbound_actor_context(),
11127 crate::orchestration::pipeline(
11128 registry,
11129 &context_input,
11130 &stages,
11131 config.timeout_ms,
11132 Some(&*self.hooks),
11133 Some(&context_values),
11134 ),
11135 ),
11136 )
11137 .await?;
11138
11139 let duration_ms = start.elapsed().as_millis() as u64;
11140
11141 let _ = self.context_manager.set(
11143 "pipeline.result",
11144 serde_json::Value::String(result.response.content.clone()),
11145 );
11146
11147 let stages_json: Vec<serde_json::Value> = result
11149 .stage_outputs
11150 .iter()
11151 .map(|s| {
11152 serde_json::json!({
11153 "agent_id": s.agent_id,
11154 "output": s.output,
11155 "duration_ms": s.duration_ms,
11156 "skipped": s.skipped,
11157 })
11158 })
11159 .collect();
11160
11161 let _ = self.context_manager.set(
11163 "orchestration",
11164 serde_json::json!({
11165 "type": "pipeline",
11166 "result": result.response.content,
11167 "stages": stages_json,
11168 "duration_ms": duration_ms,
11169 }),
11170 );
11171
11172 self.commit_root_user_message(input).await?;
11173
11174 let post_result = self
11175 .post_loop_processing(input, result.response.content.clone())
11176 .await?;
11177 let final_content = self
11178 .apply_post_loop_result(input, post_result)
11179 .await?
11180 .content;
11181
11182 let mut response = AgentResponse::new(final_content);
11183 let metadata = serde_json::json!({
11184 "orchestration": {
11185 "type": "pipeline",
11186 "result": result.response.content,
11187 "stages": stages_json,
11188 "duration_ms": duration_ms,
11189 }
11190 });
11191 response.metadata = Some(
11192 serde_json::from_value::<std::collections::HashMap<String, serde_json::Value>>(
11193 metadata,
11194 )
11195 .unwrap_or_default(),
11196 );
11197
11198 self.finish_turn_if_root(&response).await?;
11199 Ok(response)
11200 }
11201
11202 async fn handle_handoff_state(
11204 &self,
11205 input: &str,
11206 config: &ai_agents_state::HandoffStateConfig,
11207 ) -> Result<AgentResponse> {
11208 use std::time::Instant;
11209
11210 let registry = self.spawner_registry.as_ref().ok_or_else(|| {
11211 AgentError::Config(
11212 "Handoff state requires an agent registry. Add a spawner section.".into(),
11213 )
11214 })?;
11215
11216 let llm = self
11217 .llm_registry
11218 .get("router")
11219 .map_err(|_| AgentError::Config("Handoff state requires a router LLM.".into()))?;
11220
11221 let start = Instant::now();
11222
11223 let context_mode = config.context_mode.clone().unwrap_or_default();
11225 let context_input = self
11226 .observe_purpose(
11227 ObservationPurpose::OrchestrationRouting,
11228 crate::orchestration::context::prepare_delegate_input(
11229 input,
11230 &context_mode,
11231 &*self.memory,
11232 self.llm_registry.get("router").ok().as_deref(),
11233 ),
11234 )
11235 .await?;
11236
11237 let effective_input = if let Some(ref tmpl) = config.input {
11239 render_concurrent_template(tmpl, &context_input, &self.build_context_with_overlays())
11240 .unwrap_or_else(|_| context_input.clone())
11241 } else {
11242 context_input
11243 };
11244
11245 let result = self
11246 .observe_purpose(
11247 ObservationPurpose::OrchestrationRouting,
11248 scope_actor_context(
11249 self.outbound_actor_context(),
11250 crate::orchestration::handoff(
11251 registry,
11252 &effective_input,
11253 &config.initial_agent,
11254 &config.available_agents,
11255 config.max_handoffs,
11256 llm.as_ref(),
11257 Some(&*self.hooks),
11258 ),
11259 ),
11260 )
11261 .await?;
11262
11263 let duration_ms = start.elapsed().as_millis() as u64;
11264
11265 let _ = self.context_manager.set(
11267 "handoff.result",
11268 serde_json::Value::String(result.response.content.clone()),
11269 );
11270
11271 let chain_json: Vec<serde_json::Value> = result
11273 .handoff_chain
11274 .iter()
11275 .map(|h| {
11276 serde_json::json!({
11277 "from": h.from_agent,
11278 "to": h.to_agent,
11279 "reason": h.reason,
11280 })
11281 })
11282 .collect();
11283
11284 let _ = self.context_manager.set(
11286 "orchestration",
11287 serde_json::json!({
11288 "type": "handoff",
11289 "result": result.response.content,
11290 "final_agent": result.final_agent,
11291 "handoff_chain": chain_json,
11292 "duration_ms": duration_ms,
11293 }),
11294 );
11295
11296 self.commit_root_user_message(input).await?;
11297
11298 let post_result = self
11299 .post_loop_processing(input, result.response.content.clone())
11300 .await?;
11301 let final_content = self
11302 .apply_post_loop_result(input, post_result)
11303 .await?
11304 .content;
11305
11306 let mut response = AgentResponse::new(final_content);
11307 let metadata = serde_json::json!({
11308 "orchestration": {
11309 "type": "handoff",
11310 "result": result.response.content,
11311 "final_agent": result.final_agent,
11312 "handoff_chain": chain_json,
11313 "duration_ms": duration_ms,
11314 }
11315 });
11316 response.metadata = Some(
11317 serde_json::from_value::<std::collections::HashMap<String, serde_json::Value>>(
11318 metadata,
11319 )
11320 .unwrap_or_default(),
11321 );
11322
11323 self.finish_turn_if_root(&response).await?;
11324 Ok(response)
11325 }
11326
11327 async fn run_loop_internal(&self, input: &str) -> Result<AgentResponse> {
11329 self.begin_root_turn();
11330 self.pre_turn_session_lifecycle().await;
11332
11333 let input_data = self.process_input(input).await?;
11334 self.update_active_turn_context(&input_data.content, input_data.context.clone());
11335
11336 for (key, value) in &input_data.context {
11339 let _ = self.context_manager.set(key, value.clone());
11340 }
11341
11342 if input_data.metadata.rejected {
11343 let reason = input_data
11344 .metadata
11345 .rejection_reason
11346 .unwrap_or_else(|| "Input rejected".to_string());
11347 warn!(reason = %reason, "Input rejected");
11348 let response = AgentResponse::new(reason);
11349 self.finish_turn_if_root(&response).await?;
11350 return Ok(response);
11351 }
11352
11353 let processed_input = &input_data.content;
11354
11355 if let Some(response) = self.try_pre_response_transition(processed_input).await? {
11356 return Ok(response);
11357 }
11358
11359 if let Some(ref sm) = self.state_machine
11361 && let Some(def) = sm.current_definition()
11362 {
11363 if let Some(ref delegate_id) = def.delegate {
11364 return self
11365 .handle_delegated_state(processed_input, delegate_id, &def)
11366 .await;
11367 }
11368 if let Some(ref concurrent_config) = def.concurrent {
11369 return self
11370 .handle_concurrent_state(processed_input, concurrent_config)
11371 .await;
11372 }
11373 if let Some(ref group_chat_config) = def.group_chat {
11374 return self
11375 .handle_group_chat_state(processed_input, group_chat_config)
11376 .await;
11377 }
11378 if let Some(ref pipeline_config) = def.pipeline {
11379 return self
11380 .handle_pipeline_state(processed_input, pipeline_config)
11381 .await;
11382 }
11383 if let Some(ref handoff_config) = def.handoff {
11384 return self
11385 .handle_handoff_state(processed_input, handoff_config)
11386 .await;
11387 }
11388 }
11389
11390 if let Some(response) =
11395 Box::pin(self.try_speculative_branches(processed_input, &input_data.context)).await?
11396 {
11397 return Ok(response);
11398 }
11399
11400 match self.try_skill_route(processed_input).await? {
11401 SkillRouteResult::Response { skill_id, content } => {
11402 self.commit_root_user_message(processed_input).await?;
11403 return self
11404 .handle_skill_response(processed_input, &skill_id, content, &input_data.context)
11405 .await;
11406 }
11407 SkillRouteResult::NeedsClarification {
11408 response,
11409 ownership,
11410 } => {
11411 let admission = self
11412 .admit_optional_disambiguation_ownership(ownership)
11413 .await?;
11414 self.commit_root_user_message(processed_input).await?;
11415 if Self::skill_clarification_needs_memory_record(&response) {
11416 self.memory
11419 .add_message(ChatMessage::assistant(&response.content))
11420 .await?;
11421 }
11422 drop(admission);
11423 self.finish_turn_if_root(&response).await?;
11424 return Ok(response);
11425 }
11426 SkillRouteResult::NoMatch => {} }
11428
11429 let effective_reasoning = self.get_effective_reasoning_config();
11430 let reasoning_mode = self.determine_reasoning_mode(processed_input).await?;
11431 let auto_detected = matches!(effective_reasoning.mode, ReasoningMode::Auto);
11432 let reflection_enabled = self.routing_reflection_mode();
11434
11435 info!(
11436 reasoning_mode = ?reasoning_mode,
11437 auto_detected = auto_detected,
11438 reflection_enabled = ?reflection_enabled,
11439 "Reasoning mode determined"
11440 );
11441
11442 if matches!(reasoning_mode, ReasoningMode::PlanAndExecute) {
11443 self.commit_root_user_message(processed_input).await?;
11444 return self
11445 .handle_plan_and_execute(processed_input, &input_data.context, auto_detected)
11446 .await;
11447 }
11448
11449 self.commit_root_user_message(processed_input).await?;
11450
11451 let mut iterations = 0u32;
11452 let mut all_tool_calls: Vec<ToolCall> = Vec::new();
11453 let mut thinking_content: Option<String> = None;
11454
11455 let llm = self.get_state_llm()?;
11456
11457 loop {
11458 let effective_max = if reasoning_mode != ReasoningMode::None {
11460 let rc = self.get_effective_reasoning_config();
11461 self.max_iterations.min(rc.max_iterations)
11462 } else {
11463 self.max_iterations
11464 };
11465
11466 if iterations >= effective_max {
11467 let err = AgentError::Other(format!("Max iterations ({}) exceeded", effective_max));
11468 self.hooks.on_error(&err).await;
11469 error!(iterations = iterations, "Max iterations exceeded");
11470 return Err(err);
11471 }
11472 iterations += 1;
11473 *self.iteration_count.write() = iterations;
11474
11475 debug!(iteration = iterations, max = effective_max, "LLM call");
11476
11477 let protocol = self.main_tool_protocol(llm.as_ref(), false).await?;
11478 let mut messages = self
11479 .build_messages_internal(true, None, protocol.choice.is_none())
11480 .await?;
11481 self.inject_reasoning_prompt(&mut messages, &reasoning_mode, iterations == 1);
11482
11483 self.hooks.on_llm_start(&messages).await;
11484 let llm_start = Instant::now();
11485 let response = self
11486 .complete_main_llm_with_recovery(Arc::clone(&llm), &messages, &protocol)
11487 .await?;
11488
11489 let llm_duration_ms = llm_start.elapsed().as_millis() as u64;
11490 self.hooks.on_llm_complete(&response, llm_duration_ms).await;
11491
11492 let content = response.content.trim();
11493
11494 if let Some(tool_calls) = self.parse_main_tool_calls(content, &protocol)? {
11495 match self
11496 .handle_tool_calls(
11497 processed_input,
11498 content,
11499 tool_calls,
11500 &mut all_tool_calls,
11501 None,
11502 )
11503 .await?
11504 {
11505 ToolCallOutcome::Continue | ToolCallOutcome::TransitionFired => continue,
11506 ToolCallOutcome::Rejected(resp) => {
11507 self.finish_turn_if_root(&resp).await?;
11508 return Ok(resp);
11509 }
11510 }
11511 }
11512
11513 let (extracted_thinking, answer) = self.extract_thinking(content);
11514 if extracted_thinking.is_some() {
11515 thinking_content = extracted_thinking;
11516 }
11517
11518 let output_data = self.process_output(&answer, &input_data.context).await?;
11519
11520 let mut final_content = if output_data.metadata.rejected {
11521 output_data
11522 .metadata
11523 .rejection_reason
11524 .unwrap_or_else(|| answer.to_string())
11525 } else {
11526 output_data.content
11527 };
11528
11529 let reflection_metadata;
11531 (final_content, reflection_metadata) = self
11532 .run_reflection(&*llm, processed_input, final_content)
11533 .await?;
11534
11535 final_content =
11536 self.format_response_with_thinking(thinking_content.as_deref(), &final_content);
11537
11538 let final_content = {
11542 let result = self
11543 .post_loop_processing(processed_input, final_content)
11544 .await?;
11545 self.apply_post_loop_result(processed_input, result)
11546 .await?
11547 .content
11548 };
11549
11550 let reflected = reflection_metadata.is_some();
11551 let reasoning_mode_debug = format!("{:?}", reasoning_mode);
11552
11553 let response = self.build_agent_response(AgentResponseParts {
11554 content: final_content,
11555 all_tool_calls,
11556 reasoning_mode,
11557 auto_detected,
11558 iterations,
11559 thinking: thinking_content,
11560 reflection_metadata,
11561 });
11562
11563 self.finish_turn_if_root(&response).await?;
11564
11565 let tool_call_count = response.tool_calls.as_ref().map(|tc| tc.len()).unwrap_or(0);
11566 info!(
11567 tool_calls = tool_call_count,
11568 response_len = response.content.len(),
11569 reasoning_mode = %reasoning_mode_debug,
11570 reflected = reflected,
11571 "Chat completed"
11572 );
11573 return Ok(response);
11574 }
11575 }
11576
11577 async fn generate_buffered_streaming_draft(
11578 &self,
11579 processed_input: &str,
11580 routing_resolved: Arc<AtomicBool>,
11581 ) -> Result<StreamingDraftResult> {
11582 let llm = self.get_state_llm()?;
11583 if llm.configured_tool_choice().is_some() {
11584 let draft = self
11585 .generate_main_response_draft(processed_input, &ReasoningMode::None)
11586 .await?;
11587 return Ok(StreamingDraftResult::new(draft, Vec::new()));
11588 }
11589 let protocol = self.main_tool_protocol(llm.as_ref(), true).await?;
11591 let messages = self.build_messages_for_draft(processed_input).await?;
11592 let source = self
11593 .open_main_stream_with_recovery(Arc::clone(&llm), &messages, &protocol)
11594 .await?;
11595 let mut buffer = crate::optimization::StreamBranchBuffer::new(self.streaming.buffer_size)?;
11596 let mut chunks = Vec::new();
11597 let mut accumulated = String::new();
11598 match source {
11599 MainStreamSource::StaticResponse(text) => {
11600 accumulated.push_str(&text);
11601 let stream_chunk = StreamChunk::content(text);
11602 if routing_resolved.load(Ordering::SeqCst) {
11603 chunks.push(stream_chunk);
11604 } else {
11605 buffer.push(stream_chunk)?;
11606 }
11607 }
11608 MainStreamSource::Stream(mut stream) => {
11609 while let Some(chunk_result) = stream.next().await {
11610 let chunk = chunk_result.map_err(|e| AgentError::LLM(e.to_string()))?;
11611 accumulated.push_str(&chunk.delta);
11612 let stream_chunk = StreamChunk::content(chunk.delta);
11613 if routing_resolved.load(Ordering::SeqCst) {
11614 chunks.push(stream_chunk);
11615 } else {
11616 buffer.push(stream_chunk)?;
11617 }
11618 }
11619 }
11620 }
11621 chunks.splice(0..0, buffer.drain());
11622 let content = accumulated.trim().to_string();
11623 let draft = if let Some(calls) = self.parse_tool_calls(&content)? {
11624 MainResponseDraft::ToolCalls {
11625 raw_content: content,
11626 calls,
11627 thinking: None,
11628 }
11629 } else {
11630 MainResponseDraft::Text {
11631 raw_content: content,
11632 thinking: None,
11633 }
11634 };
11635 Ok(StreamingDraftResult::new(draft, chunks))
11636 }
11637
11638 async fn try_buffered_streaming_branches(
11639 &self,
11640 processed_input: &str,
11641 input_context: &HashMap<String, Value>,
11642 ) -> Result<Option<(AgentResponse, Vec<StreamChunk>)>> {
11643 let optimization = &self.runtime_config.optimization;
11644 if !optimization.enabled {
11645 return Ok(None);
11646 }
11647 if !matches!(
11653 self.get_effective_reasoning_config().mode,
11654 ReasoningMode::None
11655 ) {
11656 return Ok(None);
11657 }
11658 let transition_enabled =
11659 optimization.speculative_state_transitions && self.has_parallel_transition_candidates();
11660 if !transition_enabled {
11661 return Ok(None);
11662 }
11663 let mut branch_scheduler =
11664 TurnBranchScheduler::new(optimization.max_parallel_runtime_tasks)?;
11665 if !branch_scheduler.reserve_task() {
11666 return Ok(None);
11667 }
11668 if !self
11669 .reserve_active_speculative_llm_call(RuntimeOptimizationKind::BufferedStreamingRouting)
11670 {
11671 branch_scheduler.release_task();
11672 return Ok(None);
11673 }
11674 if !branch_scheduler.reserve_task() {
11675 branch_scheduler.release_task();
11676 return Ok(None);
11677 }
11678 let mut main_branch = RuntimeBranch::new(
11679 RuntimeTaskPurpose::MainResponse,
11680 RuntimeOptimizationKind::BufferedStreamingRouting,
11681 RuntimeTaskPriority::Normal,
11682 RuntimeCommitBehavior::FinalResponse,
11683 );
11684 let mut transition_branch = RuntimeBranch::new(
11685 RuntimeTaskPurpose::StateTransition,
11686 RuntimeOptimizationKind::ParallelStateTransition,
11687 RuntimeTaskPriority::Critical,
11688 RuntimeCommitBehavior::TransitionDecision,
11689 );
11690 let main_id = main_branch.branch_id();
11691 let transition_id = transition_branch.branch_id();
11692 let routing_resolved = Arc::new(AtomicBool::new(false));
11693 let mut main_future =
11694 Box::pin(crate::optimization::observability::with_branch_observation(
11695 &main_id,
11696 RuntimeOptimizationKind::BufferedStreamingRouting,
11697 RuntimeCommitBehavior::FinalResponse,
11698 self.generate_buffered_streaming_draft(
11699 processed_input,
11700 Arc::clone(&routing_resolved),
11701 ),
11702 ));
11703 let mut transition_future =
11704 Box::pin(crate::optimization::observability::with_branch_observation(
11705 &transition_id,
11706 RuntimeOptimizationKind::ParallelStateTransition,
11707 RuntimeCommitBehavior::TransitionDecision,
11708 self.select_parallel_transition_candidate(processed_input),
11709 ));
11710 let mut main_pending = true;
11711 let mut transition_pending = true;
11712 let mut main_result: Option<Result<StreamingDraftResult>> = None;
11713 let mut transition_finalized = false;
11714 let mut transition_candidate: Option<TransitionCandidate> = None;
11715 loop {
11716 if let Some(candidate) = transition_candidate.take() {
11717 if self
11718 .approve_transition_target(&candidate.from_state, candidate.target())
11719 .await?
11720 {
11721 drop(main_future);
11723 drop(transition_future);
11724 self.finalize_branch_loss(
11725 &main_id,
11726 RuntimeOptimizationKind::BufferedStreamingRouting,
11727 RuntimeCommitBehavior::FinalResponse,
11728 main_pending,
11729 main_result.as_ref().map(|result| result.is_err()),
11730 );
11731 if !self
11732 .apply_pre_response_transition_candidate(
11733 &candidate,
11734 &HashMap::new(),
11735 processed_input,
11736 )
11737 .await?
11738 {
11739 self.finalize_optional_branch(
11740 &transition_id,
11741 RuntimeOptimizationKind::ParallelStateTransition,
11742 RuntimeCommitBehavior::TransitionDecision,
11743 "discarded",
11744 false,
11745 );
11746 return Ok(None);
11747 }
11748 self.finalize_optional_branch(
11749 &transition_id,
11750 RuntimeOptimizationKind::ParallelStateTransition,
11751 RuntimeCommitBehavior::TransitionDecision,
11752 "committed",
11753 true,
11754 );
11755 let response = self.redispatch_current_state(processed_input).await?;
11756 return Ok(Some((
11757 response.clone(),
11758 vec![StreamChunk::content(response.content)],
11759 )));
11760 }
11761 self.finalize_optional_branch(
11762 &transition_id,
11763 RuntimeOptimizationKind::ParallelStateTransition,
11764 RuntimeCommitBehavior::TransitionDecision,
11765 "discarded",
11766 false,
11767 );
11768 transition_finalized = true;
11769 }
11770 if transition_finalized && !routing_resolved.load(Ordering::SeqCst) {
11776 match self
11777 .resolve_buffered_skill_after_transition(processed_input, &routing_resolved)
11778 .await
11779 {
11780 Ok(Some(candidate)) => {
11781 drop(main_future);
11783 drop(transition_future);
11784 self.finalize_branch_loss(
11785 &main_id,
11786 RuntimeOptimizationKind::BufferedStreamingRouting,
11787 RuntimeCommitBehavior::FinalResponse,
11788 main_pending,
11789 main_result.as_ref().map(|result| result.is_err()),
11790 );
11791 return match self
11792 .commit_winning_skill_candidate(
11793 candidate,
11794 processed_input,
11795 input_context,
11796 )
11797 .await?
11798 {
11799 Some(response) => Ok(Some((
11800 response.clone(),
11801 vec![StreamChunk::content(response.content)],
11802 ))),
11803 None => Ok(None),
11804 };
11805 }
11806 Ok(None) => {}
11807 Err(error) => {
11808 drop(main_future);
11809 drop(transition_future);
11810 self.finalize_branch_loss(
11811 &main_id,
11812 RuntimeOptimizationKind::BufferedStreamingRouting,
11813 RuntimeCommitBehavior::FinalResponse,
11814 main_pending,
11815 main_result.as_ref().map(|result| result.is_err()),
11816 );
11817 return Err(error);
11818 }
11819 }
11820 }
11821 if transition_finalized
11822 && routing_resolved.load(Ordering::SeqCst)
11823 && let Some(result) = main_result.take()
11824 {
11825 let stream_draft = match result {
11826 Ok(stream_draft) => stream_draft,
11827 Err(error) => {
11828 self.finalize_optional_branch(
11829 &main_id,
11830 RuntimeOptimizationKind::BufferedStreamingRouting,
11831 RuntimeCommitBehavior::FinalResponse,
11832 "failed",
11833 false,
11834 );
11835 return Err(error);
11836 }
11837 };
11838 let raw_draft_content = stream_draft.draft.raw_content().to_string();
11839 let buffered_chunks = stream_draft.chunks;
11840 self.finalize_optional_branch(
11841 &main_id,
11842 RuntimeOptimizationKind::BufferedStreamingRouting,
11843 RuntimeCommitBehavior::FinalResponse,
11844 "committed",
11845 true,
11846 );
11847 let response = self
11848 .commit_main_response_draft(
11849 processed_input,
11850 input_context,
11851 stream_draft.draft,
11852 ReasoningMode::None,
11853 false,
11854 )
11855 .await?;
11856 let chunks = if response.content == raw_draft_content {
11857 buffered_chunks
11858 } else {
11859 vec![StreamChunk::content(response.content.clone())]
11860 };
11861 return Ok(Some((response, chunks)));
11862 }
11863 tokio::select! {
11864 result = &mut main_future, if main_pending => {
11865 main_pending = false;
11866 main_branch.transition_to(RuntimeBranchStatus::Completed)?;
11867 main_result = Some(result);
11868 }
11869 result = &mut transition_future, if transition_pending => {
11870 transition_pending = false;
11871 transition_branch.transition_to(RuntimeBranchStatus::Completed)?;
11872 match result {
11873 Ok(ParallelTransitionSelection::Candidate(candidate)) => {
11874 transition_candidate = Some(candidate)
11875 }
11876 Ok(ParallelTransitionSelection::NoMatch) => {
11877 self.finalize_optional_branch(
11878 &transition_id,
11879 RuntimeOptimizationKind::ParallelStateTransition,
11880 RuntimeCommitBehavior::TransitionDecision,
11881 "discarded",
11882 false,
11883 );
11884 transition_finalized = true;
11885 }
11886 Ok(ParallelTransitionSelection::ReservationExhausted) => {
11887 self.finalize_optional_branch(
11888 &transition_id,
11889 RuntimeOptimizationKind::ParallelStateTransition,
11890 RuntimeCommitBehavior::TransitionDecision,
11891 "cancelled",
11892 false,
11893 );
11894 routing_resolved.store(true, Ordering::SeqCst);
11895 self.finalize_branch_loss(
11896 &main_id,
11897 RuntimeOptimizationKind::BufferedStreamingRouting,
11898 RuntimeCommitBehavior::FinalResponse,
11899 main_pending,
11900 main_result.as_ref().map(|result| result.is_err()),
11901 );
11902 return Ok(None);
11903 }
11904 Err(_) => {
11905 self.finalize_optional_branch(
11906 &transition_id,
11907 RuntimeOptimizationKind::ParallelStateTransition,
11908 RuntimeCommitBehavior::TransitionDecision,
11909 "failed",
11910 false,
11911 );
11912 transition_finalized = true;
11913 }
11914 }
11915 }
11916 }
11917 }
11918 }
11919
11920 async fn resolve_buffered_skill_after_transition(
11926 &self,
11927 processed_input: &str,
11928 routing_resolved: &AtomicBool,
11929 ) -> Result<Option<SkillCandidate>> {
11930 let candidate = if self.skill_router.is_some() {
11931 self.select_skill_candidate(processed_input).await?
11932 } else {
11933 None
11934 };
11935 if candidate.is_none() {
11936 routing_resolved.store(true, Ordering::SeqCst);
11937 }
11938 Ok(candidate)
11939 }
11940
11941 fn run_loop_internal_stream<'a>(
11945 &'a self,
11946 input: &'a str,
11947 terminal: RuntimeStreamTerminalSlot,
11948 ) -> Pin<Box<dyn Stream<Item = StreamChunk> + Send + 'a>> {
11949 let include_state_events = self.streaming.include_state_events;
11950
11951 Box::pin(async_stream::stream! {
11952 self.begin_root_turn();
11953 self.pre_turn_session_lifecycle().await;
11955
11956 let input_data = match self.process_input(input).await {
11957 Ok(data) => data,
11958 Err(e) => {
11959 yield StreamChunk::error(e.to_string());
11960 return;
11961 }
11962 };
11963 self.update_active_turn_context(&input_data.content, input_data.context.clone());
11964
11965 for (key, value) in &input_data.context {
11967 let _ = self.context_manager.set(key, value.clone());
11968 }
11969
11970 if input_data.metadata.rejected {
11971 let reason = input_data
11972 .metadata
11973 .rejection_reason
11974 .unwrap_or_else(|| "Input rejected".to_string());
11975 warn!(reason = %reason, "Input rejected (stream)");
11976 let response = AgentResponse::new(&reason);
11979 if let Err(e) = self.finish_turn_if_root(&response).await {
11980 yield StreamChunk::error(e.to_string());
11981 return;
11982 }
11983 yield StreamChunk::content(&reason);
11984 record_runtime_stream_final(&terminal, response);
11985 yield StreamChunk::Done {};
11986 return;
11987 }
11988
11989 let processed_input = &input_data.content;
11990
11991 let streaming_policy = self.runtime_config.optimization.streaming_policy;
11992
11993 if self.runtime_config.optimization.enabled
12000 && !matches!(
12001 streaming_policy,
12002 crate::optimization::StreamingOptimizationPolicy::Disabled
12003 )
12004 {
12005 match self.try_pre_response_transition(processed_input).await {
12006 Ok(Some(response)) => {
12007 yield StreamChunk::content(&response.content);
12008 record_runtime_stream_final(&terminal, response);
12009 yield StreamChunk::Done {};
12010 return;
12011 }
12012 Ok(None) => {}
12013 Err(e) => {
12014 yield StreamChunk::error(e.to_string());
12015 return;
12016 }
12017 }
12018 }
12019
12020 if self.runtime_config.optimization.enabled
12021 && matches!(
12022 streaming_policy,
12023 crate::optimization::StreamingOptimizationPolicy::BufferUntilRoutingDone
12024 )
12025 {
12026 match Box::pin(self.try_buffered_streaming_branches(processed_input, &input_data.context)).await {
12031 Ok(Some((response, chunks))) => {
12032 for chunk in chunks {
12033 yield chunk;
12034 }
12035 record_runtime_stream_final(&terminal, response);
12036 yield StreamChunk::Done {};
12037 return;
12038 }
12039 Ok(None) => {}
12040 Err(e) => {
12041 yield StreamChunk::error(e.to_string());
12042 return;
12043 }
12044 }
12045 }
12046
12047 if let Some(ref sm) = self.state_machine
12049 && let Some(def) = sm.current_definition()
12050 {
12051 let orchestration_result = if let Some(ref delegate_id) = def.delegate {
12052 Some(self.handle_delegated_state(processed_input, delegate_id, &def).await)
12053 } else if let Some(ref concurrent_config) = def.concurrent {
12054 Some(self.handle_concurrent_state(processed_input, concurrent_config).await)
12055 } else if let Some(ref group_chat_config) = def.group_chat {
12056 Some(self.handle_group_chat_state(processed_input, group_chat_config).await)
12057 } else if let Some(ref pipeline_config) = def.pipeline {
12058 Some(self.handle_pipeline_state(processed_input, pipeline_config).await)
12059 } else if let Some(ref handoff_config) = def.handoff {
12060 Some(self.handle_handoff_state(processed_input, handoff_config).await)
12061 } else {
12062 None
12063 };
12064
12065 if let Some(result) = orchestration_result {
12066 match result {
12067 Ok(response) => {
12068 yield StreamChunk::content(&response.content);
12069 record_runtime_stream_final(&terminal, response);
12070 yield StreamChunk::Done {};
12071 }
12072 Err(e) => {
12073 yield StreamChunk::error(e.to_string());
12074 }
12075 }
12076 return;
12077 }
12078 }
12079
12080 match self.try_skill_route(processed_input).await {
12082 Ok(SkillRouteResult::Response { skill_id, content }) => {
12083 if let Err(e) = self.commit_root_user_message(processed_input).await {
12084 yield StreamChunk::error(e.to_string());
12085 return;
12086 }
12087 match self.handle_skill_response(processed_input, &skill_id, content, &input_data.context).await {
12088 Ok(resp) => {
12089 yield StreamChunk::content(&resp.content);
12090 record_runtime_stream_final(&terminal, resp);
12091 yield StreamChunk::Done {};
12092 return;
12093 }
12094 Err(e) => {
12095 yield StreamChunk::error(e.to_string());
12096 return;
12097 }
12098 }
12099 }
12100 Ok(SkillRouteResult::NeedsClarification {
12101 response,
12102 ownership,
12103 }) => {
12104 let admission = match self
12105 .admit_optional_disambiguation_ownership(ownership)
12106 .await
12107 {
12108 Ok(admission) => admission,
12109 Err(e) => {
12110 yield StreamChunk::error(e.to_string());
12111 return;
12112 }
12113 };
12114 if let Err(e) = self.commit_root_user_message(processed_input).await {
12115 yield StreamChunk::error(e.to_string());
12116 return;
12117 }
12118 if Self::skill_clarification_needs_memory_record(&response)
12120 && let Err(e) = self.memory.add_message(ChatMessage::assistant(&response.content)).await
12121 {
12122 yield StreamChunk::error(e.to_string());
12123 return;
12124 }
12125 drop(admission);
12126 if let Err(e) = self.finish_turn_if_root(&response).await {
12127 yield StreamChunk::error(e.to_string());
12128 return;
12129 }
12130 yield StreamChunk::content(&response.content);
12131 record_runtime_stream_final(&terminal, response);
12132 yield StreamChunk::Done {};
12133 return;
12134 }
12135 Ok(SkillRouteResult::NoMatch) => {} Err(e) => {
12137 yield StreamChunk::error(e.to_string());
12138 return;
12139 }
12140 }
12141
12142 let effective_reasoning = self.get_effective_reasoning_config();
12144 let reasoning_mode = match self.determine_reasoning_mode(processed_input).await {
12145 Ok(mode) => mode,
12146 Err(e) => {
12147 yield StreamChunk::error(e.to_string());
12148 return;
12149 }
12150 };
12151 let auto_detected = matches!(effective_reasoning.mode, ReasoningMode::Auto);
12152
12153 info!(
12154 reasoning_mode = ?reasoning_mode,
12155 auto_detected = auto_detected,
12156 "Reasoning mode determined (stream)"
12157 );
12158
12159 if matches!(reasoning_mode, ReasoningMode::PlanAndExecute) {
12161 if let Err(e) = self.commit_root_user_message(processed_input).await {
12162 yield StreamChunk::error(e.to_string());
12163 return;
12164 }
12165 match self.handle_plan_and_execute(processed_input, &input_data.context, auto_detected).await {
12166 Ok(resp) => {
12167 yield StreamChunk::content(&resp.content);
12168 record_runtime_stream_final(&terminal, resp);
12169 yield StreamChunk::Done {};
12170 return;
12171 }
12172 Err(e) => {
12173 yield StreamChunk::error(e.to_string());
12174 return;
12175 }
12176 }
12177 }
12178
12179 if let Err(e) = self.commit_root_user_message(processed_input).await {
12180 yield StreamChunk::error(e.to_string());
12181 return;
12182 }
12183
12184 let llm = match self.get_state_llm() {
12185 Ok(llm) => llm,
12186 Err(e) => {
12187 yield StreamChunk::error(e.to_string());
12188 return;
12189 }
12190 };
12191
12192 let mut iterations = 0u32;
12193 let mut all_tool_calls: Vec<ToolCall> = Vec::new();
12194 let mut thinking_content: Option<String> = None;
12195
12196 loop {
12197 let effective_max = if reasoning_mode != ReasoningMode::None {
12199 let rc = self.get_effective_reasoning_config();
12200 self.max_iterations.min(rc.max_iterations)
12201 } else {
12202 self.max_iterations
12203 };
12204
12205 if iterations >= effective_max {
12206 let err_msg = format!("Max iterations ({}) exceeded", effective_max);
12207 let err = AgentError::Other(err_msg.clone());
12208 self.hooks.on_error(&err).await;
12209 error!(iterations = iterations, "Max iterations exceeded (stream)");
12210 yield StreamChunk::error(err_msg);
12211 return;
12212 }
12213 iterations += 1;
12214 *self.iteration_count.write() = iterations;
12215
12216 debug!(iteration = iterations, max = effective_max, "LLM call (stream)");
12217
12218 let protocol = match self.main_tool_protocol(llm.as_ref(), false).await {
12219 Ok(protocol) => protocol,
12220 Err(e) => {
12221 yield StreamChunk::error(e.to_string());
12222 return;
12223 }
12224 };
12225 let mut messages = match self
12226 .build_messages_internal(true, None, protocol.choice.is_none())
12227 .await
12228 {
12229 Ok(m) => m,
12230 Err(e) => {
12231 yield StreamChunk::error(e.to_string());
12232 return;
12233 }
12234 };
12235 self.inject_reasoning_prompt(&mut messages, &reasoning_mode, iterations == 1);
12236
12237 self.hooks.on_llm_start(&messages).await;
12238 let llm_start = Instant::now();
12239
12240 let buffered_decision = self.main_stream_must_buffer(&reasoning_mode, &protocol);
12241 let content = if buffered_decision {
12242 let response = match self
12246 .complete_main_llm_with_recovery(
12247 Arc::clone(&llm),
12248 &messages,
12249 &protocol,
12250 )
12251 .await
12252 {
12253 Ok(r) => r,
12254 Err(e) => {
12255 yield StreamChunk::error(e.to_string());
12256 return;
12257 }
12258 };
12259 let llm_duration_ms = llm_start.elapsed().as_millis() as u64;
12260 self.hooks.on_llm_complete(&response, llm_duration_ms).await;
12261 response.content.trim().to_string()
12262 } else {
12263 let source = match self
12265 .open_main_stream_with_recovery(Arc::clone(&llm), &messages, &protocol)
12266 .await
12267 {
12268 Ok(source) => source,
12269 Err(e) => {
12270 yield StreamChunk::error(e.to_string());
12271 return;
12272 }
12273 };
12274 let mut accumulated = String::new();
12275 match source {
12276 MainStreamSource::StaticResponse(text) => {
12277 accumulated.push_str(&text);
12278 yield StreamChunk::content(text);
12279 }
12280 MainStreamSource::Stream(mut stream_inner) => {
12281 while let Some(chunk_result) = stream_inner.next().await {
12282 match chunk_result {
12283 Ok(chunk) => {
12284 accumulated.push_str(&chunk.delta);
12285 yield StreamChunk::content(chunk.delta);
12286 }
12287 Err(e) => {
12288 yield StreamChunk::error(e.to_string());
12290 return;
12291 }
12292 }
12293 }
12294 }
12295 }
12296 let llm_duration_ms = llm_start.elapsed().as_millis() as u64;
12297 let llm_response = ai_agents_core::LLMResponse::new(
12299 accumulated.trim(),
12300 ai_agents_core::FinishReason::Stop,
12301 );
12302 self.hooks.on_llm_complete(&llm_response, llm_duration_ms).await;
12303 accumulated.trim().to_string()
12304 };
12305
12306 let parsed_tool_calls = match self.parse_main_tool_calls(&content, &protocol) {
12308 Ok(calls) => calls,
12309 Err(error) => {
12310 yield StreamChunk::error(error.to_string());
12311 return;
12312 }
12313 };
12314 if let Some(tool_calls) = parsed_tool_calls {
12315 let mut events = Vec::new();
12318 let outcome = self
12319 .handle_tool_calls(
12320 processed_input,
12321 &content,
12322 tool_calls,
12323 &mut all_tool_calls,
12324 Some(&mut events),
12325 )
12326 .await;
12327 for chunk in events.drain(..) {
12328 yield chunk;
12329 }
12330 match outcome {
12331 Ok(ToolCallOutcome::Continue) | Ok(ToolCallOutcome::TransitionFired) => continue,
12332 Ok(ToolCallOutcome::Rejected(response)) => {
12333 if let Err(finalize_error) = self.finish_turn_if_root(&response).await {
12334 yield StreamChunk::error(finalize_error.to_string());
12335 return;
12336 }
12337 let legacy_error = response.content.clone();
12338 record_runtime_stream_final(&terminal, response);
12339 yield StreamChunk::error(legacy_error);
12340 yield StreamChunk::Done {};
12341 return;
12342 }
12343 Err(e) => {
12344 yield StreamChunk::error(e.to_string());
12345 return;
12346 }
12347 }
12348 }
12349
12350 let (extracted_thinking, answer) = self.extract_thinking(&content);
12352 if extracted_thinking.is_some() {
12353 thinking_content = extracted_thinking;
12354 }
12355
12356 let output_data = match self.process_output(&answer, &input_data.context).await {
12357 Ok(d) => d,
12358 Err(e) => {
12359 yield StreamChunk::error(e.to_string());
12360 return;
12361 }
12362 };
12363
12364 let final_content = if output_data.metadata.rejected {
12365 output_data
12366 .metadata
12367 .rejection_reason
12368 .unwrap_or_else(|| answer.to_string())
12369 } else {
12370 output_data.content
12371 };
12372
12373 let (final_content, reflection_metadata) = match self
12375 .run_reflection(&*llm, processed_input, final_content)
12376 .await
12377 {
12378 Ok(r) => r,
12379 Err(e) => {
12380 yield StreamChunk::error(e.to_string());
12381 return;
12382 }
12383 };
12384
12385 let final_content = self.format_response_with_thinking(
12386 thinking_content.as_deref(),
12387 &final_content,
12388 );
12389
12390 if buffered_decision {
12392 yield StreamChunk::content(&final_content);
12393 }
12394
12395 let post_result = match self
12399 .post_loop_processing(processed_input, final_content)
12400 .await
12401 {
12402 Ok(r) => r,
12403 Err(e) => {
12404 yield StreamChunk::error(e.to_string());
12405 return;
12406 }
12407 };
12408
12409 let applied = match self.apply_post_loop_result(processed_input, post_result).await {
12410 Ok(applied) => applied,
12411 Err(e) => {
12412 yield StreamChunk::error(e.to_string());
12413 return;
12414 }
12415 };
12416
12417 if applied.transitioned {
12418 if include_state_events
12419 && let Some(state) = self.current_state()
12420 {
12421 yield StreamChunk::state_transition(None, state);
12422 }
12423 if applied.regenerated {
12429 yield StreamChunk::content(&applied.content);
12430 }
12431 }
12432 let final_content = applied.content;
12433
12434 let final_response = self.build_agent_response(AgentResponseParts {
12436 content: final_content,
12437 all_tool_calls,
12438 reasoning_mode,
12439 auto_detected,
12440 iterations,
12441 thinking: thinking_content,
12442 reflection_metadata,
12443 });
12444 if let Err(e) = self.finish_turn_if_root(&final_response).await {
12445 yield StreamChunk::error(e.to_string());
12446 return;
12447 }
12448
12449 record_runtime_stream_final(&terminal, final_response);
12450 yield StreamChunk::Done {};
12451 return;
12452 }
12453 })
12454 }
12455
12456 fn run_loop_stream<'a>(
12459 &'a self,
12460 input: &'a str,
12461 terminal: RuntimeStreamTerminalSlot,
12462 ) -> Pin<Box<dyn Stream<Item = StreamChunk> + Send + 'a>> {
12463 Box::pin(async_stream::stream! {
12464 self.begin_root_turn();
12465 let _root_cleanup = RootTurnCleanup::new(self);
12466 self.hooks.on_message_received(input).await;
12467
12468 if let Err(e) = self.prepare_turn_context().await {
12470 yield StreamChunk::error(e.to_string());
12471 return;
12472 }
12473
12474 self.clear_disambiguation_context();
12476
12477 let input_to_run = match self.resolve_disambiguation(input).await {
12481 Err(e) => {
12482 yield StreamChunk::error(e.to_string());
12483 return;
12484 }
12485 Ok(DisambiguationDispatch::Terminal(response)) => {
12486 yield StreamChunk::content(&response.content);
12487 record_runtime_stream_final(&terminal, response);
12488 yield StreamChunk::Done {};
12489 return;
12490 }
12491 Ok(DisambiguationDispatch::RecheckSkill {
12492 skill_id,
12493 enriched_input,
12494 disambiguation_epoch,
12495 state_generation,
12496 }) => {
12497 match self
12498 .recheck_skill_disambiguation(
12499 &skill_id,
12500 &enriched_input,
12501 disambiguation_epoch,
12502 state_generation,
12503 )
12504 .await
12505 {
12506 Ok(resp) => {
12507 yield StreamChunk::content(&resp.content);
12508 record_runtime_stream_final(&terminal, resp);
12509 yield StreamChunk::Done {};
12510 return;
12511 }
12512 Err(e) => {
12513 yield StreamChunk::error(e.to_string());
12514 return;
12515 }
12516 }
12517 }
12518 Ok(DisambiguationDispatch::Proceed(input)) => input,
12519 };
12520
12521 let mut inner = self.run_loop_internal_stream(&input_to_run, Arc::clone(&terminal));
12522 while let Some(chunk) = inner.next().await {
12523 yield chunk;
12524 }
12525 })
12526 }
12527
12528 pub fn info(&self) -> AgentInfo {
12529 self.info.clone()
12530 }
12531
12532 pub fn skills(&self) -> &[SkillDefinition] {
12533 &self.skills
12534 }
12535
12536 async fn reset_runtime_state(&self) -> Result<()> {
12538 let _admission = self.disambiguation_admission.write().await;
12539 if self.state_transition_reserved.load(Ordering::SeqCst) {
12540 return Err(AgentError::Other(
12541 "Cannot reset while a state transition is in progress".to_string(),
12542 ));
12543 }
12544 self.disambiguation_epoch.fetch_add(1, Ordering::SeqCst);
12545 *self.pending_skill_id.write() = None;
12546 if let Some(disambiguator) = self.disambiguation_manager.as_ref() {
12547 disambiguator.clear_pending().await;
12548 }
12549 self.memory.clear().await?;
12550 self.active_native_exchanges.write().clear();
12551 *self.iteration_count.write() = 0;
12552 self.tool_call_history.write().clear();
12553 if let Some(ref sm) = self.state_machine {
12554 sm.reset();
12555 }
12556 Ok(())
12557 }
12558
12559 pub async fn reset(&self) -> Result<()> {
12561 self.reset_runtime_state().await
12562 }
12563
12564 pub fn max_context_tokens(&self) -> u32 {
12565 self.max_context_tokens
12566 }
12567
12568 pub fn llm_registry(&self) -> &Arc<LLMRegistry> {
12569 &self.llm_registry
12570 }
12571
12572 pub fn state_machine(&self) -> Option<&Arc<StateMachine>> {
12573 self.state_machine.as_ref()
12574 }
12575
12576 pub fn context_manager(&self) -> &Arc<ContextManager> {
12577 &self.context_manager
12578 }
12579
12580 pub fn tool_call_history(&self) -> Vec<ToolCallRecord> {
12581 self.tool_call_history.read().clone()
12582 }
12583
12584 pub fn memory_token_budget(&self) -> Option<&MemoryTokenBudget> {
12585 self.memory_token_budget.as_ref()
12586 }
12587
12588 pub fn parallel_tools_config(&self) -> &ParallelToolsConfig {
12589 &self.parallel_tools
12590 }
12591
12592 pub fn streaming_config(&self) -> &StreamingConfig {
12593 &self.streaming
12594 }
12595
12596 pub fn hooks(&self) -> &Arc<dyn AgentHooks> {
12597 &self.hooks
12598 }
12599
12600 pub fn hitl_engine(&self) -> Option<&HITLEngine> {
12601 self.hitl_engine.as_ref()
12602 }
12603
12604 pub fn approval_handler(&self) -> &Arc<dyn ApprovalHandler> {
12605 &self.approval_handler
12606 }
12607
12608 fn build_hitl_language_context(&self) -> HashMap<String, Value> {
12610 let mut ctx = HashMap::new();
12611 for key in &["user.language", "input.detected.language", "language"] {
12612 if let Some(val) = self.context_manager.get(key) {
12613 ctx.insert(key.to_string(), val);
12614 }
12615 }
12616 ctx
12617 }
12618
12619 async fn request_hitl_approval(&self, check_result: HITLCheckResult) -> Result<ApprovalResult> {
12621 let Some(request) = check_result.into_request() else {
12622 return Ok(ApprovalResult::Approved);
12623 };
12624
12625 self.hooks.on_approval_requested(&request).await;
12626
12627 let timeout = request.timeout;
12628
12629 let raw_result = if let Some(duration) = timeout {
12630 match tokio::time::timeout(
12631 duration,
12632 self.approval_handler.request_approval(request.clone()),
12633 )
12634 .await
12635 {
12636 Ok(result) => result,
12637 Err(_) => ApprovalResult::timeout(),
12638 }
12639 } else {
12640 self.approval_handler
12641 .request_approval(request.clone())
12642 .await
12643 };
12644
12645 self.hooks
12646 .on_approval_result(&request.id, &raw_result)
12647 .await;
12648
12649 let (outcome, effective_result): (ApprovalResolvedOutcome, Result<ApprovalResult>) =
12650 match &raw_result {
12651 ApprovalResult::Approved => (
12652 ApprovalResolvedOutcome::Approved,
12653 Ok(ApprovalResult::Approved),
12654 ),
12655 ApprovalResult::Rejected { reason } => (
12656 ApprovalResolvedOutcome::Rejected {
12657 reason: reason.clone(),
12658 },
12659 Ok(ApprovalResult::Rejected {
12660 reason: reason.clone(),
12661 }),
12662 ),
12663 ApprovalResult::Modified { changes } => (
12664 ApprovalResolvedOutcome::Modified {
12665 changes: changes.clone(),
12666 },
12667 Ok(ApprovalResult::Modified {
12668 changes: changes.clone(),
12669 }),
12670 ),
12671 ApprovalResult::Timeout => {
12672 if let Some(ref engine) = self.hitl_engine {
12673 match engine.config().on_timeout {
12674 TimeoutAction::Approve => (
12675 ApprovalResolvedOutcome::Approved,
12676 Ok(ApprovalResult::Approved),
12677 ),
12678 TimeoutAction::Reject => {
12679 let reason = Some("Timeout".to_string());
12680 (
12681 ApprovalResolvedOutcome::Rejected {
12682 reason: reason.clone(),
12683 },
12684 Ok(ApprovalResult::Rejected { reason }),
12685 )
12686 }
12687 TimeoutAction::Error => {
12688 let message = "HITL approval timeout".to_string();
12689 (
12690 ApprovalResolvedOutcome::Error {
12691 message: message.clone(),
12692 },
12693 Err(AgentError::Other(message)),
12694 )
12695 }
12696 }
12697 } else {
12698 let reason = Some("Timeout (no engine)".to_string());
12699 (
12700 ApprovalResolvedOutcome::Rejected {
12701 reason: reason.clone(),
12702 },
12703 Ok(ApprovalResult::Rejected { reason }),
12704 )
12705 }
12706 }
12707 };
12708
12709 self.hooks
12710 .on_approval_resolved(&request, &raw_result, &outcome)
12711 .await;
12712
12713 effective_result
12714 }
12715
12716 pub async fn check_state_hitl(&self, from: Option<&str>, to: &str) -> Result<bool> {
12717 if let Some(ref hitl_engine) = self.hitl_engine {
12718 let hitl_lang_ctx = self.build_hitl_language_context();
12719 let check_result = self
12720 .observe_purpose(
12721 ObservationPurpose::HitlLocalization,
12722 hitl_engine.check_state_transition_with_localization(
12723 from,
12724 to,
12725 &hitl_lang_ctx,
12726 self.approval_handler.as_ref(),
12727 Some(&self.llm_registry),
12728 ),
12729 )
12730 .await?;
12731 if check_result.is_required() {
12732 let result = self.request_hitl_approval(check_result).await?;
12733 return Ok(matches!(
12734 result,
12735 ApprovalResult::Approved | ApprovalResult::Modified { .. }
12736 ));
12737 }
12738 }
12739 Ok(true)
12740 }
12741
12742 async fn execute_tools_parallel(
12744 &self,
12745 tool_calls: &[ToolCall],
12746 ) -> Vec<(String, Result<String>)> {
12747 let can_run_parallel = tool_calls.iter().all(|tc| {
12748 self.tools
12749 .resolve(&tc.name)
12750 .map(|resolved| resolved.tool.classify_call(&tc.arguments).concurrency_safe)
12751 .unwrap_or(false)
12752 });
12753
12754 if !self.parallel_tools.enabled || tool_calls.len() <= 1 || !can_run_parallel {
12755 let mut results = Vec::new();
12756 for tc in tool_calls {
12757 let result = self
12758 .observe_purpose(
12759 current_observation_context()
12760 .map(|context| context.purpose)
12761 .unwrap_or_default(),
12762 self.execute_tool_smart(tc),
12763 )
12764 .await;
12765 results.push((tc.id.clone(), result));
12766 }
12767 return results;
12768 }
12769
12770 let chunks: Vec<_> = tool_calls
12771 .chunks(self.parallel_tools.max_parallel)
12772 .collect();
12773
12774 let mut all_results = Vec::new();
12775
12776 for chunk in chunks {
12777 let futures: Vec<_> = chunk
12778 .iter()
12779 .map(|tc| {
12780 let tc = tc.clone();
12781 async move {
12782 let result = self.execute_tool_smart(&tc).await;
12783 (tc.id.clone(), result)
12784 }
12785 })
12786 .collect();
12787
12788 let results = futures::future::join_all(futures).await;
12789 all_results.extend(results);
12790 }
12791
12792 all_results
12793 }
12794
12795 pub async fn chat_stream<'a>(
12799 &'a self,
12800 input: &'a str,
12801 ) -> Result<Pin<Box<dyn Stream<Item = StreamChunk> + Send + 'a>>> {
12802 let RootTurnAdmission {
12803 guard: root_turn_guard,
12804 identity_stack,
12805 } = self.acquire_root_turn().await?;
12806 scope_runtime_gate_identity_stack(&identity_stack, self.init_storage()).await?;
12810 info!(input_len = input.len(), "Starting streaming chat");
12811 let terminal = new_runtime_stream_terminal_slot();
12812 let inner = self.run_loop_stream(input, terminal);
12813 let observation_context = self.build_observation_context(None);
12814 let stream: Pin<Box<dyn Stream<Item = StreamChunk> + Send + 'a>> =
12815 Box::pin(async_stream::stream! {
12816 let mut root_turn_guard = Some(root_turn_guard);
12817 let mut inner = inner;
12818 loop {
12819 let next = scope_runtime_gate_identity_stack(&identity_stack, async {
12820 if let Some(context) = observation_context.as_ref() {
12821 with_observation_context(context.clone(), inner.next()).await
12822 } else {
12823 inner.next().await
12824 }
12825 })
12826 .await;
12827 match next {
12828 Some(StreamChunk::Done {}) => {
12829 while scope_runtime_gate_identity_stack(&identity_stack, async {
12830 if let Some(context) = observation_context.as_ref() {
12831 with_observation_context(context.clone(), inner.next())
12832 .await
12833 .is_some()
12834 } else {
12835 inner.next().await.is_some()
12836 }
12837 })
12838 .await
12839 {}
12840 if observation_context.is_some() {
12841 scope_runtime_gate_identity_stack(
12842 &identity_stack,
12843 self.export_observability_if_configured(),
12844 )
12845 .await;
12846 }
12847 drop(root_turn_guard.take());
12848 yield StreamChunk::Done {};
12849 return;
12850 }
12851 Some(chunk) => yield chunk,
12852 None => {
12853 if observation_context.is_some() {
12854 scope_runtime_gate_identity_stack(
12855 &identity_stack,
12856 self.export_observability_if_configured(),
12857 )
12858 .await;
12859 }
12860 drop(root_turn_guard.take());
12861 return;
12862 }
12863 }
12864 }
12865 });
12866 Ok(stream)
12867 }
12868
12869 pub async fn chat_stream_events<'a>(
12873 &'a self,
12874 input: &'a str,
12875 ) -> Result<Pin<Box<dyn Stream<Item = AgentStreamEvent> + Send + 'a>>> {
12876 let RootTurnAdmission {
12877 guard,
12878 identity_stack,
12879 } = self.acquire_root_turn().await?;
12880 scope_runtime_gate_identity_stack(&identity_stack, self.init_storage()).await?;
12884 info!(input_len = input.len(), "Starting streaming chat events");
12885 let terminal = new_runtime_stream_terminal_slot();
12886 let inner = self.run_loop_stream(input, Arc::clone(&terminal));
12887 let observation_context = self.build_observation_context(None);
12888 Ok(self.drive_event_stream(
12889 inner,
12890 terminal,
12891 guard,
12892 identity_stack,
12893 observation_context,
12894 None,
12895 ))
12896 }
12897
12898 pub async fn chat_stream_events_with_actor_context<'a>(
12904 &'a self,
12905 input: &'a str,
12906 actor_context: crate::TurnActorContext,
12907 ) -> Result<Pin<Box<dyn Stream<Item = AgentStreamEvent> + Send + 'a>>> {
12908 let RootTurnAdmission {
12909 guard,
12910 identity_stack,
12911 } = self.acquire_root_turn().await?;
12912 scope_runtime_gate_identity_stack(&identity_stack, self.init_storage()).await?;
12913 info!(
12914 input_len = input.len(),
12915 "Starting streaming chat events with actor context"
12916 );
12917 let actor_id = actor_context.effective_actor_id().map(str::to_string);
12918 let terminal = new_runtime_stream_terminal_slot();
12919 let inner = self.run_loop_stream(input, Arc::clone(&terminal));
12920 let observation_context = self.build_observation_context(actor_id);
12921 Ok(self.drive_event_stream(
12922 inner,
12923 terminal,
12924 guard,
12925 identity_stack,
12926 observation_context,
12927 Some(actor_context),
12928 ))
12929 }
12930
12931 fn drive_event_stream<'a>(
12940 &'a self,
12941 mut inner: Pin<Box<dyn Stream<Item = StreamChunk> + Send + 'a>>,
12942 terminal: RuntimeStreamTerminalSlot,
12943 root_turn_guard: tokio::sync::OwnedMutexGuard<()>,
12944 identity_stack: RootTurnGateIdentityStack,
12945 observation_context: Option<SpanContext>,
12946 actor_context: Option<crate::TurnActorContext>,
12947 ) -> Pin<Box<dyn Stream<Item = AgentStreamEvent> + Send + 'a>> {
12948 Box::pin(async_stream::stream! {
12949 let mut root_turn_guard = Some(root_turn_guard);
12950 loop {
12951 let next = poll_scoped_chunk(
12952 &mut inner,
12953 &identity_stack,
12954 observation_context.as_ref(),
12955 actor_context.as_ref(),
12956 )
12957 .await;
12958 match next {
12959 Some(StreamChunk::Done {}) => {
12960 let terminal_event = { terminal.write().take() };
12961 if let Some(response) = terminal_event {
12962 while poll_scoped_chunk(
12963 &mut inner,
12964 &identity_stack,
12965 observation_context.as_ref(),
12966 actor_context.as_ref(),
12967 )
12968 .await
12969 .is_some()
12970 {}
12971 if observation_context.is_some() {
12972 scope_runtime_gate_identity_stack(
12973 &identity_stack,
12974 self.export_observability_if_configured(),
12975 )
12976 .await;
12977 }
12978 drop(root_turn_guard.take());
12979 yield AgentStreamEvent::Final(response);
12980 return;
12981 }
12982 }
12983 Some(StreamChunk::Error { message }) => {
12984 let finalized = { terminal.read().is_some() };
12985 if finalized {
12986 continue;
12987 }
12988 while poll_scoped_chunk(
12989 &mut inner,
12990 &identity_stack,
12991 observation_context.as_ref(),
12992 actor_context.as_ref(),
12993 )
12994 .await
12995 .is_some()
12996 {}
12997 if observation_context.is_some() {
12998 scope_runtime_gate_identity_stack(
12999 &identity_stack,
13000 self.export_observability_if_configured(),
13001 )
13002 .await;
13003 }
13004 drop(root_turn_guard.take());
13005 yield AgentStreamEvent::Chunk(StreamChunk::Error { message });
13006 return;
13007 }
13008 Some(chunk) => yield AgentStreamEvent::Chunk(chunk),
13009 None => {
13010 if observation_context.is_some() {
13011 scope_runtime_gate_identity_stack(
13012 &identity_stack,
13013 self.export_observability_if_configured(),
13014 )
13015 .await;
13016 }
13017 drop(root_turn_guard.take());
13018 return;
13019 }
13020 }
13021 }
13022 })
13023 }
13024}
13025
13026async fn poll_scoped_chunk<'a>(
13032 inner: &mut Pin<Box<dyn Stream<Item = StreamChunk> + Send + 'a>>,
13033 identity_stack: &RootTurnGateIdentityStack,
13034 observation_context: Option<&SpanContext>,
13035 actor_context: Option<&crate::TurnActorContext>,
13036) -> Option<StreamChunk> {
13037 scope_runtime_gate_identity_stack(identity_stack, async {
13038 let next = inner.next();
13039 match (observation_context, actor_context) {
13040 (Some(observation), Some(actor)) => {
13041 with_observation_context(
13042 observation.clone(),
13043 scope_actor_context(actor.clone(), next),
13044 )
13045 .await
13046 }
13047 (Some(observation), None) => with_observation_context(observation.clone(), next).await,
13048 (None, Some(actor)) => scope_actor_context(actor.clone(), next).await,
13049 (None, None) => next.await,
13050 }
13051 })
13052 .await
13053}
13054
13055#[async_trait]
13056impl ToolInvoker for RuntimeAgent {
13057 async fn invoke_tool(&self, request: ToolExecutionRequest) -> Result<ToolExecutionRecord> {
13058 self.execute_tool_record(request).await
13059 }
13060}
13061
13062#[async_trait]
13063impl Agent for RuntimeAgent {
13064 async fn chat(&self, input: &str) -> Result<AgentResponse> {
13066 let RootTurnAdmission {
13067 guard,
13068 identity_stack,
13069 } = self.acquire_root_turn().await?;
13070 let result = scope_runtime_gate_identity_stack(&identity_stack, async {
13071 let result = if let Some(context) = self.build_observation_context(None) {
13072 with_observation_context(context, self.run_loop(input)).await
13073 } else {
13074 self.run_loop(input).await
13075 };
13076 self.export_observability_if_configured().await;
13077 result
13078 })
13079 .await;
13080 drop(guard);
13081 result
13082 }
13083
13084 fn info(&self) -> AgentInfo {
13085 self.info.clone()
13086 }
13087
13088 async fn reset(&self) -> Result<()> {
13090 self.reset_runtime_state().await
13091 }
13092}
13093
13094fn background_maintenance_tags(
13104 label: &str,
13105 stage: &str,
13106 reason: Option<&str>,
13107 policy: Option<&crate::optimization::config::MaintenanceTaskPolicy>,
13108) -> HashMap<String, String> {
13109 let mut tags = HashMap::new();
13110 tags.insert("runtime.background".to_string(), "true".to_string());
13111 tags.insert("runtime.maintenance".to_string(), label.to_string());
13112 tags.insert("runtime.maintenance_stage".to_string(), stage.to_string());
13113 if let Some(policy) = policy {
13114 tags.insert(
13115 "runtime.await_before_next_turn".to_string(),
13116 await_before_next_turn_label(policy.await_before_next_turn).to_string(),
13117 );
13118 tags.insert(
13119 "runtime.maintenance_mode".to_string(),
13120 maintenance_mode_label(policy.mode).to_string(),
13121 );
13122 }
13123 if let Some(reason) = reason {
13124 tags.insert("runtime.reason".to_string(), reason.to_string());
13125 }
13126 tags
13127}
13128
13129fn await_before_next_turn_label(policy: AwaitBeforeNextTurn) -> &'static str {
13130 match policy {
13131 AwaitBeforeNextTurn::Never => "never",
13132 AwaitBeforeNextTurn::SameActor => "same_actor",
13133 AwaitBeforeNextTurn::Always => "always",
13134 }
13135}
13136
13137fn maintenance_mode_label(mode: MaintenanceMode) -> &'static str {
13138 match mode {
13139 MaintenanceMode::InlineSerial => "inline_serial",
13140 MaintenanceMode::InlineParallel => "inline_parallel",
13141 MaintenanceMode::Background => "background",
13142 }
13143}
13144
13145fn record_background_maintenance_event(
13147 manager: Option<&Arc<ObservabilityManager>>,
13148 label: &str,
13149 status: EventStatus,
13150 duration_ms: u64,
13151 stage: &str,
13152 reason: Option<String>,
13153 policy: Option<&crate::optimization::config::MaintenanceTaskPolicy>,
13154) {
13155 if let Some(manager) = manager {
13156 manager.record_lifecycle_event(
13157 EventType::MemoryOperation {
13158 operation: format!("{}_background_{}", label, stage),
13159 },
13160 ObservationPurpose::Other(format!("{}_maintenance", label)),
13161 status,
13162 duration_ms,
13163 background_maintenance_tags(label, stage, reason.as_deref(), policy),
13164 None,
13165 );
13166 }
13167}
13168
13169fn effective_maintenance_mode(mode: MaintenanceMode, force_parallel: bool) -> MaintenanceMode {
13170 if force_parallel && matches!(mode, MaintenanceMode::InlineSerial) {
13171 MaintenanceMode::InlineParallel
13172 } else {
13173 mode
13174 }
13175}
13176
13177fn observation_purpose_for_process(hint: ProcessPurposeHint) -> ObservationPurpose {
13178 match hint {
13179 ProcessPurposeHint::Detect => ObservationPurpose::ProcessDetect,
13180 ProcessPurposeHint::Extract => ObservationPurpose::ProcessExtract,
13181 ProcessPurposeHint::Validate => ObservationPurpose::ProcessValidate,
13182 ProcessPurposeHint::Transform | ProcessPurposeHint::Other => {
13183 ObservationPurpose::ProcessTransform
13184 }
13185 }
13186}
13187
13188fn new_tool_resource_locks() -> ToolResourceLocks {
13189 Arc::new(RwLock::new(HashMap::new()))
13190}
13191
13192fn tool_resource_lock_keys(
13197 _canonical_id: &str,
13198 args: &Value,
13199 bindings: &ai_agents_core::ToolPolicyBindings,
13200 classification: &ai_agents_core::ToolCallClassification,
13201) -> Vec<String> {
13202 if classification.concurrency_safe {
13203 return Vec::new();
13204 }
13205
13206 let mut keys = Vec::new();
13207 let mut has_path_resource = false;
13208 for binding in &bindings.path_fields {
13209 let value = value_at_argument_path(args, &binding.field)
13210 .cloned()
13211 .or_else(|| {
13212 binding
13213 .default_path
13214 .as_ref()
13215 .map(|path| Value::String(path.clone()))
13216 });
13217 if let Some(value) = value {
13218 collect_resource_strings(&value, |_| {
13219 has_path_resource = true;
13220 });
13221 }
13222 }
13223 for binding in &bindings.domain_fields {
13224 if let Some(value) = value_at_argument_path(args, &binding.field) {
13225 collect_resource_strings(value, |domain| {
13226 let normalized = if binding.is_url {
13227 normalized_url_resource_key(domain)
13228 } else {
13229 domain.trim().trim_end_matches('.').to_ascii_lowercase()
13230 };
13231 keys.push(format!("domain:{}", normalized));
13232 });
13233 }
13234 }
13235 for binding in &bindings.command_fields {
13236 if !matches!(binding.kind, ai_agents_core::CommandBindingKind::Cwd) {
13237 continue;
13238 }
13239 if let Some(value) = value_at_argument_path(args, &binding.field) {
13240 collect_resource_strings(value, |_| {
13241 has_path_resource = true;
13242 });
13243 }
13244 }
13245 if has_path_resource {
13246 keys.push("path-mutation:global".to_string());
13247 }
13248 if keys.is_empty() {
13249 keys.push("side-effect:unbound".to_string());
13250 }
13251 keys.sort();
13252 keys.dedup();
13253 keys
13254}
13255
13256fn value_at_argument_path<'a>(value: &'a Value, field: &str) -> Option<&'a Value> {
13257 let mut current = value;
13258 for segment in field.split('.') {
13259 if segment.is_empty() {
13260 return None;
13261 }
13262 current = current.get(segment)?;
13263 }
13264 Some(current)
13265}
13266
13267fn collect_resource_strings(value: &Value, mut collect: impl FnMut(&str)) {
13268 match value {
13269 Value::String(value) => collect(value),
13270 Value::Array(values) => {
13271 for value in values {
13272 if let Some(value) = value.as_str() {
13273 collect(value);
13274 }
13275 }
13276 }
13277 _ => {}
13278 }
13279}
13280
13281fn normalized_url_resource_key(value: &str) -> String {
13282 let value = value.trim();
13283 let Some((scheme, remainder)) = value.split_once("://") else {
13284 return value.to_ascii_lowercase();
13285 };
13286 let authority_end = remainder.find(['/', '?', '#']).unwrap_or(remainder.len());
13287 let (authority, suffix) = remainder.split_at(authority_end);
13288 format!(
13289 "{}://{}{}",
13290 scheme.to_ascii_lowercase(),
13291 authority.to_ascii_lowercase(),
13292 suffix
13293 )
13294}
13295
13296fn render_concurrent_template(
13297 template: &str,
13298 user_input: &str,
13299 context_values: &std::collections::HashMap<String, serde_json::Value>,
13300) -> Result<String> {
13301 let mut env = minijinja::Environment::new();
13302 env.add_template("concurrent", template)
13303 .map_err(|e| AgentError::Other(format!("Concurrent template parse error: {}", e)))?;
13304
13305 let mut ctx = std::collections::BTreeMap::new();
13306 ctx.insert("user_input".to_string(), minijinja::Value::from(user_input));
13307
13308 let context_obj = minijinja::Value::from_serialize(context_values);
13310 ctx.insert("context".to_string(), context_obj);
13311
13312 let tmpl = env
13313 .get_template("concurrent")
13314 .map_err(|e| AgentError::Other(format!("Concurrent template error: {}", e)))?;
13315
13316 tmpl.render(minijinja::Value::from_serialize(&ctx))
13317 .map_err(|e| AgentError::Other(format!("Concurrent template render error: {}", e)))
13318}
13319
13320#[cfg(test)]
13321mod tests {
13322 use super::*;
13323 use crate::AgentBuilder;
13324 use ai_agents_core::{LLMChunk, LLMConfig, LLMError, LLMFeature, Tool};
13325 use ai_agents_llm::mock::MockLLMProvider;
13326 use ai_agents_skills::{SkillDefinition, SkillStep};
13327 use ai_agents_tools::{
13328 CalculatorTool, CopyPathTool, DeletePathTool, FileWriteTool, MovePathTool, ToolAliases,
13329 ToolDescriptor, ToolProvider, ToolProviderError, ToolProviderType, WebFetchResolver,
13330 WebFetchTool, WebFetchTransport, WebFetchTransportRequest, WebFetchTransportResponse,
13331 };
13332
13333 fn mock_with_response(response: &str) -> MockLLMProvider {
13334 let mut mock = MockLLMProvider::new("test");
13335 mock.set_response(response);
13336 mock
13337 }
13338
13339 fn mock_with_responses(responses: Vec<&str>) -> MockLLMProvider {
13340 let mut mock = MockLLMProvider::new("test");
13341 mock.set_responses(responses.into_iter().map(String::from).collect(), true);
13342 mock
13343 }
13344
13345 async fn collect_stream_events(
13347 agent: &RuntimeAgent,
13348 input: &str,
13349 ) -> (String, Vec<StreamChunk>, Option<AgentResponse>) {
13350 use futures::StreamExt;
13351 let mut events = agent.chat_stream_events(input).await.expect("stream opens");
13352 let mut content = String::new();
13353 let mut chunks = Vec::new();
13354 let mut final_response = None;
13355 while let Some(event) = events.next().await {
13356 match event {
13357 AgentStreamEvent::Chunk(chunk) => {
13358 if let StreamChunk::Content { text } = &chunk {
13359 content.push_str(text);
13360 }
13361 chunks.push(chunk);
13362 }
13363 AgentStreamEvent::Final(response) => final_response = Some(response),
13364 }
13365 }
13366 (content, chunks, final_response)
13367 }
13368
13369 fn metadata_keys(response: &AgentResponse) -> std::collections::BTreeSet<String> {
13370 response
13371 .metadata
13372 .as_ref()
13373 .map(|m| m.keys().cloned().collect())
13374 .unwrap_or_default()
13375 }
13376
13377 async fn assert_blocking_streaming_parity<F>(
13380 build: F,
13381 input: &str,
13382 ) -> (AgentResponse, AgentResponse, Vec<StreamChunk>)
13383 where
13384 F: Fn() -> RuntimeAgent,
13385 {
13386 let blocking_agent = build();
13387 let streaming_agent = build();
13388
13389 let blocking = blocking_agent
13390 .chat(input)
13391 .await
13392 .expect("blocking chat succeeds");
13393 let (_, chunks, final_response) = collect_stream_events(&streaming_agent, input).await;
13394 let streamed = final_response.expect("streaming must emit Final when blocking succeeds");
13395
13396 assert_eq!(
13397 blocking.content, streamed.content,
13398 "committed content differs"
13399 );
13400 assert_eq!(
13401 metadata_keys(&blocking),
13402 metadata_keys(&streamed),
13403 "metadata key sets differ"
13404 );
13405 assert_eq!(
13406 blocking.tool_calls.as_ref().map(Vec::len),
13407 streamed.tool_calls.as_ref().map(Vec::len),
13408 "tool call counts differ"
13409 );
13410 assert_eq!(
13411 blocking_agent.current_state(),
13412 streaming_agent.current_state(),
13413 "final states differ"
13414 );
13415 (blocking, streamed, chunks)
13416 }
13417
13418 fn signed_calculator_response(
13419 exchange_id: &str,
13420 call_id: &str,
13421 expression: &str,
13422 ) -> LLMResponse {
13423 let call = ToolCall {
13424 id: call_id.to_string(),
13425 name: "calculator".to_string(),
13426 arguments: serde_json::json!({"expression": expression}),
13427 };
13428 let state = ai_agents_core::NativeProviderState::new(
13429 exchange_id,
13430 "fixture",
13431 "native-tools",
13432 ai_agents_core::NativeProviderTarget::new("https://fixture.invalid/", "fixture-model")
13433 .unwrap(),
13434 serde_json::json!({
13435 "role": "model",
13436 "parts": [{
13437 "functionCall": {"name": "calculator", "args": {"expression": expression}},
13438 "thoughtSignature": format!("signature-{exchange_id}")
13439 }]
13440 }),
13441 vec![ai_agents_core::NativeCallBinding::new(call_id, 0).unwrap()],
13442 )
13443 .unwrap();
13444 LLMResponse::new("", FinishReason::ToolCall)
13445 .with_provider_state(state)
13446 .unwrap()
13447 .with_tool_calls(vec![call])
13448 .unwrap()
13449 }
13450
13451 struct TerminalHistoryProvider {
13452 calls: Arc<std::sync::atomic::AtomicU32>,
13453 }
13454
13455 struct DroppingSignedAssistantMemory {
13456 messages: RwLock<Vec<ChatMessage>>,
13457 }
13458
13459 struct DroppingEarlierSequentialMemory {
13460 messages: RwLock<Vec<ChatMessage>>,
13461 signed_seen: std::sync::atomic::AtomicUsize,
13462 }
13463
13464 #[async_trait]
13465 impl ai_agents_core::Memory for DroppingSignedAssistantMemory {
13466 async fn add_message(&self, message: ChatMessage) -> Result<()> {
13467 let signed = message.role == ai_agents_core::Role::Assistant
13468 && ai_agents_core::decode_native_tool_call_markers(&message.content)
13469 .map_err(|error| AgentError::LLM(error.to_string()))?
13470 .is_some_and(|batch| batch.provider_state().is_some());
13471 if !signed {
13472 self.messages.write().push(message);
13473 }
13474 Ok(())
13475 }
13476
13477 async fn get_messages(&self, limit: Option<usize>) -> Result<Vec<ChatMessage>> {
13478 let messages = self.messages.read();
13479 let start = limit
13480 .map(|limit| messages.len().saturating_sub(limit))
13481 .unwrap_or(0);
13482 Ok(messages[start..].to_vec())
13483 }
13484
13485 async fn clear(&self) -> Result<()> {
13486 self.messages.write().clear();
13487 Ok(())
13488 }
13489
13490 fn len(&self) -> usize {
13491 self.messages.read().len()
13492 }
13493
13494 async fn restore(&self, snapshot: ai_agents_core::MemorySnapshot) -> Result<()> {
13495 *self.messages.write() = snapshot.messages;
13496 Ok(())
13497 }
13498 }
13499
13500 #[async_trait]
13501 impl ai_agents_memory::Memory for DroppingSignedAssistantMemory {}
13502
13503 #[async_trait]
13504 impl ai_agents_core::Memory for DroppingEarlierSequentialMemory {
13505 async fn add_message(&self, message: ChatMessage) -> Result<()> {
13506 let signed = message.role == ai_agents_core::Role::Assistant
13507 && ai_agents_core::decode_native_tool_call_markers(&message.content)
13508 .map_err(|error| AgentError::LLM(error.to_string()))?
13509 .is_some_and(|batch| batch.provider_state().is_some());
13510 let mut messages = self.messages.write();
13511 if signed && self.signed_seen.fetch_add(1, Ordering::SeqCst) == 1 {
13512 messages.retain(|stored| !stored.content.contains("seq-call-1"));
13513 }
13514 messages.push(message);
13515 Ok(())
13516 }
13517
13518 async fn get_messages(&self, limit: Option<usize>) -> Result<Vec<ChatMessage>> {
13519 let messages = self.messages.read();
13520 let start = limit
13521 .map(|limit| messages.len().saturating_sub(limit))
13522 .unwrap_or(0);
13523 Ok(messages[start..].to_vec())
13524 }
13525
13526 async fn clear(&self) -> Result<()> {
13527 self.messages.write().clear();
13528 self.signed_seen.store(0, Ordering::SeqCst);
13529 Ok(())
13530 }
13531
13532 fn len(&self) -> usize {
13533 self.messages.read().len()
13534 }
13535
13536 async fn restore(&self, snapshot: ai_agents_core::MemorySnapshot) -> Result<()> {
13537 *self.messages.write() = snapshot.messages;
13538 self.signed_seen.store(0, Ordering::SeqCst);
13539 Ok(())
13540 }
13541 }
13542
13543 #[async_trait]
13544 impl ai_agents_memory::Memory for DroppingEarlierSequentialMemory {}
13545
13546 #[async_trait]
13547 impl LLMProvider for TerminalHistoryProvider {
13548 async fn complete(
13549 &self,
13550 _messages: &[ChatMessage],
13551 _config: Option<&LLMConfig>,
13552 ) -> std::result::Result<LLMResponse, LLMError> {
13553 self.calls.fetch_add(1, Ordering::SeqCst);
13554 Err(LLMError::Serialization(
13555 "native history integrity failure".to_string(),
13556 ))
13557 }
13558
13559 async fn complete_stream(
13560 &self,
13561 _messages: &[ChatMessage],
13562 _config: Option<&LLMConfig>,
13563 ) -> std::result::Result<
13564 Box<dyn Stream<Item = std::result::Result<LLMChunk, LLMError>> + Unpin + Send>,
13565 LLMError,
13566 > {
13567 Err(LLMError::Serialization(
13568 "native history integrity failure".to_string(),
13569 ))
13570 }
13571
13572 fn provider_name(&self) -> &str {
13573 "terminal-history"
13574 }
13575
13576 fn supports(&self, _feature: LLMFeature) -> bool {
13577 false
13578 }
13579
13580 fn is_terminal_error(&self, error: &LLMError) -> bool {
13581 matches!(error, LLMError::Serialization(_))
13582 }
13583 }
13584
13585 fn disambiguation_state_machine(
13587 state_enabled: Option<bool>,
13588 require_confirmation: bool,
13589 ) -> Arc<StateMachine> {
13590 let definition = ai_agents_state::StateDefinition {
13591 prompt: Some("Handle the resolved request.".to_string()),
13592 disambiguation: Some(ai_agents_disambiguation::StateDisambiguationOverride {
13593 enabled: state_enabled,
13594 require_confirmation,
13595 ..Default::default()
13596 }),
13597 ..Default::default()
13598 };
13599 let review = ai_agents_state::StateDefinition {
13600 prompt: Some("Review a fresh request.".to_string()),
13601 ..Default::default()
13602 };
13603 Arc::new(
13604 StateMachine::new(ai_agents_state::StateConfig {
13605 initial: "active".to_string(),
13606 states: std::collections::HashMap::from([
13607 ("active".to_string(), definition),
13608 ("review".to_string(), review),
13609 ]),
13610 global_transitions: Vec::new(),
13611 fallback: None,
13612 max_no_transition: None,
13613 regenerate_on_transition: true,
13614 })
13615 .unwrap(),
13616 )
13617 }
13618
13619 fn state_disambiguation_agent(
13621 responses: Vec<&str>,
13622 manager_enabled: bool,
13623 state_enabled: Option<bool>,
13624 require_confirmation: bool,
13625 ) -> (RuntimeAgent, MockLLMProvider) {
13626 state_disambiguation_agent_with_skills(
13627 responses,
13628 manager_enabled,
13629 state_enabled,
13630 require_confirmation,
13631 Vec::new(),
13632 )
13633 }
13634
13635 fn state_disambiguation_agent_with_skills(
13637 responses: Vec<&str>,
13638 manager_enabled: bool,
13639 state_enabled: Option<bool>,
13640 require_confirmation: bool,
13641 skills: Vec<SkillDefinition>,
13642 ) -> (RuntimeAgent, MockLLMProvider) {
13643 let mut mock = MockLLMProvider::new("state-confirmation");
13644 mock.set_responses(responses.into_iter().map(String::from).collect(), false);
13645 let observed = mock.clone();
13646 let agent = AgentBuilder::new()
13647 .system_prompt("Handle requests.")
13648 .llm(Arc::new(mock.clone()))
13649 .llm_alias("router", Arc::new(mock))
13650 .state_machine(disambiguation_state_machine(
13651 state_enabled,
13652 require_confirmation,
13653 ))
13654 .skills(skills)
13655 .build()
13656 .unwrap()
13657 .with_disambiguation(DisambiguationConfig {
13658 enabled: manager_enabled,
13659 ..Default::default()
13660 });
13661 (agent, observed)
13662 }
13663
13664 fn confirmation_skill() -> SkillDefinition {
13666 SkillDefinition {
13667 id: "send_report".to_string(),
13668 description: "Send a report after clarification".to_string(),
13669 trigger: "When the user asks to send a report".to_string(),
13670 steps: vec![SkillStep::Prompt {
13671 prompt: "Execute confirmed report skill for: {{ input }}".to_string(),
13672 llm: None,
13673 }],
13674 reasoning: None,
13675 reflection: None,
13676 disambiguation: Some(ai_agents_disambiguation::SkillDisambiguationOverride {
13677 enabled: Some(true),
13678 ..Default::default()
13679 }),
13680 }
13681 }
13682
13683 fn confirmation_skill_call_count(observed: &MockLLMProvider) -> usize {
13685 observed
13686 .call_history()
13687 .iter()
13688 .filter(|call| {
13689 call.messages
13690 .iter()
13691 .any(|message| message.content.contains("Execute confirmed report skill"))
13692 })
13693 .count()
13694 }
13695
13696 struct BlockingRuntimeConfirmationObserver {
13697 entered: tokio::sync::Barrier,
13698 release: tokio::sync::Notify,
13699 }
13700
13701 impl BlockingRuntimeConfirmationObserver {
13702 fn new() -> Self {
13703 Self {
13704 entered: tokio::sync::Barrier::new(2),
13705 release: tokio::sync::Notify::new(),
13706 }
13707 }
13708 }
13709
13710 struct ResetOnTransitionHooks {
13711 agent: parking_lot::Mutex<Option<Weak<RuntimeAgent>>>,
13712 invoked: AtomicBool,
13713 }
13714
13715 #[async_trait]
13716 impl AgentHooks for ResetOnTransitionHooks {
13717 async fn on_state_transition(&self, _from: Option<&str>, _to: &str, _reason: &str) {
13718 if self.invoked.swap(true, Ordering::SeqCst) {
13719 return;
13720 }
13721 let agent = self.agent.lock().as_ref().and_then(Weak::upgrade);
13722 if let Some(agent) = agent {
13723 agent.reset().await.unwrap();
13724 }
13725 }
13726 }
13727
13728 impl ClarificationObserver for BlockingRuntimeConfirmationObserver {
13729 fn observe_question<'a>(
13730 &'a self,
13731 future: ClarificationQuestionFuture<'a>,
13732 ) -> ClarificationQuestionFuture<'a> {
13733 future
13734 }
13735
13736 fn observe_parse<'a>(
13737 &'a self,
13738 future: ClarificationParseFuture<'a>,
13739 ) -> ClarificationParseFuture<'a> {
13740 future
13741 }
13742
13743 fn observe_confirmation_parse<'a>(
13744 &'a self,
13745 future: ConfirmationParseFuture<'a>,
13746 ) -> ConfirmationParseFuture<'a> {
13747 Box::pin(async move {
13748 self.entered.wait().await;
13749 self.release.notified().await;
13750 future.await
13751 })
13752 }
13753 }
13754
13755 #[tokio::test]
13756 async fn state_confirmation_blocks_redispatch_until_explicit_agreement() {
13757 let (agent, observed) = state_disambiguation_agent(
13758 vec![
13759 r#"{"is_ambiguous":true,"confidence":0.2,"ambiguity_type":"missing_target","reasoning":"target missing","what_is_unclear":["target"],"detected_language":"en"}"#,
13760 r#"{"question":"What should I send?","options":null}"#,
13761 r#"{"status":"answered","selected_option":null,"enriched_input":"Send the report to Ada","resolved":{"intent":"send_report"}}"#,
13762 r#"{"question":"Should I send the report to Ada?"}"#,
13763 r#"{"status":"confirmed"}"#,
13764 "Request executed.",
13765 ],
13766 true,
13767 None,
13768 true,
13769 );
13770
13771 let clarification = agent.chat("Send it").await.unwrap();
13772 assert_eq!(clarification.content, "What should I send?");
13773 assert_eq!(observed.call_count(), 2);
13774
13775 let confirmation = agent.chat("The report to Ada").await.unwrap();
13776 assert_eq!(confirmation.content, "Should I send the report to Ada?");
13777 assert_eq!(
13778 confirmation
13779 .metadata
13780 .as_ref()
13781 .and_then(|metadata| metadata.get("disambiguation"))
13782 .and_then(|metadata| metadata.get("status"))
13783 .and_then(Value::as_str),
13784 Some("awaiting_confirmation")
13785 );
13786 assert_eq!(observed.call_count(), 4);
13787
13788 let completed = agent.chat("Yes").await.unwrap();
13789 assert_eq!(completed.content, "Request executed.");
13790 assert_eq!(observed.call_count(), 6);
13791 }
13792
13793 #[tokio::test]
13794 async fn streaming_state_confirmation_ends_the_turn_before_redispatch() {
13795 let (agent, observed) = state_disambiguation_agent(
13796 vec![
13797 r#"{"is_ambiguous":true,"confidence":0.2,"ambiguity_type":"missing_target","reasoning":"target missing","what_is_unclear":["target"],"detected_language":"en"}"#,
13798 r#"{"question":"What should I send?","options":null}"#,
13799 r#"{"status":"answered","selected_option":null,"enriched_input":"Send the report to Ada","resolved":{"intent":"send_report"}}"#,
13800 r#"{"question":"Should I send the report to Ada?"}"#,
13801 r#"{"status":"confirmed"}"#,
13802 "Request executed.",
13803 ],
13804 true,
13805 None,
13806 true,
13807 );
13808
13809 let mut clarification_stream = agent.chat_stream("Send it").await.unwrap();
13810 let mut clarification = String::new();
13811 while let Some(chunk) = clarification_stream.next().await {
13812 match chunk {
13813 StreamChunk::Content { text } => clarification.push_str(&text),
13814 StreamChunk::Done {} => break,
13815 StreamChunk::Error { message } => panic!("unexpected stream error: {message}"),
13816 _ => {}
13817 }
13818 }
13819 assert_eq!(clarification, "What should I send?");
13820 assert_eq!(observed.call_count(), 2);
13821
13822 let mut confirmation_stream = agent.chat_stream_events("The report to Ada").await.unwrap();
13823 let mut confirmation = None;
13824 while let Some(event) = confirmation_stream.next().await {
13825 match event {
13826 AgentStreamEvent::Final(response) => confirmation = Some(response),
13827 AgentStreamEvent::Chunk(StreamChunk::Error { message }) => {
13828 panic!("unexpected stream error: {message}")
13829 }
13830 AgentStreamEvent::Chunk(_) => {}
13831 }
13832 }
13833 let confirmation = confirmation.expect("confirmation must finalize");
13834 assert_eq!(confirmation.content, "Should I send the report to Ada?");
13835 assert_eq!(
13836 confirmation
13837 .metadata
13838 .as_ref()
13839 .and_then(|metadata| metadata.get("disambiguation"))
13840 .and_then(|metadata| metadata.get("status"))
13841 .and_then(Value::as_str),
13842 Some("awaiting_confirmation")
13843 );
13844 assert_eq!(observed.call_count(), 4);
13845
13846 let mut completed_stream = agent.chat_stream("Yes").await.unwrap();
13847 let mut completed = String::new();
13848 while let Some(chunk) = completed_stream.next().await {
13849 match chunk {
13850 StreamChunk::Content { text } => completed.push_str(&text),
13851 StreamChunk::Done {} => break,
13852 StreamChunk::Error { message } => panic!("unexpected stream error: {message}"),
13853 _ => {}
13854 }
13855 }
13856 assert_eq!(completed, "Request executed.");
13857 assert_eq!(observed.call_count(), 6);
13858 }
13859
13860 #[tokio::test]
13862 async fn root_turn_gate_serializes_blocking_and_streaming_entry_points() {
13863 let (complete_entered, mut complete_events) = tokio::sync::mpsc::unbounded_channel();
13864 let agent = Arc::new(
13865 AgentBuilder::new()
13866 .system_prompt("Serialize root turns.")
13867 .llm(Arc::new(RootTurnProbeProvider { complete_entered }))
13868 .build()
13869 .unwrap(),
13870 );
13871 let blocking_agent = Arc::clone(&agent);
13872
13873 let legacy_stream = agent.chat_stream("stream owner").await.unwrap();
13874 assert!(agent.root_turn_gate.try_lock().is_err());
13875 let blocking = tokio::spawn(async move { blocking_agent.chat("blocked").await.unwrap() });
13876 assert!(
13877 tokio::time::timeout(std::time::Duration::from_millis(50), complete_events.recv())
13878 .await
13879 .is_err(),
13880 "blocking turn reached the provider while the legacy stream owned the root gate"
13881 );
13882
13883 drop(legacy_stream);
13884 assert_eq!(
13885 tokio::time::timeout(std::time::Duration::from_secs(2), complete_events.recv())
13886 .await
13887 .expect("blocking turn did not enter after stream drop"),
13888 Some(())
13889 );
13890 let response = tokio::time::timeout(std::time::Duration::from_secs(2), blocking)
13891 .await
13892 .expect("blocking turn did not finish after stream drop")
13893 .unwrap();
13894 assert_eq!(response.content, "blocking complete");
13895
13896 let mut event_stream = agent.chat_stream_events("event terminal").await.unwrap();
13897 assert!(agent.root_turn_gate.try_lock().is_err());
13898 let mut saw_final = false;
13899 while let Some(event) = event_stream.next().await {
13900 if matches!(event, AgentStreamEvent::Final(_)) {
13901 saw_final = true;
13902 break;
13903 }
13904 }
13905 assert!(saw_final);
13906 assert!(
13907 agent.root_turn_gate.try_lock().is_ok(),
13908 "authoritative terminal event retained the root gate"
13909 );
13910 }
13911
13912 #[tokio::test]
13914 async fn response_hook_rejects_same_runtime_chat_reentry() {
13915 let hooks = Arc::new(ResponseChatHooks {
13916 target: parking_lot::Mutex::new(None),
13917 invoked: AtomicBool::new(false),
13918 nested_result: parking_lot::Mutex::new(None),
13919 });
13920 let agent = Arc::new(
13921 AgentBuilder::new()
13922 .system_prompt("Reject response hook reentry.")
13923 .llm(Arc::new(mock_with_response("outer response")))
13924 .hooks(hooks.clone())
13925 .build()
13926 .unwrap(),
13927 );
13928 *hooks.target.lock() = Some(Arc::downgrade(&agent));
13929
13930 let response = tokio::time::timeout(
13931 std::time::Duration::from_secs(2),
13932 agent.chat("outer request"),
13933 )
13934 .await
13935 .expect("same-runtime response hook reentry must fail without deadlocking")
13936 .unwrap();
13937
13938 assert_eq!(response.content, "outer response");
13939 let nested_result = hooks
13940 .nested_result
13941 .lock()
13942 .clone()
13943 .expect("response hook must record its nested call");
13944 let error = nested_result.expect_err("same-runtime nested chat must be rejected");
13945 assert!(error.contains("reentrant root turn ownership"));
13946 }
13947
13948 #[tokio::test]
13950 async fn root_turn_gate_allows_nested_runtime_and_rejects_cycles() {
13951 let agent_a = AgentBuilder::new()
13952 .system_prompt("Runtime A.")
13953 .llm(Arc::new(mock_with_response("response A")))
13954 .build()
13955 .unwrap();
13956 let agent_b = AgentBuilder::new()
13957 .system_prompt("Runtime B.")
13958 .llm(Arc::new(mock_with_response("response B")))
13959 .build()
13960 .unwrap();
13961 let RootTurnAdmission {
13962 guard: guard_a,
13963 identity_stack: stack_a,
13964 } = agent_a.acquire_root_turn().await.unwrap();
13965
13966 let cycle_error = scope_runtime_gate_identity_stack(&stack_a, async {
13967 let RootTurnAdmission {
13968 guard: guard_b,
13969 identity_stack: stack_b,
13970 } = agent_b
13971 .acquire_root_turn()
13972 .await
13973 .expect("runtime B must acquire a different gate");
13974 let result =
13975 scope_runtime_gate_identity_stack(&stack_b, agent_a.acquire_root_turn()).await;
13976 drop(guard_b);
13977 match result {
13978 Err(error) => error,
13979 Ok(_) => panic!("runtime A accepted a repeated gate identity"),
13980 }
13981 })
13982 .await;
13983 drop(guard_a);
13984
13985 assert!(
13986 cycle_error
13987 .to_string()
13988 .contains("reentrant root turn ownership")
13989 );
13990 }
13991
13992 #[tokio::test]
13994 async fn concurrent_orchestration_propagates_root_gate_ancestry() {
13995 let registry = Arc::new(crate::spawner::AgentRegistry::new());
13996 let hooks_a = Arc::new(ConcurrentResponseHooks {
13997 registry: Arc::downgrade(®istry),
13998 child_id: "runtime-b".to_string(),
13999 invoked: AtomicBool::new(false),
14000 nested_result: parking_lot::Mutex::new(None),
14001 });
14002 let hooks_b = Arc::new(ResponseChatHooks {
14003 target: parking_lot::Mutex::new(None),
14004 invoked: AtomicBool::new(false),
14005 nested_result: parking_lot::Mutex::new(None),
14006 });
14007 let agent_a = AgentBuilder::new()
14008 .system_prompt("Runtime A dispatches runtime B concurrently.")
14009 .llm(Arc::new(mock_with_response("response A")))
14010 .hooks(hooks_a.clone())
14011 .build()
14012 .unwrap();
14013 let agent_b = AgentBuilder::new()
14014 .system_prompt("Runtime B attempts to re-enter runtime A.")
14015 .llm(Arc::new(mock_with_response("response B")))
14016 .hooks(hooks_b.clone())
14017 .build()
14018 .unwrap();
14019 let spec_a = crate::spec::AgentSpec {
14020 name: "runtime-a".to_string(),
14021 system_prompt: "Runtime A dispatches runtime B concurrently.".to_string(),
14022 ..crate::spec::AgentSpec::default()
14023 };
14024 let spec_b = crate::spec::AgentSpec {
14025 name: "runtime-b".to_string(),
14026 system_prompt: "Runtime B attempts to re-enter runtime A.".to_string(),
14027 ..crate::spec::AgentSpec::default()
14028 };
14029 registry
14030 .register(crate::spawner::SpawnedAgent::from_runtime(
14031 "runtime-a".to_string(),
14032 agent_a,
14033 spec_a,
14034 ))
14035 .await
14036 .unwrap();
14037 registry
14038 .register(crate::spawner::SpawnedAgent::from_runtime(
14039 "runtime-b".to_string(),
14040 agent_b,
14041 spec_b,
14042 ))
14043 .await
14044 .unwrap();
14045 let runtime_a = registry.get("runtime-a").unwrap();
14046 *hooks_b.target.lock() = Some(Arc::downgrade(&runtime_a));
14047
14048 let response = tokio::time::timeout(
14049 std::time::Duration::from_secs(2),
14050 runtime_a.chat("outer concurrent request"),
14051 )
14052 .await
14053 .expect("concurrent orchestration cycle must fail without deadlocking")
14054 .unwrap();
14055
14056 assert_eq!(response.content, "response A");
14057 let child_result = hooks_a
14058 .nested_result
14059 .lock()
14060 .clone()
14061 .expect("runtime A hook must record runtime B completion");
14062 assert_eq!(child_result.unwrap(), "response B");
14063 let cycle_result = hooks_b
14064 .nested_result
14065 .lock()
14066 .clone()
14067 .expect("runtime B hook must record runtime A reentry");
14068 assert!(
14069 cycle_result
14070 .expect_err("runtime A accepted a repeated gate identity")
14071 .contains("reentrant root turn ownership")
14072 );
14073 }
14074
14075 fn skill_clarification_responses() -> Vec<&'static str> {
14078 vec![
14079 r#"{"is_ambiguous":false,"confidence":0.99,"ambiguity_type":null,"reasoning":"top-level clear","what_is_unclear":[],"detected_language":"en"}"#,
14080 "send_report",
14081 r#"{"is_ambiguous":true,"confidence":0.2,"ambiguity_type":"missing_target","reasoning":"target missing","what_is_unclear":["target"],"detected_language":"en"}"#,
14082 r#"{"question":"What should I send?","options":null}"#,
14083 ]
14084 }
14085
14086 #[tokio::test]
14092 async fn test_stream_skill_clarification_memory_matches_blocking() {
14093 let (blocking_agent, _) = state_disambiguation_agent_with_skills(
14094 skill_clarification_responses(),
14095 true,
14096 None,
14097 true,
14098 vec![confirmation_skill()],
14099 );
14100 let blocking = blocking_agent.chat("Send it").await.unwrap();
14101 let blocking_messages = blocking_agent.memory.get_messages(None).await.unwrap();
14102
14103 let (streaming_agent, _) = state_disambiguation_agent_with_skills(
14104 skill_clarification_responses(),
14105 true,
14106 None,
14107 true,
14108 vec![confirmation_skill()],
14109 );
14110 let (content, chunks, streamed) = collect_stream_events(&streaming_agent, "Send it").await;
14111 assert!(!chunks.iter().any(StreamChunk::is_error), "{chunks:?}");
14112 let streamed = streamed.expect("skill clarification must finalize as Final");
14113 let streaming_messages = streaming_agent.memory.get_messages(None).await.unwrap();
14114
14115 assert_eq!(blocking.content, "What should I send?");
14116 assert_eq!(streamed.content, blocking.content);
14117 assert_eq!(content, streamed.content);
14118 assert_eq!(
14119 blocking
14120 .metadata
14121 .as_ref()
14122 .and_then(|m| m.get("disambiguation")),
14123 streamed
14124 .metadata
14125 .as_ref()
14126 .and_then(|m| m.get("disambiguation")),
14127 );
14128 assert_eq!(
14129 streamed
14130 .metadata
14131 .as_ref()
14132 .and_then(|m| m.get("disambiguation"))
14133 .and_then(|d| d.get("status"))
14134 .and_then(Value::as_str),
14135 Some("awaiting_clarification"),
14136 );
14137 let shape = |messages: &[ChatMessage]| {
14138 messages
14139 .iter()
14140 .map(|m| (format!("{:?}", m.role), m.content.clone()))
14141 .collect::<Vec<_>>()
14142 };
14143 assert_eq!(shape(&blocking_messages), shape(&streaming_messages));
14144 assert_eq!(
14145 shape(&streaming_messages),
14146 vec![
14147 ("User".to_string(), "Send it".to_string()),
14148 ("Assistant".to_string(), "What should I send?".to_string()),
14149 ],
14150 );
14151 assert_eq!(
14152 *streaming_agent.pending_skill_id.read(),
14153 Some("send_report".to_string()),
14154 );
14155 }
14156
14157 #[tokio::test]
14159 async fn test_stream_skill_clarification_memory_failure_surfaces_as_error() {
14160 let build = || {
14162 let mut mock = MockLLMProvider::new("skill-clarification");
14163 mock.set_responses(
14164 skill_clarification_responses()
14165 .into_iter()
14166 .map(String::from)
14167 .collect(),
14168 false,
14169 );
14170 AgentBuilder::new()
14171 .system_prompt("Handle requests.")
14172 .llm(Arc::new(mock.clone()))
14173 .llm_alias("router", Arc::new(mock))
14174 .state_machine(disambiguation_state_machine(None, true))
14175 .skills(vec![confirmation_skill()])
14176 .memory(Arc::new(FailingMemory {
14177 messages: parking_lot::RwLock::new(Vec::new()),
14178 fail_on_add: 2,
14179 adds: std::sync::atomic::AtomicUsize::new(0),
14180 }))
14181 .build()
14182 .unwrap()
14183 .with_disambiguation(DisambiguationConfig {
14184 enabled: true,
14185 ..Default::default()
14186 })
14187 };
14188
14189 let blocking = build().chat("Send it").await;
14190 assert!(
14191 blocking.is_err(),
14192 "blocking must surface the failed clarification write: {blocking:?}"
14193 );
14194
14195 let (_, chunks, streamed) = collect_stream_events(&build(), "Send it").await;
14196 assert!(
14197 streamed.is_none(),
14198 "a failed write must not finalize the turn"
14199 );
14200 assert!(
14201 chunks.iter().any(|chunk| matches!(
14202 chunk,
14203 StreamChunk::Error { message } if message.contains("simulated memory failure")
14204 )),
14205 "streaming must surface the failed clarification write: {chunks:?}"
14206 );
14207 }
14208
14209 #[tokio::test]
14211 async fn confirmed_skill_route_executes_exactly_once() {
14212 let (agent, observed) = state_disambiguation_agent_with_skills(
14213 vec![
14214 r#"{"is_ambiguous":false,"confidence":0.99,"ambiguity_type":null,"reasoning":"top-level clear","what_is_unclear":[],"detected_language":"en"}"#,
14215 "send_report",
14216 r#"{"is_ambiguous":true,"confidence":0.2,"ambiguity_type":"missing_target","reasoning":"target missing","what_is_unclear":["target"],"detected_language":"en"}"#,
14217 r#"{"question":"What should I send?","options":null}"#,
14218 r#"{"status":"answered","selected_option":null,"enriched_input":"Send the report to Ada","resolved":{"intent":"send_report"}}"#,
14219 r#"{"question":"Should I send the report to Ada?"}"#,
14220 r#"{"status":"confirmed"}"#,
14221 r#"{"is_ambiguous":false,"confidence":0.99,"ambiguity_type":null,"reasoning":"resolved","what_is_unclear":[],"detected_language":"en"}"#,
14222 "Report skill executed.",
14223 ],
14224 true,
14225 None,
14226 true,
14227 vec![confirmation_skill()],
14228 );
14229
14230 let clarification = agent.chat("Send it").await.unwrap();
14231 assert_eq!(clarification.content, "What should I send?");
14232 assert_eq!(confirmation_skill_call_count(&observed), 0);
14233
14234 let confirmation = agent.chat("The report to Ada").await.unwrap();
14235 assert_eq!(confirmation.content, "Should I send the report to Ada?");
14236 assert_eq!(
14237 confirmation
14238 .metadata
14239 .as_ref()
14240 .and_then(|metadata| metadata.get("disambiguation"))
14241 .and_then(|metadata| metadata.get("status"))
14242 .and_then(Value::as_str),
14243 Some("awaiting_confirmation")
14244 );
14245 assert_eq!(confirmation_skill_call_count(&observed), 0);
14246
14247 let completed = agent.chat("Yes").await.unwrap();
14248 assert_eq!(completed.content, "Report skill executed.");
14249 assert_eq!(confirmation_skill_call_count(&observed), 1);
14250 assert!(agent.pending_skill_id.read().is_none());
14251 let messages = agent.memory.get_messages(None).await.unwrap();
14252 assert!(!messages.iter().any(|message| message.content == "Yes"));
14253 }
14254
14255 #[tokio::test]
14257 async fn confirmed_skill_recheck_preserves_new_clarification_metadata() {
14258 let (agent, observed) = state_disambiguation_agent_with_skills(
14259 vec![
14260 r#"{"is_ambiguous":false,"confidence":0.99,"ambiguity_type":null,"reasoning":"top-level clear","what_is_unclear":[],"detected_language":"en"}"#,
14261 "send_report",
14262 r#"{"is_ambiguous":true,"confidence":0.2,"ambiguity_type":"missing_target","reasoning":"target missing","what_is_unclear":["target"],"detected_language":"en"}"#,
14263 r#"{"question":"What should I send?","options":null}"#,
14264 r#"{"status":"answered","selected_option":null,"enriched_input":"Send the report to Ada","resolved":{"intent":"send_report"}}"#,
14265 r#"{"question":"Should I send the report to Ada?"}"#,
14266 r#"{"status":"confirmed"}"#,
14267 r#"{"is_ambiguous":true,"confidence":0.3,"ambiguity_type":"missing_parameters","reasoning":"timing missing","what_is_unclear":["timing"],"detected_language":"en"}"#,
14268 r#"{"question":"When should I send it?","options":null}"#,
14269 ],
14270 true,
14271 None,
14272 true,
14273 vec![confirmation_skill()],
14274 );
14275
14276 agent.chat("Send it").await.unwrap();
14277 agent.chat("The report to Ada").await.unwrap();
14278 let follow_up = agent.chat("Yes").await.unwrap();
14279
14280 assert_eq!(follow_up.content, "When should I send it?");
14281 let metadata = follow_up
14282 .metadata
14283 .as_ref()
14284 .and_then(|metadata| metadata.get("disambiguation"))
14285 .unwrap();
14286 assert_eq!(
14287 metadata.get("status").and_then(Value::as_str),
14288 Some("awaiting_clarification")
14289 );
14290 assert_eq!(
14291 metadata.get("skill_id").and_then(Value::as_str),
14292 Some("send_report")
14293 );
14294 assert!(metadata.get("detection").is_some());
14295 assert_eq!(confirmation_skill_call_count(&observed), 0);
14296 }
14297
14298 #[tokio::test]
14300 async fn rejected_skill_confirmation_never_executes() {
14301 let (agent, observed) = state_disambiguation_agent_with_skills(
14302 vec![
14303 r#"{"is_ambiguous":false,"confidence":0.99,"ambiguity_type":null,"reasoning":"top-level clear","what_is_unclear":[],"detected_language":"en"}"#,
14304 "send_report",
14305 r#"{"is_ambiguous":true,"confidence":0.2,"ambiguity_type":"missing_target","reasoning":"target missing","what_is_unclear":["target"],"detected_language":"en"}"#,
14306 r#"{"question":"What should I send?","options":null}"#,
14307 r#"{"status":"answered","selected_option":null,"enriched_input":"Send the report to Ada","resolved":{"intent":"send_report"}}"#,
14308 r#"{"question":"Should I send the report to Ada?"}"#,
14309 r#"{"status":"rejected"}"#,
14310 "Confirmation rejected.",
14311 ],
14312 true,
14313 None,
14314 true,
14315 vec![confirmation_skill()],
14316 );
14317
14318 agent.chat("Send it").await.unwrap();
14319 agent.chat("The report to Ada").await.unwrap();
14320 let rejected = agent.chat("No").await.unwrap();
14321
14322 assert_eq!(rejected.content, "Confirmation rejected.");
14323 assert_eq!(confirmation_skill_call_count(&observed), 0);
14324 assert!(agent.pending_skill_id.read().is_none());
14325 }
14326
14327 #[tokio::test]
14329 async fn reset_invalidates_pending_skill_confirmation_before_streaming_input() {
14330 let (agent, observed) = state_disambiguation_agent_with_skills(
14331 vec![
14332 r#"{"is_ambiguous":false,"confidence":0.99,"ambiguity_type":null,"reasoning":"top-level clear","what_is_unclear":[],"detected_language":"en"}"#,
14333 "send_report",
14334 r#"{"is_ambiguous":true,"confidence":0.2,"ambiguity_type":"missing_target","reasoning":"target missing","what_is_unclear":["target"],"detected_language":"en"}"#,
14335 r#"{"question":"What should I send?","options":null}"#,
14336 r#"{"status":"answered","selected_option":null,"enriched_input":"Send the report to Ada","resolved":{"intent":"send_report"}}"#,
14337 r#"{"question":"Should I send the report to Ada?"}"#,
14338 r#"{"is_ambiguous":false,"confidence":0.99,"ambiguity_type":null,"reasoning":"fresh input","what_is_unclear":[],"detected_language":"en"}"#,
14339 "none",
14340 "Fresh response.",
14341 ],
14342 true,
14343 None,
14344 true,
14345 vec![confirmation_skill()],
14346 );
14347
14348 agent.chat("Send it").await.unwrap();
14349 agent.chat("The report to Ada").await.unwrap();
14350 agent.reset().await.unwrap();
14351 assert!(agent.pending_skill_id.read().is_none());
14352 assert!(
14353 !agent
14354 .disambiguation_manager()
14355 .unwrap()
14356 .has_pending_clarification()
14357 .await
14358 );
14359
14360 let mut stream = agent.chat_stream("Yes").await.unwrap();
14361 let mut content = String::new();
14362 while let Some(chunk) = stream.next().await {
14363 match chunk {
14364 StreamChunk::Content { text } => content.push_str(&text),
14365 StreamChunk::Done {} => break,
14366 StreamChunk::Error { message } => panic!("unexpected stream error: {message}"),
14367 _ => {}
14368 }
14369 }
14370
14371 assert_eq!(content, "Fresh response.");
14372 assert_eq!(confirmation_skill_call_count(&observed), 0);
14373 }
14374
14375 #[tokio::test]
14377 async fn trait_reset_clears_pending_skill_confirmation() {
14378 let (agent, _) = state_disambiguation_agent_with_skills(
14379 vec![
14380 r#"{"is_ambiguous":false,"confidence":0.99,"ambiguity_type":null,"reasoning":"top-level clear","what_is_unclear":[],"detected_language":"en"}"#,
14381 "send_report",
14382 r#"{"is_ambiguous":true,"confidence":0.2,"ambiguity_type":"missing_target","reasoning":"target missing","what_is_unclear":["target"],"detected_language":"en"}"#,
14383 r#"{"question":"What should I send?","options":null}"#,
14384 r#"{"status":"answered","selected_option":null,"enriched_input":"Send the report to Ada","resolved":{"intent":"send_report"}}"#,
14385 r#"{"question":"Should I send the report to Ada?"}"#,
14386 ],
14387 true,
14388 None,
14389 true,
14390 vec![confirmation_skill()],
14391 );
14392
14393 agent.chat("Send it").await.unwrap();
14394 agent.chat("The report to Ada").await.unwrap();
14395 <RuntimeAgent as Agent>::reset(&agent).await.unwrap();
14396
14397 assert!(agent.pending_skill_id.read().is_none());
14398 assert!(
14399 !agent
14400 .disambiguation_manager()
14401 .unwrap()
14402 .has_pending_clarification()
14403 .await
14404 );
14405 }
14406
14407 #[tokio::test]
14409 async fn state_change_invalidates_pending_skill_confirmation() {
14410 let (agent, observed) = state_disambiguation_agent_with_skills(
14411 vec![
14412 r#"{"is_ambiguous":false,"confidence":0.99,"ambiguity_type":null,"reasoning":"top-level clear","what_is_unclear":[],"detected_language":"en"}"#,
14413 "send_report",
14414 r#"{"is_ambiguous":true,"confidence":0.2,"ambiguity_type":"missing_target","reasoning":"target missing","what_is_unclear":["target"],"detected_language":"en"}"#,
14415 r#"{"question":"What should I send?","options":null}"#,
14416 r#"{"status":"answered","selected_option":null,"enriched_input":"Send the report to Ada","resolved":{"intent":"send_report"}}"#,
14417 r#"{"question":"Should I send the report to Ada?"}"#,
14418 r#"{"is_ambiguous":false,"confidence":0.99,"ambiguity_type":null,"reasoning":"fresh input","what_is_unclear":[],"detected_language":"en"}"#,
14419 "none",
14420 "Fresh response.",
14421 ],
14422 true,
14423 None,
14424 true,
14425 vec![confirmation_skill()],
14426 );
14427
14428 agent.chat("Send it").await.unwrap();
14429 agent.chat("The report to Ada").await.unwrap();
14430 agent.transition_to("review").await.unwrap();
14431 let cancelled = agent.chat("Yes").await.unwrap();
14432
14433 assert_eq!(cancelled.content, "Fresh response.");
14434 assert_eq!(confirmation_skill_call_count(&observed), 0);
14435 assert!(agent.pending_skill_id.read().is_none());
14436 }
14437
14438 #[tokio::test]
14440 async fn in_flight_confirmation_cannot_redispatch_after_reset() {
14441 let (mut agent, observed) = state_disambiguation_agent_with_skills(
14442 vec![
14443 r#"{"is_ambiguous":false,"confidence":0.99,"ambiguity_type":null,"reasoning":"top-level clear","what_is_unclear":[],"detected_language":"en"}"#,
14444 "send_report",
14445 r#"{"is_ambiguous":true,"confidence":0.2,"ambiguity_type":"missing_target","reasoning":"target missing","what_is_unclear":["target"],"detected_language":"en"}"#,
14446 r#"{"question":"What should I send?","options":null}"#,
14447 r#"{"status":"answered","selected_option":null,"enriched_input":"Send the report to Ada","resolved":{"intent":"send_report"}}"#,
14448 r#"{"question":"Should I send the report to Ada?"}"#,
14449 r#"{"status":"confirmed"}"#,
14450 "Confirmation cancelled.",
14451 ],
14452 true,
14453 None,
14454 true,
14455 vec![confirmation_skill()],
14456 );
14457 let observer = Arc::new(BlockingRuntimeConfirmationObserver::new());
14458 let manager = agent
14459 .disambiguation_manager
14460 .take()
14461 .unwrap()
14462 .with_clarification_observer(observer.clone());
14463 agent.disambiguation_manager = Some(manager);
14464 let agent = Arc::new(agent);
14465
14466 agent.chat("Send it").await.unwrap();
14467 agent.chat("The report to Ada").await.unwrap();
14468
14469 let confirming_agent = Arc::clone(&agent);
14470 let confirmation = tokio::spawn(async move { confirming_agent.chat("Yes").await });
14471 observer.entered.wait().await;
14472 agent.reset().await.unwrap();
14473 observer.release.notify_one();
14474
14475 let response = confirmation.await.unwrap().unwrap();
14476 assert_eq!(response.content, "Confirmation cancelled.");
14477 assert_eq!(confirmation_skill_call_count(&observed), 0);
14478 assert!(agent.pending_skill_id.read().is_none());
14479 }
14480
14481 #[tokio::test]
14483 async fn queued_reset_prevents_stale_confirmation_question_publication() {
14484 let (agent, observed) = state_disambiguation_agent(
14485 vec![
14486 r#"{"is_ambiguous":true,"confidence":0.2,"ambiguity_type":"missing_target","reasoning":"target missing","what_is_unclear":["target"],"detected_language":"en"}"#,
14487 r#"{"question":"What should I send?","options":null}"#,
14488 r#"{"status":"answered","selected_option":null,"enriched_input":"Send the report to Ada","resolved":{"intent":"send_report"}}"#,
14489 r#"{"question":"Should I send the report to Ada?"}"#,
14490 ],
14491 true,
14492 None,
14493 true,
14494 );
14495 let agent = Arc::new(agent);
14496 agent.chat("Send it").await.unwrap();
14497
14498 let admission = agent.disambiguation_admission.write().await;
14499 let (started_tx, started_rx) = tokio::sync::oneshot::channel();
14500 let resetting_agent = Arc::clone(&agent);
14501 let reset = tokio::spawn(async move {
14502 let _ = started_tx.send(());
14503 resetting_agent.reset().await
14504 });
14505 started_rx.await.unwrap();
14506 tokio::task::yield_now().await;
14507
14508 let responding_agent = Arc::clone(&agent);
14509 let response =
14510 tokio::spawn(async move { responding_agent.chat("The report to Ada").await });
14511 tokio::time::timeout(std::time::Duration::from_secs(2), async {
14512 while observed.call_count() < 4 {
14513 tokio::task::yield_now().await;
14514 }
14515 })
14516 .await
14517 .expect("clarification processing must reach terminal publication");
14518 drop(admission);
14519
14520 reset.await.unwrap().unwrap();
14521 let error = response.await.unwrap().unwrap_err();
14522 assert!(error.to_string().contains("ownership changed"));
14523 assert!(
14524 !agent
14525 .disambiguation_manager()
14526 .unwrap()
14527 .has_pending_clarification()
14528 .await
14529 );
14530 assert!(agent.memory.get_messages(None).await.unwrap().is_empty());
14531 }
14532
14533 #[tokio::test]
14535 async fn queued_reset_prevents_stale_skill_clarification_publication() {
14536 let (agent, observed) = state_disambiguation_agent_with_skills(
14537 vec![
14538 r#"{"is_ambiguous":false,"confidence":0.99,"ambiguity_type":null,"reasoning":"top-level clear","what_is_unclear":[],"detected_language":"en"}"#,
14539 "send_report",
14540 r#"{"is_ambiguous":true,"confidence":0.2,"ambiguity_type":"missing_target","reasoning":"target missing","what_is_unclear":["target"],"detected_language":"en"}"#,
14541 r#"{"question":"What should I send?","options":null}"#,
14542 ],
14543 true,
14544 None,
14545 true,
14546 vec![confirmation_skill()],
14547 );
14548 let agent = Arc::new(agent);
14549 let admission = agent.disambiguation_admission.write().await;
14550 let (started_tx, started_rx) = tokio::sync::oneshot::channel();
14551 let resetting_agent = Arc::clone(&agent);
14552 let reset = tokio::spawn(async move {
14553 let _ = started_tx.send(());
14554 resetting_agent.reset().await
14555 });
14556 started_rx.await.unwrap();
14557 tokio::task::yield_now().await;
14558
14559 let responding_agent = Arc::clone(&agent);
14560 let response = tokio::spawn(async move { responding_agent.chat("Send it").await });
14561 tokio::time::timeout(std::time::Duration::from_secs(2), async {
14562 while observed.call_count() < 4 {
14563 tokio::task::yield_now().await;
14564 }
14565 })
14566 .await
14567 .expect("skill clarification must reach terminal publication");
14568 drop(admission);
14569
14570 reset.await.unwrap().unwrap();
14571 let error = response.await.unwrap().unwrap_err();
14572 assert!(error.to_string().contains("ownership changed"));
14573 assert_eq!(confirmation_skill_call_count(&observed), 0);
14574 assert!(agent.pending_skill_id.read().is_none());
14575 assert!(agent.memory.get_messages(None).await.unwrap().is_empty());
14576 }
14577
14578 #[tokio::test]
14580 async fn transition_hook_can_reset_without_admission_deadlock() {
14581 let hooks = Arc::new(ResetOnTransitionHooks {
14582 agent: parking_lot::Mutex::new(None),
14583 invoked: AtomicBool::new(false),
14584 });
14585 let agent = Arc::new(
14586 AgentBuilder::new()
14587 .system_prompt("Test transition hook reentrancy.")
14588 .llm(Arc::new(mock_with_response("done")))
14589 .state_machine(disambiguation_state_machine(None, false))
14590 .build()
14591 .unwrap()
14592 .with_hooks(hooks.clone()),
14593 );
14594 *hooks.agent.lock() = Some(Arc::downgrade(&agent));
14595
14596 let transitioned = tokio::time::timeout(
14597 std::time::Duration::from_secs(2),
14598 agent.apply_transition_target("active", "review", "test transition", None),
14599 )
14600 .await
14601 .expect("transition hook reset must not deadlock")
14602 .unwrap();
14603
14604 assert!(transitioned);
14605 assert!(hooks.invoked.load(Ordering::SeqCst));
14606 assert_eq!(agent.current_state().as_deref(), Some("active"));
14607 }
14608
14609 #[tokio::test]
14611 async fn concurrent_transition_cannot_duplicate_exit_actions() {
14612 let gate = PathMutationGate::new();
14613 let active = ai_agents_state::StateDefinition {
14614 on_exit: vec![StateAction::Tool {
14615 tool: "transition_exit".to_string(),
14616 args: Some(serde_json::json!({"path": "./transition-exit.txt"})),
14617 }],
14618 ..Default::default()
14619 };
14620 let state_machine = Arc::new(
14621 StateMachine::new(ai_agents_state::StateConfig {
14622 initial: "active".to_string(),
14623 states: HashMap::from([
14624 ("active".to_string(), active),
14625 (
14626 "review".to_string(),
14627 ai_agents_state::StateDefinition::default(),
14628 ),
14629 ]),
14630 global_transitions: Vec::new(),
14631 fallback: None,
14632 max_no_transition: None,
14633 regenerate_on_transition: true,
14634 })
14635 .unwrap(),
14636 );
14637 let agent = Arc::new(
14638 AgentBuilder::new()
14639 .system_prompt("Test transition reservation.")
14640 .llm(Arc::new(mock_with_response("done")))
14641 .tool(Arc::new(BlockingPathMutationTool {
14642 id: "transition_exit",
14643 path_fields: vec![ai_agents_core::PathPolicyBinding::write("path")],
14644 gate: gate.clone(),
14645 }))
14646 .state_machine(state_machine)
14647 .build()
14648 .unwrap(),
14649 );
14650
14651 let first_agent = Arc::clone(&agent);
14652 let first = tokio::spawn(async move { first_agent.transition_to("review").await });
14653 tokio::time::timeout(std::time::Duration::from_secs(2), gate.wait_until_entered())
14654 .await
14655 .expect("reserved transition must enter its exit action");
14656
14657 let second = tokio::time::timeout(
14658 std::time::Duration::from_secs(2),
14659 agent.transition_to("review"),
14660 )
14661 .await
14662 .expect("competing transition must fail without waiting for the exit action")
14663 .unwrap_err();
14664 assert!(second.to_string().contains("already in progress"));
14665
14666 gate.release();
14667 first.await.unwrap().unwrap();
14668 assert_eq!(agent.current_state().as_deref(), Some("review"));
14669 }
14670
14671 #[tokio::test]
14673 async fn concurrent_transition_cannot_overtake_enter_actions() {
14674 let gate = PathMutationGate::new();
14675 let review = ai_agents_state::StateDefinition {
14676 on_enter: vec![StateAction::Tool {
14677 tool: "transition_enter".to_string(),
14678 args: Some(serde_json::json!({"path": "./transition-enter.txt"})),
14679 }],
14680 ..Default::default()
14681 };
14682 let state_machine = Arc::new(
14683 StateMachine::new(ai_agents_state::StateConfig {
14684 initial: "active".to_string(),
14685 states: HashMap::from([
14686 (
14687 "active".to_string(),
14688 ai_agents_state::StateDefinition::default(),
14689 ),
14690 ("review".to_string(), review),
14691 ]),
14692 global_transitions: Vec::new(),
14693 fallback: None,
14694 max_no_transition: None,
14695 regenerate_on_transition: true,
14696 })
14697 .unwrap(),
14698 );
14699 let agent = Arc::new(
14700 AgentBuilder::new()
14701 .system_prompt("Test transition lifecycle reservation.")
14702 .llm(Arc::new(mock_with_response("done")))
14703 .tool(Arc::new(BlockingPathMutationTool {
14704 id: "transition_enter",
14705 path_fields: vec![ai_agents_core::PathPolicyBinding::write("path")],
14706 gate: gate.clone(),
14707 }))
14708 .state_machine(state_machine)
14709 .build()
14710 .unwrap(),
14711 );
14712
14713 let first_agent = Arc::clone(&agent);
14714 let first = tokio::spawn(async move { first_agent.transition_to("review").await });
14715 tokio::time::timeout(std::time::Duration::from_secs(2), gate.wait_until_entered())
14716 .await
14717 .expect("committed transition must enter its destination action");
14718
14719 let second = agent.transition_to("active").await.unwrap_err();
14720 assert!(second.to_string().contains("already in progress"));
14721 assert!(agent.reset().await.is_err());
14722
14723 gate.release();
14724 first.await.unwrap().unwrap();
14725 assert_eq!(agent.current_state().as_deref(), Some("review"));
14726 }
14727
14728 #[tokio::test]
14730 async fn same_state_restore_invalidates_pending_skill_confirmation() {
14731 let (agent, observed) = state_disambiguation_agent_with_skills(
14732 vec![
14733 r#"{"is_ambiguous":false,"confidence":0.99,"ambiguity_type":null,"reasoning":"top-level clear","what_is_unclear":[],"detected_language":"en"}"#,
14734 "send_report",
14735 r#"{"is_ambiguous":true,"confidence":0.2,"ambiguity_type":"missing_target","reasoning":"target missing","what_is_unclear":["target"],"detected_language":"en"}"#,
14736 r#"{"question":"What should I send?","options":null}"#,
14737 r#"{"status":"answered","selected_option":null,"enriched_input":"Send the report to Ada","resolved":{"intent":"send_report"}}"#,
14738 r#"{"question":"Should I send the report to Ada?"}"#,
14739 ],
14740 true,
14741 None,
14742 true,
14743 vec![confirmation_skill()],
14744 );
14745
14746 agent.chat("Send it").await.unwrap();
14747 agent.chat("The report to Ada").await.unwrap();
14748 let snapshot = agent.save_state().await.unwrap();
14749 assert_eq!(agent.current_state().as_deref(), Some("active"));
14750
14751 agent.restore_state(snapshot).await.unwrap();
14752
14753 assert_eq!(agent.current_state().as_deref(), Some("active"));
14754 assert!(agent.pending_skill_id.read().is_none());
14755 assert!(
14756 !agent
14757 .disambiguation_manager()
14758 .unwrap()
14759 .has_pending_clarification()
14760 .await
14761 );
14762 assert_eq!(confirmation_skill_call_count(&observed), 0);
14763 }
14764
14765 #[tokio::test]
14767 async fn direct_state_generation_change_invalidates_confirmation() {
14768 let (agent, observed) = state_disambiguation_agent_with_skills(
14769 vec![
14770 r#"{"is_ambiguous":false,"confidence":0.99,"ambiguity_type":null,"reasoning":"top-level clear","what_is_unclear":[],"detected_language":"en"}"#,
14771 "send_report",
14772 r#"{"is_ambiguous":true,"confidence":0.2,"ambiguity_type":"missing_target","reasoning":"target missing","what_is_unclear":["target"],"detected_language":"en"}"#,
14773 r#"{"question":"What should I send?","options":null}"#,
14774 r#"{"status":"answered","selected_option":null,"enriched_input":"Send the report to Ada","resolved":{"intent":"send_report"}}"#,
14775 r#"{"question":"Should I send the report to Ada?"}"#,
14776 "Confirmation cancelled.",
14777 ],
14778 true,
14779 None,
14780 true,
14781 vec![confirmation_skill()],
14782 );
14783
14784 agent.chat("Send it").await.unwrap();
14785 agent.chat("The report to Ada").await.unwrap();
14786 let state_machine = agent.state_machine().unwrap();
14787 state_machine
14788 .transition_to("review", "external test")
14789 .unwrap();
14790 state_machine
14791 .transition_to("active", "external test")
14792 .unwrap();
14793
14794 let response = agent.chat("Yes").await.unwrap();
14795
14796 assert_eq!(response.content, "Confirmation cancelled.");
14797 assert_eq!(confirmation_skill_call_count(&observed), 0);
14798 assert!(agent.pending_skill_id.read().is_none());
14799 }
14800
14801 #[tokio::test]
14802 async fn state_confirmation_does_not_add_a_question_for_clear_input() {
14803 let (agent, observed) = state_disambiguation_agent(
14804 vec![
14805 r#"{"is_ambiguous":false,"confidence":0.99,"ambiguity_type":null,"reasoning":"clear","what_is_unclear":[],"detected_language":"en"}"#,
14806 "Request executed.",
14807 ],
14808 true,
14809 None,
14810 true,
14811 );
14812
14813 let response = agent.chat("Send the report to Ada").await.unwrap();
14814
14815 assert_eq!(response.content, "Request executed.");
14816 assert_eq!(observed.call_count(), 2);
14817 }
14818
14819 #[tokio::test]
14820 async fn state_override_cannot_activate_a_disabled_top_level_manager() {
14821 let (agent, observed) =
14822 state_disambiguation_agent(vec!["Request executed."], false, Some(true), true);
14823
14824 assert!(!agent.has_disambiguation());
14825 let response = agent.chat("Send it").await.unwrap();
14826
14827 assert_eq!(response.content, "Request executed.");
14828 assert_eq!(observed.call_count(), 1);
14829 }
14830
14831 #[tokio::test]
14832 async fn native_required_choice_executes_through_the_shared_tool_path() {
14833 let mut mock = MockLLMProvider::new("native-required");
14834 mock.set_tool_choice(Some(ToolChoice::Required));
14835 let native_call = ToolCall {
14836 id: "provider-call-1".to_string(),
14837 name: "calculator".to_string(),
14838 arguments: serde_json::json!({"expression": "2 + 2"}),
14839 };
14840 let provider_state = ai_agents_core::NativeProviderState::new(
14841 "fixture-exchange-1",
14842 "fixture",
14843 "native-tools",
14844 ai_agents_core::NativeProviderTarget::new("https://fixture.invalid/", "fixture-model")
14845 .unwrap(),
14846 serde_json::json!({
14847 "role": "model",
14848 "parts": [{
14849 "functionCall": {"name": "calculator", "args": {"expression": "2 + 2"}},
14850 "thoughtSignature": "fixture-signature"
14851 }]
14852 }),
14853 vec![ai_agents_core::NativeCallBinding::new("provider-call-1", 0).unwrap()],
14854 )
14855 .unwrap();
14856 mock.add_response(
14857 LLMResponse::new("", FinishReason::ToolCall)
14858 .with_provider_state(provider_state)
14859 .unwrap()
14860 .with_tool_calls(vec![native_call])
14861 .unwrap(),
14862 );
14863 mock.add_response(LLMResponse::new("The answer is 4.", FinishReason::Stop));
14864 let observed = mock.clone();
14865 let agent = AgentBuilder::new()
14866 .system_prompt("Use the calculator when needed.")
14867 .llm(Arc::new(mock))
14868 .tool(Arc::new(CalculatorTool::new()))
14869 .build()
14870 .unwrap();
14871
14872 let response = agent.chat("What is 2 + 2?").await.unwrap();
14873
14874 assert_eq!(response.content, "The answer is 4.");
14875 assert_eq!(
14876 response.tool_calls.as_ref().unwrap()[0].id,
14877 "provider-call-1"
14878 );
14879 let calls = observed.call_history();
14880 assert_eq!(calls.len(), 2);
14881 assert!(matches!(
14882 calls[0].request.as_ref().map(|request| &request.choice),
14883 Some(ToolChoice::Required)
14884 ));
14885 assert!(matches!(
14886 calls[1].request.as_ref().map(|request| &request.choice),
14887 Some(ToolChoice::Auto)
14888 ));
14889 let replay_batch = calls[1]
14890 .messages
14891 .iter()
14892 .find_map(|message| {
14893 ai_agents_core::decode_native_tool_call_markers(&message.content).unwrap()
14894 })
14895 .expect("signed native call marker must be replayed");
14896 assert_eq!(
14897 replay_batch.provider_state().unwrap().exchange_id(),
14898 "fixture-exchange-1"
14899 );
14900 assert!(calls[1].messages.iter().any(|message| {
14901 ai_agents_core::decode_native_tool_result_markers(&message.content)
14902 .is_ok_and(|results| results.is_some())
14903 }));
14904 }
14905
14906 #[tokio::test]
14907 async fn custom_memory_loss_stops_before_signed_tool_execution() {
14908 let mut mock = MockLLMProvider::new("native-custom-memory");
14909 mock.set_tool_choice(Some(ToolChoice::Required));
14910 let call = ToolCall {
14911 id: "provider-call-drop".to_string(),
14912 name: "calculator".to_string(),
14913 arguments: serde_json::json!({"expression": "3 + 4"}),
14914 };
14915 let state = ai_agents_core::NativeProviderState::new(
14916 "fixture-exchange-drop",
14917 "fixture",
14918 "native-tools",
14919 ai_agents_core::NativeProviderTarget::new("https://fixture.invalid/", "fixture-model")
14920 .unwrap(),
14921 serde_json::json!({
14922 "role": "model",
14923 "parts": [{
14924 "functionCall": {"name": "calculator", "args": {"expression": "3 + 4"}},
14925 "thoughtSignature": "fixture-signature-drop"
14926 }]
14927 }),
14928 vec![ai_agents_core::NativeCallBinding::new("provider-call-drop", 0).unwrap()],
14929 )
14930 .unwrap();
14931 mock.add_response(
14932 LLMResponse::new("", FinishReason::ToolCall)
14933 .with_provider_state(state)
14934 .unwrap()
14935 .with_tool_calls(vec![call])
14936 .unwrap(),
14937 );
14938 let agent = AgentBuilder::new()
14939 .system_prompt("Use the calculator.")
14940 .llm(Arc::new(mock))
14941 .memory(Arc::new(DroppingSignedAssistantMemory {
14942 messages: RwLock::new(Vec::new()),
14943 }))
14944 .tool(Arc::new(CalculatorTool::new()))
14945 .build()
14946 .unwrap();
14947
14948 let error = agent.chat("What is 3 + 4?").await.unwrap_err();
14949
14950 assert!(
14951 error
14952 .to_string()
14953 .contains("removed before provider continuation")
14954 );
14955 assert!(agent.tool_call_history.read().is_empty());
14956 }
14957
14958 #[tokio::test]
14959 async fn sequential_signed_history_validates_every_prior_exchange() {
14960 let mut mock = MockLLMProvider::new("native-sequential-memory");
14961 mock.set_tool_choice(Some(ToolChoice::Required));
14962 mock.add_response(signed_calculator_response(
14963 "seq-exchange-1",
14964 "seq-call-1",
14965 "1 + 1",
14966 ));
14967 mock.add_response(signed_calculator_response(
14968 "seq-exchange-2",
14969 "seq-call-2",
14970 "2 + 2",
14971 ));
14972 let agent = AgentBuilder::new()
14973 .system_prompt("Use the calculator sequentially.")
14974 .llm(Arc::new(mock))
14975 .memory(Arc::new(DroppingEarlierSequentialMemory {
14976 messages: RwLock::new(Vec::new()),
14977 signed_seen: std::sync::atomic::AtomicUsize::new(0),
14978 }))
14979 .tool(Arc::new(CalculatorTool::new()))
14980 .build()
14981 .unwrap();
14982
14983 let error = agent.chat("Calculate twice.").await.unwrap_err();
14984
14985 assert!(error.to_string().contains("seq-exchange-1"));
14986 assert_eq!(agent.tool_call_history.read().len(), 1);
14987 }
14988
14989 #[tokio::test]
14990 async fn post_transition_signed_hitl_rejection_stops_before_continuation() {
14991 let mut native = MockLLMProvider::new("post-transition-native");
14992 native.set_tool_choice(Some(ToolChoice::Auto));
14993 let call = ToolCall {
14994 id: "post-transition-call".to_string(),
14995 name: "echo".to_string(),
14996 arguments: serde_json::json!({"message": "hello"}),
14997 };
14998 let state = ai_agents_core::NativeProviderState::new(
14999 "post-transition-exchange",
15000 "fixture",
15001 "native-tools",
15002 ai_agents_core::NativeProviderTarget::new("https://fixture.invalid/", "fixture-model")
15003 .unwrap(),
15004 serde_json::json!({
15005 "role": "model",
15006 "parts": [{
15007 "functionCall": {"name": "echo", "args": {"message": "hello"}},
15008 "thoughtSignature": "post-transition-signature"
15009 }]
15010 }),
15011 vec![ai_agents_core::NativeCallBinding::new("post-transition-call", 0).unwrap()],
15012 )
15013 .unwrap();
15014 native.add_response(
15015 LLMResponse::new("", FinishReason::ToolCall)
15016 .with_provider_state(state)
15017 .unwrap()
15018 .with_tool_calls(vec![call])
15019 .unwrap(),
15020 );
15021 let observed_native = native.clone();
15022 let yaml = r#"
15023name: PostTransitionNativeReject
15024system_prompt: test
15025tools: [echo]
15026hitl:
15027 tools:
15028 echo:
15029 require_approval: true
15030states:
15031 initial: intake
15032 states:
15033 intake:
15034 prompt: intake
15035 transitions:
15036 - to: active
15037 guard:
15038 context:
15039 route:
15040 eq: active
15041 active:
15042 prompt: active
15043 llm: native
15044"#;
15045 let agent = AgentBuilder::from_yaml(yaml)
15046 .unwrap()
15047 .llm(Arc::new(mock_with_response("stale intake response")))
15048 .llm_alias("native", Arc::new(native))
15049 .auto_configure_features()
15050 .unwrap()
15051 .build()
15052 .unwrap();
15053 agent
15054 .set_context("route", serde_json::json!("active"))
15055 .unwrap();
15056
15057 let error = agent.chat("move to active").await.unwrap_err();
15058
15059 assert!(matches!(error, AgentError::HITLRejected(_)));
15060 assert_eq!(observed_native.call_count(), 1);
15061 }
15062
15063 #[test]
15064 fn runtime_overflow_removes_a_past_signed_user_turn_as_one_prefix() {
15065 let call = ToolCall {
15066 id: "overflow-call".to_string(),
15067 name: "calculator".to_string(),
15068 arguments: serde_json::json!({"expression": "1 + 1"}),
15069 };
15070 let state = ai_agents_core::NativeProviderState::new(
15071 "overflow-exchange",
15072 "google",
15073 "generateContent",
15074 ai_agents_core::NativeProviderTarget::new("https://example.invalid/", "gemini-3")
15075 .unwrap(),
15076 serde_json::json!({
15077 "role": "model",
15078 "parts": [{
15079 "functionCall": {"name": "calculator", "args": {"expression": "1 + 1"}},
15080 "thoughtSignature": "overflow-signature"
15081 }]
15082 }),
15083 vec![ai_agents_core::NativeCallBinding::new("overflow-call", 0).unwrap()],
15084 )
15085 .unwrap();
15086 let call_marker = ai_agents_core::encode_native_tool_call_markers(
15087 std::slice::from_ref(&call),
15088 Some(&state),
15089 )
15090 .unwrap();
15091 let result_marker = ai_agents_core::encode_native_tool_result_marker(
15092 &call,
15093 serde_json::json!({"result": 2}),
15094 )
15095 .unwrap();
15096 let history = vec![
15097 ChatMessage::user("old question"),
15098 ChatMessage::assistant(call_marker),
15099 ChatMessage::function("calculator", result_marker),
15100 ChatMessage::assistant("old answer"),
15101 ChatMessage::user("new question"),
15102 ];
15103
15104 let removable = RuntimeAgent::native_safe_prefix_at_least(&history, 1).unwrap();
15105
15106 assert_eq!(removable, 4);
15107 }
15108
15109 #[test]
15110 fn auxiliary_projection_does_not_interpret_user_marker_text() {
15111 let user_text = serde_json::json!({
15112 "_ai_agents_native_tool_call": true,
15113 "id": "",
15114 "tool": "user-data",
15115 "arguments": {}
15116 })
15117 .to_string();
15118
15119 let projected =
15120 RuntimeAgent::readable_native_messages(vec![ChatMessage::user(&user_text)]).unwrap();
15121
15122 assert_eq!(projected[0].content, user_text);
15123 }
15124
15125 #[tokio::test]
15126 async fn terminal_provider_history_error_skips_retry_and_static_fallback() {
15127 let calls = Arc::new(std::sync::atomic::AtomicU32::new(0));
15128 let recovery = RecoveryManager::new(ai_agents_recovery::ErrorRecoveryConfig {
15129 default: ai_agents_recovery::RetryConfig {
15130 max_retries: 3,
15131 ..Default::default()
15132 },
15133 llm: ai_agents_recovery::LLMRecoveryConfig {
15134 on_failure: LLMFailureAction::FallbackResponse {
15135 message: "must not be returned".to_string(),
15136 },
15137 ..Default::default()
15138 },
15139 ..Default::default()
15140 });
15141 let agent = AgentBuilder::new()
15142 .system_prompt("Reject corrupted native history.")
15143 .llm(Arc::new(TerminalHistoryProvider {
15144 calls: Arc::clone(&calls),
15145 }))
15146 .recovery_manager(recovery)
15147 .build()
15148 .unwrap();
15149
15150 let error = agent.chat("continue").await.unwrap_err();
15151
15152 assert!(
15153 error
15154 .to_string()
15155 .contains("native history integrity failure")
15156 );
15157 assert_eq!(calls.load(Ordering::SeqCst), 1);
15158 }
15159
15160 #[tokio::test]
15161 async fn prompt_fallback_uses_one_corrective_retry() {
15162 let mut mock = MockLLMProvider::new("prompt-required");
15163 mock.set_tool_choice(Some(ToolChoice::Required));
15164 mock.set_native_tool_support(false);
15165 mock.set_responses(
15166 vec![
15167 "I can calculate that.".to_string(),
15168 r#"{"tool":"calculator","arguments":{"expression":"2 + 2"}}"#.to_string(),
15169 "The answer is 4.".to_string(),
15170 ],
15171 false,
15172 );
15173 let observed = mock.clone();
15174 let agent = AgentBuilder::new()
15175 .system_prompt("Use tools.")
15176 .llm(Arc::new(mock))
15177 .tool(Arc::new(CalculatorTool::new()))
15178 .build()
15179 .unwrap();
15180
15181 let response = agent.chat("What is 2 + 2?").await.unwrap();
15182
15183 assert_eq!(response.content, "The answer is 4.");
15184 assert_eq!(observed.call_count(), 3);
15185 let corrective = &observed.call_history()[1].messages;
15186 assert!(
15187 corrective
15188 .last()
15189 .unwrap()
15190 .content
15191 .contains("previous response")
15192 );
15193 }
15194
15195 #[tokio::test]
15196 async fn prompt_fallback_fails_after_one_noncompliant_retry() {
15197 let mut mock = MockLLMProvider::new("prompt-required-failure");
15198 mock.set_tool_choice(Some(ToolChoice::Required));
15199 mock.set_native_tool_support(false);
15200 mock.set_responses(
15201 vec!["No tool.".to_string(), "Still no tool.".to_string()],
15202 false,
15203 );
15204 let observed = mock.clone();
15205 let agent = AgentBuilder::new()
15206 .system_prompt("Use tools.")
15207 .llm(Arc::new(mock))
15208 .tool(Arc::new(CalculatorTool::new()))
15209 .build()
15210 .unwrap();
15211
15212 let error = agent.chat("What is 2 + 2?").await.unwrap_err();
15213
15214 assert!(error.to_string().contains("one corrective retry"));
15215 assert_eq!(observed.call_count(), 2);
15216 }
15217
15218 #[tokio::test]
15219 async fn specific_choice_cannot_widen_the_effective_grant() {
15220 let mut mock = MockLLMProvider::new("specific-outside-grant");
15221 mock.set_tool_choice(Some(ToolChoice::Specific("random".to_string())));
15222 let observed = mock.clone();
15223 let agent = AgentBuilder::new()
15224 .system_prompt("Use tools.")
15225 .llm(Arc::new(mock))
15226 .tool(Arc::new(CalculatorTool::new()))
15227 .build()
15228 .unwrap();
15229
15230 let error = agent.chat("Generate a value.").await.unwrap_err();
15231
15232 assert!(error.to_string().contains("is not registered"));
15233 assert_eq!(observed.call_count(), 0);
15234 }
15235
15236 #[tokio::test]
15237 async fn none_choice_exposes_no_tool_protocol() {
15238 let mut mock = MockLLMProvider::new("no-tools");
15239 mock.set_tool_choice(Some(ToolChoice::None));
15240 mock.set_response(r#"{"tool":"calculator","arguments":{"expression":"2 + 2"}}"#);
15241 let observed = mock.clone();
15242 let agent = AgentBuilder::new()
15243 .system_prompt("Answer directly.")
15244 .llm(Arc::new(mock))
15245 .tool(Arc::new(CalculatorTool::new()))
15246 .build()
15247 .unwrap();
15248
15249 let response = agent.chat("Hello").await.unwrap();
15250
15251 assert!(response.tool_calls.is_none());
15252 assert_eq!(observed.call_count(), 1);
15253 let call = observed.last_call().unwrap();
15254 assert!(call.request.is_none());
15255 assert!(
15256 call.messages
15257 .iter()
15258 .all(|message| !message.content.contains("Available tools:"))
15259 );
15260 }
15261
15262 struct RuntimeStorage {
15263 capabilities: Box<[StorageCapability]>,
15264 snapshots: RwLock<HashMap<String, AgentSnapshot>>,
15265 metadata: RwLock<HashMap<String, ai_agents_core::SessionMetadata>>,
15266 metadata_save_calls: AtomicU64,
15267 metadata_load_calls: AtomicU64,
15268 fail_metadata_save: AtomicBool,
15269 fail_metadata_load: AtomicBool,
15270 }
15271
15272 impl RuntimeStorage {
15273 fn new(capabilities: impl IntoIterator<Item = StorageCapability>) -> Self {
15274 Self {
15275 capabilities: capabilities.into_iter().collect(),
15276 snapshots: RwLock::new(HashMap::new()),
15277 metadata: RwLock::new(HashMap::new()),
15278 metadata_save_calls: AtomicU64::new(0),
15279 metadata_load_calls: AtomicU64::new(0),
15280 fail_metadata_save: AtomicBool::new(false),
15281 fail_metadata_load: AtomicBool::new(false),
15282 }
15283 }
15284 }
15285
15286 #[async_trait]
15287 impl AgentStorage for RuntimeStorage {
15288 fn supports(&self, capability: StorageCapability) -> bool {
15289 self.capabilities.contains(&capability)
15290 }
15291
15292 async fn save(&self, session_id: &str, snapshot: &AgentSnapshot) -> Result<()> {
15293 self.snapshots
15294 .write()
15295 .insert(session_id.to_string(), snapshot.clone());
15296 Ok(())
15297 }
15298
15299 async fn load(&self, session_id: &str) -> Result<Option<AgentSnapshot>> {
15300 Ok(self.snapshots.read().get(session_id).cloned())
15301 }
15302
15303 async fn delete(&self, session_id: &str) -> Result<()> {
15304 self.snapshots.write().remove(session_id);
15305 Ok(())
15306 }
15307
15308 async fn list_sessions(&self) -> Result<Vec<String>> {
15309 Ok(self.snapshots.read().keys().cloned().collect())
15310 }
15311
15312 async fn save_snapshot_with_metadata(
15313 &self,
15314 session_id: &str,
15315 snapshot: &AgentSnapshot,
15316 metadata: &ai_agents_core::SessionMetadata,
15317 ) -> Result<()> {
15318 self.metadata_save_calls.fetch_add(1, Ordering::SeqCst);
15319 if self.fail_metadata_save.load(Ordering::SeqCst) {
15320 return Err(AgentError::Persistence("metadata save failed".into()));
15321 }
15322 self.snapshots
15323 .write()
15324 .insert(session_id.to_string(), snapshot.clone());
15325 self.metadata
15326 .write()
15327 .insert(session_id.to_string(), metadata.clone());
15328 Ok(())
15329 }
15330
15331 async fn save_metadata(
15332 &self,
15333 session_id: &str,
15334 metadata: &ai_agents_core::SessionMetadata,
15335 ) -> Result<()> {
15336 self.metadata_save_calls.fetch_add(1, Ordering::SeqCst);
15337 if self.fail_metadata_save.load(Ordering::SeqCst) {
15338 return Err(AgentError::Persistence("metadata save failed".into()));
15339 }
15340 self.metadata
15341 .write()
15342 .insert(session_id.to_string(), metadata.clone());
15343 Ok(())
15344 }
15345
15346 async fn load_metadata(
15347 &self,
15348 session_id: &str,
15349 ) -> Result<Option<ai_agents_core::SessionMetadata>> {
15350 self.metadata_load_calls.fetch_add(1, Ordering::SeqCst);
15351 if self.fail_metadata_load.load(Ordering::SeqCst) {
15352 return Err(AgentError::Persistence("metadata load failed".into()));
15353 }
15354 Ok(self.metadata.read().get(session_id).cloned())
15355 }
15356 }
15357
15358 fn runtime_storage_agent() -> RuntimeAgent {
15359 AgentBuilder::new()
15360 .system_prompt("Test runtime storage integration.")
15361 .llm(Arc::new(mock_with_response("done")))
15362 .build()
15363 .unwrap()
15364 }
15365
15366 fn restore_spec(id: &str) -> crate::spec::AgentSpec {
15367 crate::spec::AgentSpec {
15368 name: id.to_string(),
15369 system_prompt: format!("Restore child {id}."),
15370 ..crate::spec::AgentSpec::default()
15371 }
15372 }
15373
15374 fn restore_entry(id: &str) -> ai_agents_core::SpawnedAgentEntry {
15375 ai_agents_core::SpawnedAgentEntry {
15376 id: id.to_string(),
15377 name: id.to_string(),
15378 spec_yaml: serde_yaml::to_string(&restore_spec(id)).unwrap(),
15379 }
15380 }
15381
15382 fn restore_spawner(
15383 storage: Arc<RuntimeStorage>,
15384 max_agents: usize,
15385 ) -> (
15386 Arc<crate::spawner::AgentSpawner>,
15387 Arc<crate::spawner::AgentRegistry>,
15388 ) {
15389 let mut llms = LLMRegistry::new();
15390 llms.register("default", Arc::new(mock_with_response("done")));
15391 (
15392 Arc::new(
15393 crate::spawner::AgentSpawner::new()
15394 .with_shared_llms(llms)
15395 .with_shared_storage(storage)
15396 .with_max_agents(max_agents),
15397 ),
15398 Arc::new(crate::spawner::AgentRegistry::new()),
15399 )
15400 }
15401
15402 async fn save_restore_target(
15403 parent: &RuntimeAgent,
15404 storage: &RuntimeStorage,
15405 session_id: &str,
15406 entries: Vec<ai_agents_core::SpawnedAgentEntry>,
15407 ) {
15408 let mut snapshot = parent.save_state().await.unwrap();
15409 snapshot.spawned_agents = Some(entries);
15410 storage.save(session_id, &snapshot).await.unwrap();
15411 storage
15412 .save_metadata(session_id, &ai_agents_core::SessionMetadata::default())
15413 .await
15414 .unwrap();
15415 }
15416
15417 #[tokio::test]
15418 async fn storage_init_requires_storage_for_actor_facts() {
15419 let facts = ai_agents_facts::FactsConfig {
15420 enabled: true,
15421 ..Default::default()
15422 };
15423 let agent = runtime_storage_agent().with_facts_config(None, Some(facts));
15424
15425 let error = agent.init_storage().await.unwrap_err();
15426 assert!(matches!(
15427 error,
15428 AgentError::Config(message)
15429 if message.contains("actor facts or actor memory")
15430 && message.contains("none is configured or injected")
15431 ));
15432 }
15433
15434 #[tokio::test]
15435 async fn storage_init_validates_actor_facts_capability() {
15436 let storage = Arc::new(RuntimeStorage::new([StorageCapability::Snapshot]));
15437 let actor_memory = ai_agents_facts::ActorMemoryConfig {
15438 enabled: true,
15439 ..Default::default()
15440 };
15441 let agent = runtime_storage_agent()
15442 .with_storage(storage)
15443 .with_facts_config(Some(actor_memory), None);
15444
15445 assert!(matches!(
15446 agent.init_storage().await,
15447 Err(AgentError::UnsupportedStorageCapability(
15448 StorageCapability::ActorFacts
15449 ))
15450 ));
15451 }
15452
15453 #[tokio::test]
15454 async fn blocking_chat_rejects_unsupported_required_storage() {
15455 let storage = Arc::new(RuntimeStorage::new([StorageCapability::Snapshot]));
15456 let facts = ai_agents_facts::FactsConfig {
15457 enabled: true,
15458 ..Default::default()
15459 };
15460 let agent = runtime_storage_agent()
15461 .with_storage(storage)
15462 .with_facts_config(None, Some(facts));
15463
15464 assert!(matches!(
15465 agent.chat("hello").await,
15466 Err(AgentError::UnsupportedStorageCapability(
15467 StorageCapability::ActorFacts
15468 ))
15469 ));
15470 }
15471
15472 #[tokio::test]
15473 async fn streaming_chat_rejects_unsupported_required_storage_before_stream_creation() {
15474 let storage = Arc::new(RuntimeStorage::new([StorageCapability::Snapshot]));
15475 let config = ai_agents_relationships::RelationshipConfig {
15476 enabled: true,
15477 ..Default::default()
15478 };
15479 let manager = Arc::new(RelationshipManager::from_config(config).unwrap());
15480 let agent = runtime_storage_agent()
15481 .with_storage(storage)
15482 .with_relationships(manager);
15483
15484 assert!(matches!(
15485 agent.chat_stream("hello").await,
15486 Err(AgentError::UnsupportedStorageCapability(
15487 StorageCapability::ActorRelationships
15488 ))
15489 ));
15490 }
15491
15492 #[tokio::test]
15493 async fn storage_init_completes_facts_for_injected_storage() {
15494 let storage = Arc::new(RuntimeStorage::new([
15495 StorageCapability::Snapshot,
15496 StorageCapability::ActorFacts,
15497 ]));
15498 let facts = ai_agents_facts::FactsConfig {
15499 enabled: true,
15500 ..Default::default()
15501 };
15502 let agent = runtime_storage_agent()
15503 .with_storage(storage)
15504 .with_facts_config(None, Some(facts));
15505
15506 agent.init_storage().await.unwrap();
15507 assert!(agent.fact_store().is_some());
15508 }
15509
15510 #[tokio::test]
15511 async fn storage_init_requires_storage_for_persistent_relationships() {
15512 let config = ai_agents_relationships::RelationshipConfig {
15513 enabled: true,
15514 ..Default::default()
15515 };
15516 let manager = Arc::new(RelationshipManager::from_config(config).unwrap());
15517 let agent = runtime_storage_agent().with_relationships(manager);
15518
15519 let error = agent.init_storage().await.unwrap_err();
15520 assert!(matches!(
15521 error,
15522 AgentError::Config(message)
15523 if message.contains("persistent relationships")
15524 && message.contains("none is configured or injected")
15525 ));
15526 }
15527
15528 #[tokio::test]
15529 async fn storage_init_validates_persistent_relationships_capability() {
15530 let storage = Arc::new(RuntimeStorage::new([StorageCapability::Snapshot]));
15531 let config = ai_agents_relationships::RelationshipConfig {
15532 enabled: true,
15533 ..Default::default()
15534 };
15535 let manager = Arc::new(RelationshipManager::from_config(config).unwrap());
15536 let agent = runtime_storage_agent()
15537 .with_storage(storage)
15538 .with_relationships(manager);
15539
15540 assert!(matches!(
15541 agent.init_storage().await,
15542 Err(AgentError::UnsupportedStorageCapability(
15543 StorageCapability::ActorRelationships
15544 ))
15545 ));
15546 }
15547
15548 #[tokio::test]
15549 async fn session_restore_updates_identity_and_clears_stale_actor_binding() {
15550 let storage = Arc::new(RuntimeStorage::new([
15551 StorageCapability::Snapshot,
15552 StorageCapability::SessionMetadata,
15553 ]));
15554 let agent = runtime_storage_agent().with_storage(storage.clone());
15555 agent.set_actor_id("old-actor").unwrap();
15556 agent.save_session("old").await.unwrap();
15557 storage
15558 .save("target", &agent.save_state().await.unwrap())
15559 .await
15560 .unwrap();
15561 storage
15562 .save_metadata("target", &ai_agents_core::SessionMetadata::default())
15563 .await
15564 .unwrap();
15565
15566 assert!(agent.load_session("target").await.unwrap());
15567
15568 assert_eq!(agent.current_session_id.read().as_deref(), Some("target"));
15569 assert_eq!(agent.actor_id(), None);
15570 }
15571
15572 #[tokio::test]
15573 async fn complete_restore_reconciles_growth_shrink_and_empty_topologies() {
15574 let storage = Arc::new(RuntimeStorage::new([
15575 StorageCapability::Snapshot,
15576 StorageCapability::SessionMetadata,
15577 ]));
15578 let (spawner, registry) = restore_spawner(storage.clone(), 3);
15579 let parent = runtime_storage_agent()
15580 .with_storage(storage.clone())
15581 .with_spawner_handles(Arc::clone(&spawner), Arc::clone(®istry));
15582
15583 for id in ["a", "b"] {
15584 let spawned = spawner
15585 .spawn_with_id(id.to_string(), restore_spec(id))
15586 .await
15587 .unwrap();
15588 spawned.agent.save_session("grow").await.unwrap();
15589 registry.register(spawned).await.unwrap();
15590 }
15591 let staged_c = crate::spawner::storage::NamespacedStorage::new(storage.clone(), "c");
15592 staged_c
15593 .save("grow", &AgentSnapshot::new("c".into()))
15594 .await
15595 .unwrap();
15596 staged_c
15597 .save_metadata("grow", &ai_agents_core::SessionMetadata::default())
15598 .await
15599 .unwrap();
15600 save_restore_target(
15601 &parent,
15602 storage.as_ref(),
15603 "grow",
15604 vec![restore_entry("a"), restore_entry("b"), restore_entry("c")],
15605 )
15606 .await;
15607
15608 assert_eq!(parent.restore_session_full("grow").await.unwrap(), 3);
15609 assert_eq!(registry.count(), 3);
15610 assert!(registry.contains("c"));
15611 assert_eq!(spawner.spawned_count(), 3);
15612
15613 for id in ["a", "b"] {
15614 registry
15615 .get(id)
15616 .unwrap()
15617 .save_session("shrink")
15618 .await
15619 .unwrap();
15620 }
15621 save_restore_target(
15622 &parent,
15623 storage.as_ref(),
15624 "shrink",
15625 vec![restore_entry("a"), restore_entry("b")],
15626 )
15627 .await;
15628
15629 assert_eq!(parent.restore_session_full("shrink").await.unwrap(), 2);
15630 assert_eq!(registry.count(), 2);
15631 assert!(!registry.contains("c"));
15632 assert_eq!(spawner.spawned_count(), 2);
15633
15634 save_restore_target(&parent, storage.as_ref(), "empty", Vec::new()).await;
15635
15636 assert_eq!(parent.restore_session_full("empty").await.unwrap(), 0);
15637 assert_eq!(registry.count(), 0);
15638 assert_eq!(spawner.spawned_count(), 0);
15639 assert_eq!(parent.current_session_id.read().as_deref(), Some("empty"));
15640 }
15641
15642 #[tokio::test]
15643 async fn storage_session_metadata_is_called_only_when_advertised() {
15644 let storage = Arc::new(RuntimeStorage::new([StorageCapability::Snapshot]));
15645 storage.fail_metadata_save.store(true, Ordering::SeqCst);
15646 storage.fail_metadata_load.store(true, Ordering::SeqCst);
15647 let agent = runtime_storage_agent().with_storage(storage.clone());
15648
15649 agent.save_session("session").await.unwrap();
15650 assert!(agent.load_session("session").await.unwrap());
15651 assert_eq!(storage.metadata_save_calls.load(Ordering::SeqCst), 0);
15652 assert_eq!(storage.metadata_load_calls.load(Ordering::SeqCst), 0);
15653 }
15654
15655 #[cfg(feature = "sqlite")]
15656 #[tokio::test]
15657 async fn sqlite_runtime_save_filter_reopen_and_reload_stay_consistent() {
15658 let directory =
15659 std::env::temp_dir().join(format!("ai-agents-runtime-sqlite-{}", uuid::Uuid::new_v4()));
15660 let path = directory.join("sessions.sqlite");
15661 let path_string = path.to_string_lossy().into_owned();
15662 let storage = Arc::new(
15663 ai_agents_storage::SqliteStorage::new(&path_string)
15664 .await
15665 .unwrap(),
15666 );
15667 let agent = runtime_storage_agent().with_storage(storage.clone());
15668 agent.set_session_metadata(ai_agents_core::SessionMetadata {
15669 tags: vec!["initial".into()],
15670 ..Default::default()
15671 });
15672 agent.chat("persist this turn").await.unwrap();
15673 agent.save_session("session").await.unwrap();
15674
15675 agent.set_session_metadata(ai_agents_core::SessionMetadata {
15676 tags: vec!["updated".into()],
15677 ..Default::default()
15678 });
15679 agent.save_session("session").await.unwrap();
15680 assert!(
15681 agent
15682 .list_sessions_filtered(&ai_agents_core::SessionFilter {
15683 tags: Some(vec!["initial".into()]),
15684 ..Default::default()
15685 })
15686 .await
15687 .unwrap()
15688 .is_empty()
15689 );
15690 assert_eq!(
15691 agent
15692 .list_sessions_filtered(&ai_agents_core::SessionFilter {
15693 tags: Some(vec!["updated".into()]),
15694 ..Default::default()
15695 })
15696 .await
15697 .unwrap()
15698 .len(),
15699 1
15700 );
15701 drop(agent);
15702 storage.close().await;
15703 drop(storage);
15704
15705 let reopened_storage = Arc::new(
15706 ai_agents_storage::SqliteStorage::new(&path_string)
15707 .await
15708 .unwrap(),
15709 );
15710 let restored = runtime_storage_agent().with_storage(reopened_storage.clone());
15711 assert!(restored.load_session("session").await.unwrap());
15712 assert_eq!(restored.session_metadata().tags, vec!["updated"]);
15713 assert_eq!(
15714 restored.current_session_id.read().as_deref(),
15715 Some("session")
15716 );
15717 assert!(restored.save_state().await.unwrap().memory.messages.len() >= 2);
15718 assert_eq!(
15719 restored
15720 .list_sessions_filtered(&ai_agents_core::SessionFilter {
15721 tags: Some(vec!["updated".into()]),
15722 ..Default::default()
15723 })
15724 .await
15725 .unwrap()
15726 .len(),
15727 1
15728 );
15729
15730 drop(restored);
15731 reopened_storage.close().await;
15732 drop(reopened_storage);
15733 crate::remove_sqlite_test_directory(&directory)
15734 .await
15735 .unwrap();
15736 }
15737
15738 #[tokio::test]
15739 async fn storage_session_metadata_backend_failures_propagate() {
15740 let storage = Arc::new(RuntimeStorage::new([
15741 StorageCapability::Snapshot,
15742 StorageCapability::SessionMetadata,
15743 ]));
15744 let agent = runtime_storage_agent().with_storage(storage.clone());
15745
15746 agent.save_session("session").await.unwrap();
15747 storage
15748 .save("target", &agent.save_state().await.unwrap())
15749 .await
15750 .unwrap();
15751 storage.fail_metadata_load.store(true, Ordering::SeqCst);
15752 assert!(matches!(
15753 agent.load_session("target").await,
15754 Err(AgentError::Persistence(message)) if message == "metadata load failed"
15755 ));
15756 assert_eq!(agent.current_session_id.read().as_deref(), Some("session"));
15757
15758 storage.fail_metadata_save.store(true, Ordering::SeqCst);
15759 assert!(matches!(
15760 agent.save_session("session").await,
15761 Err(AgentError::Persistence(message)) if message == "metadata save failed"
15762 ));
15763 }
15764
15765 struct ProviderFutureDropSignal {
15766 dropped: Arc<AtomicBool>,
15767 }
15768
15769 impl Drop for ProviderFutureDropSignal {
15770 fn drop(&mut self) {
15771 self.dropped.store(true, Ordering::SeqCst);
15772 }
15773 }
15774
15775 struct BufferedLockingProvider {
15776 lock: Arc<tokio::sync::Mutex<()>>,
15777 stream_started: Arc<tokio::sync::Notify>,
15778 stream_dropped: Arc<AtomicBool>,
15779 committed_after_drop: Arc<AtomicBool>,
15780 }
15781
15782 #[async_trait]
15783 impl LLMProvider for BufferedLockingProvider {
15784 async fn complete(
15785 &self,
15786 _messages: &[ChatMessage],
15787 _config: Option<&LLMConfig>,
15788 ) -> std::result::Result<LLMResponse, LLMError> {
15789 let _guard = self.lock.lock().await;
15790 self.committed_after_drop
15791 .store(self.stream_dropped.load(Ordering::SeqCst), Ordering::SeqCst);
15792 Ok(LLMResponse::new(
15793 "Committed technical response.",
15794 FinishReason::Stop,
15795 ))
15796 }
15797
15798 async fn complete_stream(
15799 &self,
15800 _messages: &[ChatMessage],
15801 _config: Option<&LLMConfig>,
15802 ) -> std::result::Result<
15803 Box<dyn Stream<Item = std::result::Result<LLMChunk, LLMError>> + Unpin + Send>,
15804 LLMError,
15805 > {
15806 let _guard = self.lock.lock().await;
15807 let _drop_signal = ProviderFutureDropSignal {
15808 dropped: Arc::clone(&self.stream_dropped),
15809 };
15810 self.stream_started.notify_one();
15811 std::future::pending().await
15812 }
15813
15814 fn provider_name(&self) -> &str {
15815 "buffered-locking"
15816 }
15817
15818 fn supports(&self, _feature: LLMFeature) -> bool {
15819 false
15820 }
15821 }
15822
15823 struct PendingDropStream {
15824 dropped: Arc<AtomicBool>,
15825 dropped_notify: Arc<tokio::sync::Notify>,
15826 }
15827
15828 impl Stream for PendingDropStream {
15829 type Item = std::result::Result<LLMChunk, LLMError>;
15830
15831 fn poll_next(
15832 self: Pin<&mut Self>,
15833 _cx: &mut std::task::Context<'_>,
15834 ) -> std::task::Poll<Option<Self::Item>> {
15835 std::task::Poll::Pending
15836 }
15837 }
15838
15839 impl Drop for PendingDropStream {
15840 fn drop(&mut self) {
15841 self.dropped.store(true, Ordering::SeqCst);
15842 self.dropped_notify.notify_one();
15843 }
15844 }
15845
15846 struct EstablishedStreamProvider {
15847 stream_started: Arc<tokio::sync::Notify>,
15848 stream_dropped: Arc<AtomicBool>,
15849 stream_dropped_notify: Arc<tokio::sync::Notify>,
15850 committed_after_drop: Arc<AtomicBool>,
15851 }
15852
15853 #[async_trait]
15854 impl LLMProvider for EstablishedStreamProvider {
15855 async fn complete(
15856 &self,
15857 _messages: &[ChatMessage],
15858 _config: Option<&LLMConfig>,
15859 ) -> std::result::Result<LLMResponse, LLMError> {
15860 if !self.stream_dropped.load(Ordering::SeqCst) {
15861 self.stream_dropped_notify.notified().await;
15862 }
15863 self.committed_after_drop
15864 .store(self.stream_dropped.load(Ordering::SeqCst), Ordering::SeqCst);
15865 Ok(LLMResponse::new(
15866 "Committed technical response.",
15867 FinishReason::Stop,
15868 ))
15869 }
15870
15871 async fn complete_stream(
15872 &self,
15873 _messages: &[ChatMessage],
15874 _config: Option<&LLMConfig>,
15875 ) -> std::result::Result<
15876 Box<dyn Stream<Item = std::result::Result<LLMChunk, LLMError>> + Unpin + Send>,
15877 LLMError,
15878 > {
15879 self.stream_started.notify_one();
15880 Ok(Box::new(PendingDropStream {
15881 dropped: Arc::clone(&self.stream_dropped),
15882 dropped_notify: Arc::clone(&self.stream_dropped_notify),
15883 }))
15884 }
15885
15886 fn provider_name(&self) -> &str {
15887 "established-stream"
15888 }
15889
15890 fn supports(&self, _feature: LLMFeature) -> bool {
15891 false
15892 }
15893 }
15894
15895 struct FirstCallLockingProvider {
15896 lock: Arc<tokio::sync::Mutex<()>>,
15897 first_started: Arc<tokio::sync::Notify>,
15898 first_dropped: Arc<AtomicBool>,
15899 committed_after_drop: Arc<AtomicBool>,
15900 calls: AtomicU64,
15901 }
15902
15903 #[async_trait]
15904 impl LLMProvider for FirstCallLockingProvider {
15905 async fn complete(
15906 &self,
15907 _messages: &[ChatMessage],
15908 _config: Option<&LLMConfig>,
15909 ) -> std::result::Result<LLMResponse, LLMError> {
15910 let _guard = self.lock.lock().await;
15911 let call = self.calls.fetch_add(1, Ordering::SeqCst);
15912 if call == 0 {
15913 let _drop_signal = ProviderFutureDropSignal {
15914 dropped: Arc::clone(&self.first_dropped),
15915 };
15916 self.first_started.notify_one();
15917 return std::future::pending().await;
15918 }
15919 self.committed_after_drop
15920 .store(self.first_dropped.load(Ordering::SeqCst), Ordering::SeqCst);
15921 Ok(LLMResponse::new(
15922 "Committed technical response.",
15923 FinishReason::Stop,
15924 ))
15925 }
15926
15927 async fn complete_stream(
15928 &self,
15929 _messages: &[ChatMessage],
15930 _config: Option<&LLMConfig>,
15931 ) -> std::result::Result<
15932 Box<dyn Stream<Item = std::result::Result<LLMChunk, LLMError>> + Unpin + Send>,
15933 LLMError,
15934 > {
15935 Err(LLMError::Other(
15936 "streaming is not used in this test".to_string(),
15937 ))
15938 }
15939
15940 fn provider_name(&self) -> &str {
15941 "first-call-locking"
15942 }
15943
15944 fn supports(&self, _feature: LLMFeature) -> bool {
15945 false
15946 }
15947 }
15948
15949 struct RoutingAfterProviderStart {
15950 provider_started: Arc<tokio::sync::Notify>,
15951 }
15952
15953 #[async_trait]
15954 impl LLMProvider for RoutingAfterProviderStart {
15955 async fn complete(
15956 &self,
15957 _messages: &[ChatMessage],
15958 _config: Option<&LLMConfig>,
15959 ) -> std::result::Result<LLMResponse, LLMError> {
15960 self.provider_started.notified().await;
15961 Ok(LLMResponse::new("1", FinishReason::Stop))
15962 }
15963
15964 async fn complete_stream(
15965 &self,
15966 _messages: &[ChatMessage],
15967 _config: Option<&LLMConfig>,
15968 ) -> std::result::Result<
15969 Box<dyn Stream<Item = std::result::Result<LLMChunk, LLMError>> + Unpin + Send>,
15970 LLMError,
15971 > {
15972 Err(LLMError::Other(
15973 "streaming is not used in this test".to_string(),
15974 ))
15975 }
15976
15977 fn provider_name(&self) -> &str {
15978 "routing-after-start"
15979 }
15980
15981 fn supports(&self, _feature: LLMFeature) -> bool {
15982 false
15983 }
15984 }
15985
15986 struct ResponseCountingHooks {
15988 responses: Arc<std::sync::atomic::AtomicUsize>,
15989 }
15990
15991 struct RootTurnProbeProvider {
15993 complete_entered: tokio::sync::mpsc::UnboundedSender<()>,
15994 }
15995
15996 struct ResponseChatHooks {
15998 target: parking_lot::Mutex<Option<Weak<RuntimeAgent>>>,
15999 invoked: AtomicBool,
16000 nested_result: parking_lot::Mutex<Option<std::result::Result<String, String>>>,
16001 }
16002
16003 struct ConcurrentResponseHooks {
16005 registry: Weak<crate::spawner::AgentRegistry>,
16006 child_id: String,
16007 invoked: AtomicBool,
16008 nested_result: parking_lot::Mutex<Option<std::result::Result<String, String>>>,
16009 }
16010
16011 struct RetryDeadlineTool {
16013 calls: Arc<std::sync::atomic::AtomicUsize>,
16014 deadlines: Arc<parking_lot::Mutex<Vec<chrono::DateTime<chrono::Utc>>>>,
16015 remaining_ms: Arc<parking_lot::Mutex<Vec<i64>>>,
16016 }
16017
16018 struct ToolLifecycleRecordingHooks {
16020 events: parking_lot::Mutex<Vec<String>>,
16021 records: parking_lot::Mutex<Vec<ToolExecutionRecord>>,
16022 }
16023
16024 impl ToolLifecycleRecordingHooks {
16025 fn new() -> Self {
16027 Self {
16028 events: parking_lot::Mutex::new(Vec::new()),
16029 records: parking_lot::Mutex::new(Vec::new()),
16030 }
16031 }
16032
16033 fn events(&self) -> Vec<String> {
16035 self.events.lock().clone()
16036 }
16037
16038 fn records(&self) -> Vec<ToolExecutionRecord> {
16040 self.records.lock().clone()
16041 }
16042 }
16043
16044 struct ContextEchoTool;
16046
16047 #[async_trait]
16048 impl LLMProvider for RootTurnProbeProvider {
16049 async fn complete(
16050 &self,
16051 _messages: &[ChatMessage],
16052 _config: Option<&LLMConfig>,
16053 ) -> std::result::Result<LLMResponse, LLMError> {
16054 let _ = self.complete_entered.send(());
16055 Ok(LLMResponse::new("blocking complete", FinishReason::Stop))
16056 }
16057
16058 async fn complete_stream(
16059 &self,
16060 _messages: &[ChatMessage],
16061 _config: Option<&LLMConfig>,
16062 ) -> std::result::Result<
16063 Box<dyn Stream<Item = std::result::Result<LLMChunk, LLMError>> + Unpin + Send>,
16064 LLMError,
16065 > {
16066 Ok(Box::new(futures::stream::iter(vec![Ok(
16067 LLMChunk::final_chunk("stream complete", FinishReason::Stop, None),
16068 )])))
16069 }
16070
16071 fn provider_name(&self) -> &str {
16072 "root-turn-probe"
16073 }
16074
16075 fn supports(&self, feature: LLMFeature) -> bool {
16076 matches!(feature, LLMFeature::Streaming)
16077 }
16078 }
16079
16080 #[async_trait]
16081 impl ai_agents_core::Tool for ContextEchoTool {
16082 fn id(&self) -> &str {
16083 "context_echo"
16084 }
16085
16086 fn name(&self) -> &str {
16087 "Context Echo"
16088 }
16089
16090 fn description(&self) -> &str {
16091 "Returns selected execution context fields."
16092 }
16093
16094 fn input_schema(&self) -> Value {
16095 serde_json::json!({"type": "object"})
16096 }
16097
16098 fn policy_bindings(&self) -> ai_agents_core::ToolPolicyBindings {
16099 ai_agents_core::ToolPolicyBindings {
16100 path_fields: vec![ai_agents_core::PathPolicyBinding::read("path")],
16101 result_limit_fields: vec![ai_agents_core::ResultLimitBinding::new(
16102 "max_results",
16103 ai_agents_core::ResultLimitKind::MaxResults,
16104 )],
16105 ..Default::default()
16106 }
16107 }
16108
16109 async fn execute(
16110 &self,
16111 _args: Value,
16112 ctx: ai_agents_core::ToolExecutionContext,
16113 ) -> ToolResult {
16114 ToolResult::ok(
16115 serde_json::json!({
16116 "requested_name": ctx.requested_name,
16117 "canonical_id": ctx.canonical_id,
16118 "display_name": ctx.display_name,
16119 "max_results": ctx.limits.max_results,
16120 "custom_config": ctx.custom_config,
16121 })
16122 .to_string(),
16123 )
16124 }
16125 }
16126
16127 #[async_trait]
16128 impl ai_agents_core::Tool for RetryDeadlineTool {
16129 fn id(&self) -> &str {
16130 "retry_deadline"
16131 }
16132
16133 fn name(&self) -> &str {
16134 "Retry Deadline"
16135 }
16136
16137 fn description(&self) -> &str {
16138 "Records one deadline per retry invocation."
16139 }
16140
16141 fn input_schema(&self) -> Value {
16142 serde_json::json!({"type": "object"})
16143 }
16144
16145 fn safety_metadata(&self) -> ai_agents_core::ToolSafetyMetadata {
16146 ai_agents_core::ToolSafetyMetadata::compute()
16147 }
16148
16149 fn classify_call(&self, _args: &Value) -> ai_agents_core::ToolCallClassification {
16150 let mut classification =
16151 ai_agents_core::ToolCallClassification::from_metadata(&self.safety_metadata());
16152 classification.timeout_ms = Some(1_000);
16153 classification.safely_retryable = true;
16154 classification
16155 }
16156
16157 async fn execute(
16159 &self,
16160 _args: Value,
16161 ctx: ai_agents_core::ToolExecutionContext,
16162 ) -> ToolResult {
16163 let deadline = ctx
16164 .deadline
16165 .expect("each invocation must receive a deadline");
16166 self.remaining_ms.lock().push(
16167 deadline
16168 .signed_duration_since(chrono::Utc::now())
16169 .num_milliseconds(),
16170 );
16171 self.deadlines.lock().push(deadline);
16172 let call = self.calls.fetch_add(1, Ordering::SeqCst);
16173 if call == 0 {
16174 tokio::time::sleep(std::time::Duration::from_millis(20)).await;
16175 ToolResult::error("retry")
16176 } else {
16177 ToolResult::ok("done")
16178 }
16179 }
16180 }
16181
16182 struct ClassifiedTimeoutTool {
16184 id: &'static str,
16185 calls: Arc<std::sync::atomic::AtomicUsize>,
16186 timeout_ms: u64,
16187 sleep_ms: u64,
16188 requires_approval: bool,
16189 remaining_ms: Arc<parking_lot::Mutex<Vec<i64>>>,
16190 }
16191
16192 struct ApprovalModifiedTimeoutTool {
16194 calls: Arc<std::sync::atomic::AtomicUsize>,
16195 }
16196
16197 struct SlowTool;
16199
16200 struct FlakyWriteTool {
16202 calls: Arc<std::sync::atomic::AtomicUsize>,
16203 }
16204
16205 struct LockedWriteTool {
16207 active: Arc<std::sync::atomic::AtomicUsize>,
16208 max_active: Arc<std::sync::atomic::AtomicUsize>,
16209 }
16210
16211 struct MultiResourceWriteTool {
16212 active: Arc<std::sync::atomic::AtomicUsize>,
16213 max_active: Arc<std::sync::atomic::AtomicUsize>,
16214 }
16215
16216 #[derive(Clone)]
16217 struct PathMutationGate {
16218 entered: Arc<AtomicBool>,
16219 entered_notify: Arc<tokio::sync::Notify>,
16220 release: Arc<tokio::sync::Notify>,
16221 }
16222
16223 impl PathMutationGate {
16224 fn new() -> Self {
16225 Self {
16226 entered: Arc::new(AtomicBool::new(false)),
16227 entered_notify: Arc::new(tokio::sync::Notify::new()),
16228 release: Arc::new(tokio::sync::Notify::new()),
16229 }
16230 }
16231
16232 async fn wait_until_entered(&self) {
16233 if !self.entered.load(Ordering::SeqCst) {
16234 self.entered_notify.notified().await;
16235 }
16236 }
16237
16238 fn release(&self) {
16239 self.release.notify_one();
16240 }
16241 }
16242
16243 struct BlockingPathMutationTool {
16244 id: &'static str,
16245 path_fields: Vec<ai_agents_core::PathPolicyBinding>,
16246 gate: PathMutationGate,
16247 }
16248
16249 struct NoBindingWriteTool {
16250 active: Arc<std::sync::atomic::AtomicUsize>,
16251 max_active: Arc<std::sync::atomic::AtomicUsize>,
16252 }
16253
16254 struct RecoveryTestTool {
16255 id: String,
16256 succeeds: bool,
16257 calls: Arc<std::sync::atomic::AtomicUsize>,
16258 max_output_chars: Option<usize>,
16259 }
16260
16261 struct BlockingApprovalHandler {
16262 entered: Arc<tokio::sync::Barrier>,
16263 release: Arc<tokio::sync::Notify>,
16264 result: ApprovalResult,
16265 }
16266
16267 struct CountingApprovalHandler {
16268 calls: Arc<std::sync::atomic::AtomicUsize>,
16269 }
16270
16271 struct DriftingFallbackProvider {
16273 refreshed: AtomicBool,
16274 primary_calls: Arc<std::sync::atomic::AtomicUsize>,
16275 secondary_calls: Arc<std::sync::atomic::AtomicUsize>,
16276 }
16277
16278 struct RefreshFallbackProviderHooks {
16280 agent: parking_lot::Mutex<Option<Weak<RuntimeAgent>>>,
16281 lifecycle: Arc<ToolLifecycleRecordingHooks>,
16282 }
16283
16284 struct RuntimeWebFetchTransport {
16285 calls: Arc<std::sync::atomic::AtomicUsize>,
16286 }
16287
16288 struct RuntimeWebFetchResolver;
16289
16290 struct ReentrantToolHooks {
16291 agent: parking_lot::Mutex<Option<Weak<RuntimeAgent>>>,
16292 invoked: AtomicBool,
16293 nested_success: AtomicBool,
16294 }
16295
16296 #[async_trait]
16297 impl ai_agents_core::Tool for ClassifiedTimeoutTool {
16298 fn id(&self) -> &str {
16300 self.id
16301 }
16302
16303 fn name(&self) -> &str {
16305 "Classified Timeout"
16306 }
16307
16308 fn description(&self) -> &str {
16310 "Records and waits under one call-level timeout."
16311 }
16312
16313 fn input_schema(&self) -> Value {
16315 serde_json::json!({"type": "object"})
16316 }
16317
16318 fn classify_call(&self, _args: &Value) -> ai_agents_core::ToolCallClassification {
16320 let mut classification =
16321 ai_agents_core::ToolCallClassification::from_metadata(&self.safety_metadata());
16322 classification.timeout_ms = Some(self.timeout_ms);
16323 classification.requires_approval = self.requires_approval;
16324 classification
16325 }
16326
16327 async fn execute(
16329 &self,
16330 _args: Value,
16331 ctx: ai_agents_core::ToolExecutionContext,
16332 ) -> ToolResult {
16333 self.calls.fetch_add(1, Ordering::SeqCst);
16334 let deadline = ctx
16335 .deadline
16336 .expect("each invocation must receive a deadline");
16337 self.remaining_ms.lock().push(
16338 deadline
16339 .signed_duration_since(chrono::Utc::now())
16340 .num_milliseconds(),
16341 );
16342 tokio::time::sleep(Duration::from_millis(self.sleep_ms)).await;
16343 ToolResult::ok("done")
16344 }
16345 }
16346
16347 #[async_trait]
16348 impl ai_agents_core::Tool for ApprovalModifiedTimeoutTool {
16349 fn id(&self) -> &str {
16351 "approval_modified_timeout"
16352 }
16353
16354 fn name(&self) -> &str {
16356 "Approval Modified Timeout"
16357 }
16358
16359 fn description(&self) -> &str {
16361 "Becomes invalid only after approval modifies its arguments."
16362 }
16363
16364 fn input_schema(&self) -> Value {
16366 serde_json::json!({"type": "object"})
16367 }
16368
16369 fn policy_bindings(&self) -> ai_agents_core::ToolPolicyBindings {
16371 ai_agents_core::ToolPolicyBindings {
16372 path_fields: vec![ai_agents_core::PathPolicyBinding::write("path")],
16373 ..Default::default()
16374 }
16375 }
16376
16377 fn safety_metadata(&self) -> ai_agents_core::ToolSafetyMetadata {
16379 ai_agents_core::ToolSafetyMetadata {
16380 read_only: false,
16381 concurrency_safe: false,
16382 operation: ai_agents_core::ToolOperationKind::Write,
16383 side_effect_level: ai_agents_core::ToolSideEffectLevel::LocalWrite,
16384 requires_network: false,
16385 destructive: false,
16386 open_world: false,
16387 host_dependent: false,
16388 requires_user_interaction: false,
16389 supports_cancellation: true,
16390 default_requires_approval: true,
16391 should_defer_schema: false,
16392 max_output_chars: Some(1024),
16393 max_result_size_chars: Some(1024),
16394 }
16395 }
16396
16397 fn classify_call(&self, args: &Value) -> ai_agents_core::ToolCallClassification {
16399 let mut classification =
16400 ai_agents_core::ToolCallClassification::from_metadata(&self.safety_metadata());
16401 classification.timeout_ms = Some(if args["invalid_timeout"].as_bool() == Some(true) {
16402 u64::MAX
16403 } else {
16404 1_000
16405 });
16406 classification
16407 }
16408
16409 async fn execute(
16411 &self,
16412 _args: Value,
16413 _ctx: ai_agents_core::ToolExecutionContext,
16414 ) -> ToolResult {
16415 self.calls.fetch_add(1, Ordering::SeqCst);
16416 ToolResult::ok("unexpected")
16417 }
16418 }
16419
16420 #[async_trait]
16421 impl ai_agents_core::Tool for SlowTool {
16422 fn id(&self) -> &str {
16423 "slow"
16424 }
16425
16426 fn name(&self) -> &str {
16427 "Slow"
16428 }
16429
16430 fn description(&self) -> &str {
16431 "Waits until cancelled or timed out."
16432 }
16433
16434 fn input_schema(&self) -> Value {
16435 serde_json::json!({"type": "object"})
16436 }
16437
16438 async fn execute(
16439 &self,
16440 _args: Value,
16441 _ctx: ai_agents_core::ToolExecutionContext,
16442 ) -> ToolResult {
16443 tokio::time::sleep(std::time::Duration::from_secs(5)).await;
16444 ToolResult::ok("done")
16445 }
16446 }
16447
16448 #[async_trait]
16449 impl ai_agents_core::Tool for FlakyWriteTool {
16450 fn id(&self) -> &str {
16451 "flaky_write"
16452 }
16453
16454 fn name(&self) -> &str {
16455 "Flaky Write"
16456 }
16457
16458 fn description(&self) -> &str {
16459 "Fails on the first write attempt."
16460 }
16461
16462 fn input_schema(&self) -> Value {
16463 serde_json::json!({"type": "object", "properties": {"path": {"type": "string"}}})
16464 }
16465
16466 fn policy_bindings(&self) -> ai_agents_core::ToolPolicyBindings {
16467 ai_agents_core::ToolPolicyBindings {
16468 path_fields: vec![ai_agents_core::PathPolicyBinding::write("path")],
16469 ..Default::default()
16470 }
16471 }
16472
16473 fn safety_metadata(&self) -> ai_agents_core::ToolSafetyMetadata {
16474 ai_agents_core::ToolSafetyMetadata {
16475 read_only: false,
16476 concurrency_safe: false,
16477 operation: ai_agents_core::ToolOperationKind::Write,
16478 side_effect_level: ai_agents_core::ToolSideEffectLevel::LocalWrite,
16479 requires_network: false,
16480 destructive: false,
16481 open_world: false,
16482 host_dependent: false,
16483 requires_user_interaction: false,
16484 supports_cancellation: true,
16485 default_requires_approval: false,
16486 should_defer_schema: false,
16487 max_output_chars: Some(1024),
16488 max_result_size_chars: Some(1024),
16489 }
16490 }
16491
16492 fn classify_call(&self, _args: &Value) -> ai_agents_core::ToolCallClassification {
16493 let mut classification =
16494 ai_agents_core::ToolCallClassification::from_metadata(&self.safety_metadata());
16495 classification.safely_retryable = false;
16496 classification
16497 }
16498
16499 async fn execute(
16500 &self,
16501 _args: Value,
16502 _ctx: ai_agents_core::ToolExecutionContext,
16503 ) -> ToolResult {
16504 let call = self.calls.fetch_add(1, Ordering::SeqCst);
16505 if call == 0 {
16506 ToolResult::error("first failure")
16507 } else {
16508 ToolResult::ok("second success")
16509 }
16510 }
16511 }
16512
16513 #[async_trait]
16514 impl ai_agents_core::Tool for LockedWriteTool {
16515 fn id(&self) -> &str {
16516 "locked_write"
16517 }
16518
16519 fn name(&self) -> &str {
16520 "Locked Write"
16521 }
16522
16523 fn description(&self) -> &str {
16524 "Tracks concurrent execution on one resource."
16525 }
16526
16527 fn input_schema(&self) -> Value {
16528 serde_json::json!({"type": "object", "properties": {"path": {"type": "string"}}})
16529 }
16530
16531 fn policy_bindings(&self) -> ai_agents_core::ToolPolicyBindings {
16532 ai_agents_core::ToolPolicyBindings {
16533 path_fields: vec![ai_agents_core::PathPolicyBinding::write("path")],
16534 ..Default::default()
16535 }
16536 }
16537
16538 fn safety_metadata(&self) -> ai_agents_core::ToolSafetyMetadata {
16539 ai_agents_core::ToolSafetyMetadata {
16540 read_only: false,
16541 concurrency_safe: false,
16542 operation: ai_agents_core::ToolOperationKind::Write,
16543 side_effect_level: ai_agents_core::ToolSideEffectLevel::LocalWrite,
16544 requires_network: false,
16545 destructive: false,
16546 open_world: false,
16547 host_dependent: false,
16548 requires_user_interaction: false,
16549 supports_cancellation: true,
16550 default_requires_approval: false,
16551 should_defer_schema: false,
16552 max_output_chars: Some(1024),
16553 max_result_size_chars: Some(1024),
16554 }
16555 }
16556
16557 async fn execute(
16558 &self,
16559 _args: Value,
16560 _ctx: ai_agents_core::ToolExecutionContext,
16561 ) -> ToolResult {
16562 let active = self.active.fetch_add(1, Ordering::SeqCst) + 1;
16563 loop {
16564 let current_max = self.max_active.load(Ordering::SeqCst);
16565 if active <= current_max {
16566 break;
16567 }
16568 if self
16569 .max_active
16570 .compare_exchange(current_max, active, Ordering::SeqCst, Ordering::SeqCst)
16571 .is_ok()
16572 {
16573 break;
16574 }
16575 }
16576 tokio::time::sleep(std::time::Duration::from_millis(100)).await;
16577 self.active.fetch_sub(1, Ordering::SeqCst);
16578 ToolResult::ok("done")
16579 }
16580 }
16581
16582 #[async_trait]
16583 impl ai_agents_core::Tool for MultiResourceWriteTool {
16584 fn id(&self) -> &str {
16585 "multi_resource_write"
16586 }
16587
16588 fn name(&self) -> &str {
16589 "Multi Resource Write"
16590 }
16591
16592 fn description(&self) -> &str {
16593 "Tracks concurrent execution across source and destination resources."
16594 }
16595
16596 fn input_schema(&self) -> Value {
16597 serde_json::json!({"type": "object"})
16598 }
16599
16600 fn policy_bindings(&self) -> ai_agents_core::ToolPolicyBindings {
16601 ai_agents_core::ToolPolicyBindings {
16602 path_fields: vec![
16603 ai_agents_core::PathPolicyBinding::read_write("source_path"),
16604 ai_agents_core::PathPolicyBinding::write("destination_path"),
16605 ],
16606 ..Default::default()
16607 }
16608 }
16609
16610 fn safety_metadata(&self) -> ai_agents_core::ToolSafetyMetadata {
16611 LockedWriteTool {
16612 active: Arc::clone(&self.active),
16613 max_active: Arc::clone(&self.max_active),
16614 }
16615 .safety_metadata()
16616 }
16617
16618 async fn execute(
16619 &self,
16620 _args: Value,
16621 _ctx: ai_agents_core::ToolExecutionContext,
16622 ) -> ToolResult {
16623 let active = self.active.fetch_add(1, Ordering::SeqCst) + 1;
16624 self.max_active.fetch_max(active, Ordering::SeqCst);
16625 tokio::time::sleep(std::time::Duration::from_millis(75)).await;
16626 self.active.fetch_sub(1, Ordering::SeqCst);
16627 ToolResult::ok("done")
16628 }
16629 }
16630
16631 #[async_trait]
16632 impl ai_agents_core::Tool for BlockingPathMutationTool {
16633 fn id(&self) -> &str {
16634 self.id
16635 }
16636
16637 fn name(&self) -> &str {
16638 self.id
16639 }
16640
16641 fn description(&self) -> &str {
16642 "Blocks a path mutation until the test releases it."
16643 }
16644
16645 fn input_schema(&self) -> Value {
16646 serde_json::json!({"type": "object"})
16647 }
16648
16649 fn policy_bindings(&self) -> ai_agents_core::ToolPolicyBindings {
16650 ai_agents_core::ToolPolicyBindings {
16651 path_fields: self.path_fields.clone(),
16652 ..Default::default()
16653 }
16654 }
16655
16656 fn safety_metadata(&self) -> ai_agents_core::ToolSafetyMetadata {
16657 ai_agents_core::ToolSafetyMetadata {
16658 read_only: false,
16659 concurrency_safe: false,
16660 operation: ai_agents_core::ToolOperationKind::Write,
16661 side_effect_level: ai_agents_core::ToolSideEffectLevel::LocalWrite,
16662 requires_network: false,
16663 destructive: false,
16664 open_world: false,
16665 host_dependent: false,
16666 requires_user_interaction: false,
16667 supports_cancellation: true,
16668 default_requires_approval: false,
16669 should_defer_schema: false,
16670 max_output_chars: Some(1024),
16671 max_result_size_chars: Some(1024),
16672 }
16673 }
16674
16675 async fn execute(
16676 &self,
16677 _args: Value,
16678 _ctx: ai_agents_core::ToolExecutionContext,
16679 ) -> ToolResult {
16680 self.gate.entered.store(true, Ordering::SeqCst);
16681 self.gate.entered_notify.notify_one();
16682 self.gate.release.notified().await;
16683 ToolResult::ok("done")
16684 }
16685 }
16686
16687 #[async_trait]
16688 impl ai_agents_core::Tool for NoBindingWriteTool {
16689 fn id(&self) -> &str {
16690 "no_binding_write"
16691 }
16692
16693 fn name(&self) -> &str {
16694 "No Binding Write"
16695 }
16696
16697 fn description(&self) -> &str {
16698 "Tracks concurrent execution without resource bindings."
16699 }
16700
16701 fn input_schema(&self) -> Value {
16702 serde_json::json!({"type": "object"})
16703 }
16704
16705 fn safety_metadata(&self) -> ai_agents_core::ToolSafetyMetadata {
16706 LockedWriteTool {
16707 active: Arc::clone(&self.active),
16708 max_active: Arc::clone(&self.max_active),
16709 }
16710 .safety_metadata()
16711 }
16712
16713 async fn execute(
16714 &self,
16715 _args: Value,
16716 _ctx: ai_agents_core::ToolExecutionContext,
16717 ) -> ToolResult {
16718 let active = self.active.fetch_add(1, Ordering::SeqCst) + 1;
16719 self.max_active.fetch_max(active, Ordering::SeqCst);
16720 tokio::time::sleep(std::time::Duration::from_millis(75)).await;
16721 self.active.fetch_sub(1, Ordering::SeqCst);
16722 ToolResult::ok("done")
16723 }
16724 }
16725
16726 #[async_trait]
16727 impl ai_agents_core::Tool for RecoveryTestTool {
16728 fn id(&self) -> &str {
16729 &self.id
16730 }
16731
16732 fn name(&self) -> &str {
16733 &self.id
16734 }
16735
16736 fn description(&self) -> &str {
16737 "Records recovery execution and returns a configured result."
16738 }
16739
16740 fn input_schema(&self) -> Value {
16741 serde_json::json!({"type": "object"})
16742 }
16743
16744 fn policy_bindings(&self) -> ai_agents_core::ToolPolicyBindings {
16745 ai_agents_core::ToolPolicyBindings {
16746 path_fields: vec![ai_agents_core::PathPolicyBinding::write("path")],
16747 ..Default::default()
16748 }
16749 }
16750
16751 fn safety_metadata(&self) -> ai_agents_core::ToolSafetyMetadata {
16753 ai_agents_core::ToolSafetyMetadata {
16754 read_only: false,
16755 concurrency_safe: false,
16756 operation: ai_agents_core::ToolOperationKind::Write,
16757 side_effect_level: ai_agents_core::ToolSideEffectLevel::LocalWrite,
16758 requires_network: false,
16759 destructive: false,
16760 open_world: false,
16761 host_dependent: false,
16762 requires_user_interaction: false,
16763 supports_cancellation: true,
16764 default_requires_approval: false,
16765 should_defer_schema: false,
16766 max_output_chars: Some(self.max_output_chars.unwrap_or(1024)),
16767 max_result_size_chars: Some(1024),
16768 }
16769 }
16770
16771 async fn execute(
16773 &self,
16774 _args: Value,
16775 _ctx: ai_agents_core::ToolExecutionContext,
16776 ) -> ToolResult {
16777 self.calls.fetch_add(1, Ordering::SeqCst);
16778 let mut result = if self.succeeds {
16779 ToolResult::ok(format!("{} succeeded", self.id))
16780 } else {
16781 ToolResult::error(format!("{} failed", self.id))
16782 };
16783 result.metadata = Some(HashMap::from([(
16784 "recovery_test_tool".to_string(),
16785 Value::String(self.id.clone()),
16786 )]));
16787 result
16788 }
16789 }
16790
16791 #[async_trait]
16792 impl WebFetchTransport for RuntimeWebFetchTransport {
16793 async fn send(
16795 &self,
16796 _request: WebFetchTransportRequest,
16797 ) -> std::result::Result<WebFetchTransportResponse, String> {
16798 Err("validated addresses are required".to_string())
16799 }
16800
16801 async fn send_validated(
16803 &self,
16804 _request: WebFetchTransportRequest,
16805 _addresses: &[std::net::SocketAddr],
16806 ) -> std::result::Result<WebFetchTransportResponse, String> {
16807 self.calls.fetch_add(1, Ordering::SeqCst);
16808 Ok(WebFetchTransportResponse {
16809 status: 200,
16810 content_type: Some("text/plain".to_string()),
16811 location: None,
16812 body: b"approved".to_vec(),
16813 })
16814 }
16815 }
16816
16817 #[async_trait]
16818 impl WebFetchResolver for RuntimeWebFetchResolver {
16819 async fn resolve(
16821 &self,
16822 _host: &str,
16823 _port: u16,
16824 ) -> std::result::Result<Vec<std::net::IpAddr>, String> {
16825 Ok(vec![std::net::IpAddr::V4(std::net::Ipv4Addr::new(
16826 93, 184, 216, 34,
16827 ))])
16828 }
16829 }
16830
16831 #[async_trait]
16832 impl ToolProvider for DriftingFallbackProvider {
16833 fn id(&self) -> &str {
16835 "drifting_fallback"
16836 }
16837
16838 fn name(&self) -> &str {
16840 "Drifting Fallback"
16841 }
16842
16843 fn provider_type(&self) -> ToolProviderType {
16845 ToolProviderType::Custom
16846 }
16847
16848 async fn list_tools(&self) -> Vec<ToolDescriptor> {
16850 let alias = ToolAliases::new().with_name("en", "fallback alias");
16851 let mut primary = ToolDescriptor::new(
16852 "primary",
16853 "Primary",
16854 "Fails before fallback.",
16855 serde_json::json!({"type": "object"}),
16856 );
16857 let mut secondary = ToolDescriptor::new(
16858 "secondary",
16859 "Secondary",
16860 "Must not execute after final canonical drift.",
16861 serde_json::json!({"type": "object"}),
16862 );
16863 if self.refreshed.load(Ordering::SeqCst) {
16864 primary = primary.with_aliases(alias);
16865 } else {
16866 secondary = secondary.with_aliases(alias);
16867 }
16868 vec![primary, secondary]
16869 }
16870
16871 async fn get_tool(&self, tool_id: &str) -> Option<Arc<dyn Tool>> {
16873 let calls = match tool_id {
16874 "primary" => Arc::clone(&self.primary_calls),
16875 "secondary" => Arc::clone(&self.secondary_calls),
16876 _ => return None,
16877 };
16878 Some(Arc::new(RecoveryTestTool {
16879 id: tool_id.to_string(),
16880 succeeds: false,
16881 calls,
16882 max_output_chars: None,
16883 }))
16884 }
16885
16886 fn supports_refresh(&self) -> bool {
16888 true
16889 }
16890
16891 async fn refresh(&self) -> std::result::Result<(), ToolProviderError> {
16893 self.refreshed.store(true, Ordering::SeqCst);
16894 Ok(())
16895 }
16896 }
16897
16898 #[async_trait]
16899 impl AgentHooks for RefreshFallbackProviderHooks {
16900 async fn on_tool_start(&self, tool: &str, args: &Value) {
16902 self.lifecycle.on_tool_start(tool, args).await;
16903 if tool != "secondary" {
16904 return;
16905 }
16906 let agent = self.agent.lock().as_ref().and_then(Weak::upgrade);
16907 if let Some(agent) = agent {
16908 agent
16909 .tools
16910 .refresh_provider("drifting_fallback")
16911 .await
16912 .unwrap();
16913 }
16914 }
16915
16916 async fn on_tool_complete(&self, tool: &str, result: &ToolResult, duration_ms: u64) {
16917 self.lifecycle
16918 .on_tool_complete(tool, result, duration_ms)
16919 .await;
16920 }
16921
16922 async fn on_tool_execution_record(&self, record: &ToolExecutionRecord) {
16923 self.lifecycle.on_tool_execution_record(record).await;
16924 }
16925
16926 async fn on_error(&self, error: &AgentError) {
16927 self.lifecycle.on_error(error).await;
16928 }
16929 }
16930
16931 #[async_trait]
16932 impl ApprovalHandler for BlockingApprovalHandler {
16933 async fn request_approval(
16934 &self,
16935 _request: ai_agents_hitl::ApprovalRequest,
16936 ) -> ApprovalResult {
16937 self.entered.wait().await;
16938 self.release.notified().await;
16939 self.result.clone()
16940 }
16941 }
16942
16943 #[async_trait]
16944 impl ApprovalHandler for CountingApprovalHandler {
16945 async fn request_approval(
16946 &self,
16947 _request: ai_agents_hitl::ApprovalRequest,
16948 ) -> ApprovalResult {
16949 self.calls.fetch_add(1, Ordering::SeqCst);
16950 ApprovalResult::Approved
16951 }
16952 }
16953
16954 #[async_trait]
16955 impl AgentHooks for ReentrantToolHooks {
16956 async fn on_tool_complete(&self, tool: &str, _result: &ToolResult, _duration_ms: u64) {
16957 if tool != "reentrant_write" || self.invoked.swap(true, Ordering::SeqCst) {
16958 return;
16959 }
16960 let agent = self.agent.lock().as_ref().and_then(Weak::upgrade);
16961 if let Some(agent) = agent {
16962 let result = agent
16963 .invoke_tool(ToolExecutionRequest::new(
16964 "nested-hook-call",
16965 "reentrant_write",
16966 serde_json::json!({"path": "./hook.txt"}),
16967 ToolCallSource::Manual,
16968 ))
16969 .await;
16970 self.nested_success
16971 .store(result.is_ok_and(|record| record.success), Ordering::SeqCst);
16972 }
16973 }
16974 }
16975
16976 #[async_trait]
16977 impl AgentHooks for ResponseCountingHooks {
16978 async fn on_response(&self, _response: &AgentResponse) {
16979 self.responses.fetch_add(1, Ordering::SeqCst);
16980 }
16981 }
16982
16983 #[async_trait]
16984 impl AgentHooks for ResponseChatHooks {
16985 async fn on_response(&self, _response: &AgentResponse) {
16987 if self.invoked.swap(true, Ordering::SeqCst) {
16988 return;
16989 }
16990 let target = self.target.lock().as_ref().and_then(Weak::upgrade);
16991 let result = if let Some(target) = target {
16992 target
16993 .chat("nested response hook call")
16994 .await
16995 .map(|response| response.content)
16996 .map_err(|error| error.to_string())
16997 } else {
16998 Err("response hook target is unavailable".to_string())
16999 };
17000 *self.nested_result.lock() = Some(result);
17001 }
17002 }
17003
17004 #[async_trait]
17005 impl AgentHooks for ConcurrentResponseHooks {
17006 async fn on_response(&self, _response: &AgentResponse) {
17008 if self.invoked.swap(true, Ordering::SeqCst) {
17009 return;
17010 }
17011 let Some(registry) = self.registry.upgrade() else {
17012 *self.nested_result.lock() =
17013 Some(Err("concurrent registry is unavailable".to_string()));
17014 return;
17015 };
17016 let agents = [ai_agents_state::ConcurrentAgentRef::Id(
17017 self.child_id.clone(),
17018 )];
17019 let aggregation = ai_agents_state::AggregationConfig {
17020 strategy: ai_agents_state::AggregationStrategy::FirstWins,
17021 synthesizer_llm: None,
17022 synthesizer_prompt: None,
17023 vote: None,
17024 };
17025 let result = crate::orchestration::concurrent(
17026 ®istry,
17027 "nested concurrent response hook call",
17028 &agents,
17029 &aggregation,
17030 None,
17031 Some(1),
17032 None,
17033 ai_agents_state::PartialFailureAction::Abort,
17034 None,
17035 )
17036 .await
17037 .map(|result| result.response.content)
17038 .map_err(|error| error.to_string());
17039 *self.nested_result.lock() = Some(result);
17040 }
17041 }
17042
17043 #[async_trait]
17044 impl AgentHooks for ToolLifecycleRecordingHooks {
17045 async fn on_tool_start(&self, tool: &str, _args: &Value) {
17046 self.events.lock().push(format!("start:{tool}"));
17047 }
17048
17049 async fn on_tool_complete(&self, tool: &str, result: &ToolResult, _duration_ms: u64) {
17050 self.events
17051 .lock()
17052 .push(format!("complete:{tool}:{}", result.success));
17053 }
17054
17055 async fn on_tool_execution_record(&self, record: &ToolExecutionRecord) {
17056 self.events.lock().push(format!(
17057 "record:{}:{}",
17058 record.canonical_id, record.executed
17059 ));
17060 self.records.lock().push(record.clone());
17061 }
17062
17063 async fn on_error(&self, _error: &AgentError) {
17065 self.events.lock().push("error".to_string());
17066 }
17067 }
17068
17069 struct ApprovalRecordingHooks {
17070 events: parking_lot::Mutex<Vec<String>>,
17071 }
17072
17073 impl ApprovalRecordingHooks {
17074 fn new() -> Self {
17075 Self {
17076 events: parking_lot::Mutex::new(Vec::new()),
17077 }
17078 }
17079
17080 fn events(&self) -> Vec<String> {
17081 self.events.lock().clone()
17082 }
17083 }
17084
17085 #[async_trait]
17086 impl AgentHooks for ApprovalRecordingHooks {
17087 async fn on_approval_result(&self, request_id: &str, result: &ApprovalResult) {
17088 self.events.lock().push(format!(
17089 "raw:{}:{}",
17090 request_id,
17091 approval_result_name(result)
17092 ));
17093 }
17094
17095 async fn on_approval_resolved(
17096 &self,
17097 request: &ai_agents_hitl::ApprovalRequest,
17098 raw_result: &ApprovalResult,
17099 outcome: &ApprovalResolvedOutcome,
17100 ) {
17101 self.events.lock().push(format!(
17102 "resolved:{}:{}:{}",
17103 request.id,
17104 approval_result_name(raw_result),
17105 approval_outcome_name(outcome)
17106 ));
17107 }
17108 }
17109
17110 fn approval_result_name(result: &ApprovalResult) -> &'static str {
17111 match result {
17112 ApprovalResult::Approved => "approved",
17113 ApprovalResult::Rejected { .. } => "rejected",
17114 ApprovalResult::Modified { .. } => "modified",
17115 ApprovalResult::Timeout => "timeout",
17116 }
17117 }
17118
17119 fn approval_outcome_name(outcome: &ApprovalResolvedOutcome) -> &'static str {
17120 match outcome {
17121 ApprovalResolvedOutcome::Approved => "approved",
17122 ApprovalResolvedOutcome::Rejected { .. } => "rejected",
17123 ApprovalResolvedOutcome::Modified { .. } => "modified",
17124 ApprovalResolvedOutcome::Error { .. } => "error",
17125 }
17126 }
17127
17128 fn assert_correlated_approval_events(
17129 events: &[String],
17130 raw_status: &str,
17131 outcome_status: &str,
17132 ) {
17133 assert_eq!(events.len(), 2);
17134 let raw: Vec<_> = events[0].split(':').collect();
17135 let resolved: Vec<_> = events[1].split(':').collect();
17136 assert_eq!(raw[0], "raw");
17137 assert_eq!(resolved[0], "resolved");
17138 assert_eq!(raw[1], resolved[1]);
17139 assert_eq!(raw[2], raw_status);
17140 assert_eq!(resolved[2], raw_status);
17141 assert_eq!(resolved[3], outcome_status);
17142 }
17143
17144 fn approval_security_config(policy_enabled: bool) -> ToolSecurityConfig {
17145 let mut security = ToolSecurityConfig {
17146 enabled: true,
17147 fail_closed: true,
17148 ..Default::default()
17149 };
17150 let policy = ai_agents_tools::ToolPolicyConfig {
17151 enabled: policy_enabled,
17152 write_paths: vec![".".to_string()],
17153 require_confirmation: true,
17154 ..Default::default()
17155 };
17156 security.tools.insert("locked_write".to_string(), policy);
17157 security
17158 }
17159
17160 struct MutationTestWorkspace {
17161 root: std::path::PathBuf,
17162 }
17163
17164 impl MutationTestWorkspace {
17165 fn new() -> Self {
17166 let root = std::env::temp_dir().join(format!(
17167 "ai-agents-runtime-mutation-{}",
17168 uuid::Uuid::new_v4()
17169 ));
17170 std::fs::create_dir_all(&root).unwrap();
17171 Self { root }
17172 }
17173 }
17174
17175 impl Drop for MutationTestWorkspace {
17176 fn drop(&mut self) {
17177 let _ = std::fs::remove_dir_all(&self.root);
17178 }
17179 }
17180
17181 async fn wait_for_resource_lock_strong_count(locks: &ToolResourceLocks, minimum: usize) {
17182 tokio::time::timeout(std::time::Duration::from_secs(2), async {
17183 loop {
17184 let strong_count = locks
17185 .read()
17186 .get("path-mutation:global")
17187 .map_or(0, |lock| lock.strong_count());
17188 if strong_count >= minimum {
17189 break;
17190 }
17191 tokio::task::yield_now().await;
17192 }
17193 })
17194 .await
17195 .expect("path mutation call did not reach the shared lock");
17196 }
17197
17198 async fn assert_path_mutation_pair_serialized(
17199 first_id: &'static str,
17200 first_fields: Vec<ai_agents_core::PathPolicyBinding>,
17201 first_args: Value,
17202 second_id: &'static str,
17203 second_fields: Vec<ai_agents_core::PathPolicyBinding>,
17204 second_args: Value,
17205 ) {
17206 let locks = new_tool_resource_locks();
17207 let first_gate = PathMutationGate::new();
17208 let second_gate = PathMutationGate::new();
17209 second_gate.release();
17210 let agent = Arc::new(
17211 AgentBuilder::new()
17212 .system_prompt("Test global path mutation locking.")
17213 .llm(Arc::new(mock_with_response("done")))
17214 .tool(Arc::new(BlockingPathMutationTool {
17215 id: first_id,
17216 path_fields: first_fields,
17217 gate: first_gate.clone(),
17218 }))
17219 .tool(Arc::new(BlockingPathMutationTool {
17220 id: second_id,
17221 path_fields: second_fields,
17222 gate: second_gate.clone(),
17223 }))
17224 .build()
17225 .unwrap()
17226 .with_shared_resource_locks(Arc::clone(&locks)),
17227 );
17228
17229 let first = {
17230 let agent = Arc::clone(&agent);
17231 tokio::spawn(async move {
17232 agent
17233 .invoke_tool(ToolExecutionRequest::new(
17234 format!("{}-first", first_id),
17235 first_id,
17236 first_args,
17237 ToolCallSource::Manual,
17238 ))
17239 .await
17240 .unwrap()
17241 })
17242 };
17243 first_gate.wait_until_entered().await;
17244
17245 let second = {
17246 let agent = Arc::clone(&agent);
17247 tokio::spawn(async move {
17248 agent
17249 .invoke_tool(ToolExecutionRequest::new(
17250 format!("{}-second", second_id),
17251 second_id,
17252 second_args,
17253 ToolCallSource::Manual,
17254 ))
17255 .await
17256 .unwrap()
17257 })
17258 };
17259 wait_for_resource_lock_strong_count(&locks, 2).await;
17260 assert!(!second_gate.entered.load(Ordering::SeqCst));
17261 assert!(!second.is_finished());
17262
17263 first_gate.release();
17264 let (first, second) = tokio::time::timeout(std::time::Duration::from_secs(2), async {
17265 tokio::join!(first, second)
17266 })
17267 .await
17268 .expect("serialized path mutation calls did not finish");
17269 assert!(first.unwrap().success);
17270 assert!(second.unwrap().success);
17271 assert!(second_gate.entered.load(Ordering::SeqCst));
17272 assert!(locks.read().is_empty());
17273 }
17274
17275 #[derive(Clone, Copy)]
17276 enum MutationDenial {
17277 Policy,
17278 Approval,
17279 }
17280
17281 fn mutation_denial_security_config(
17282 tool_id: &str,
17283 workspace: &std::path::Path,
17284 denial: MutationDenial,
17285 ) -> ToolSecurityConfig {
17286 let workspace = workspace.to_string_lossy().into_owned();
17287 let mut policy = ai_agents_tools::ToolPolicyConfig {
17288 read_paths: vec![workspace.clone()],
17289 write_paths: vec![workspace.clone()],
17290 ..Default::default()
17291 };
17292 match denial {
17293 MutationDenial::Policy => policy.blocked_paths = vec![workspace],
17294 MutationDenial::Approval => policy.require_confirmation = true,
17295 }
17296
17297 let mut security = ToolSecurityConfig {
17298 enabled: true,
17299 fail_closed: true,
17300 ..Default::default()
17301 };
17302 security.tools.insert(tool_id.to_string(), policy);
17303 security
17304 }
17305
17306 async fn assert_path_mutation_denied(tool: Arc<dyn Tool>, denial: MutationDenial) {
17307 let workspace = MutationTestWorkspace::new();
17308 let tool_id = tool.id().to_string();
17309 let preserved = workspace.root.join(format!("{}-preserved.txt", tool_id));
17310 let destination = workspace.root.join(format!("{}-destination.txt", tool_id));
17311 std::fs::write(&preserved, "preserved").unwrap();
17312 let arguments = match tool_id.as_str() {
17313 "copy_path" | "move_path" => serde_json::json!({
17314 "source_path": preserved.to_string_lossy(),
17315 "destination_path": destination.to_string_lossy(),
17316 "dry_run": false
17317 }),
17318 "delete_path" => serde_json::json!({
17319 "path": preserved.to_string_lossy(),
17320 "recursive": false,
17321 "dry_run": false
17322 }),
17323 _ => panic!("unsupported mutation tool: {}", tool_id),
17324 };
17325 let security = mutation_denial_security_config(&tool_id, &workspace.root, denial);
17326 let builder = AgentBuilder::new()
17327 .system_prompt("Test mutation denial.")
17328 .llm(Arc::new(mock_with_response("done")))
17329 .tool(tool)
17330 .tool_security(ToolSecurityEngine::new(security));
17331 let builder = match denial {
17332 MutationDenial::Policy => builder,
17333 MutationDenial::Approval => builder
17334 .hitl_engine(HITLEngine::new(ai_agents_hitl::HITLConfig::default()))
17335 .approval_handler(Arc::new(RejectAllHandler::new())),
17336 };
17337 let agent = builder.build().unwrap();
17338
17339 let record = agent
17340 .invoke_tool(ToolExecutionRequest::new(
17341 format!("{}-denied", tool_id),
17342 tool_id.clone(),
17343 arguments,
17344 ToolCallSource::Manual,
17345 ))
17346 .await
17347 .unwrap();
17348
17349 assert!(!record.executed, "{} must not be invoked", tool_id);
17350 assert!(!record.success);
17351 match denial {
17352 MutationDenial::Policy => {
17353 assert_eq!(record.policy.outcome, PermissionOutcome::Deny);
17354 assert!(record.approval.as_ref().is_some_and(|approval| matches!(
17355 &approval.status,
17356 ToolApprovalStatus::NotRequired
17357 )));
17358 }
17359 MutationDenial::Approval => {
17360 assert_eq!(record.policy.outcome, PermissionOutcome::RequiresApproval);
17361 assert!(record.approval.as_ref().is_some_and(|approval| matches!(
17362 &approval.status,
17363 ToolApprovalStatus::Rejected
17364 )));
17365 }
17366 }
17367 assert_eq!(std::fs::read_to_string(&preserved).unwrap(), "preserved");
17368 assert!(!destination.exists());
17369 }
17370
17371 fn recovery_manager_with_fallbacks(
17372 fallbacks: impl IntoIterator<Item = (String, String)>,
17373 ) -> RecoveryManager {
17374 use ai_agents_recovery::{ErrorRecoveryConfig, ToolRecoveryConfig, ToolRetryConfig};
17375
17376 let per_tool = fallbacks
17377 .into_iter()
17378 .map(|(tool, fallback_tool)| {
17379 (
17380 tool,
17381 ToolRetryConfig {
17382 max_retries: 0,
17383 timeout_ms: Some(1_000),
17384 on_failure: ToolFailureAction::Fallback { fallback_tool },
17385 },
17386 )
17387 })
17388 .collect();
17389 RecoveryManager::new(ErrorRecoveryConfig {
17390 tools: ToolRecoveryConfig {
17391 per_tool,
17392 ..Default::default()
17393 },
17394 ..Default::default()
17395 })
17396 }
17397
17398 fn approval_check() -> HITLCheckResult {
17399 HITLCheckResult::required(
17400 ApprovalTrigger::tool("test", serde_json::json!({})),
17401 HashMap::new(),
17402 "Approve?",
17403 None,
17404 )
17405 }
17406
17407 fn agent_with_approval_result(
17408 raw_result: ApprovalResult,
17409 timeout_action: TimeoutAction,
17410 hooks: Arc<ApprovalRecordingHooks>,
17411 ) -> RuntimeAgent {
17412 use ai_agents_hitl::{CallbackHandler, HITLConfig};
17413
17414 let config = HITLConfig {
17415 on_timeout: timeout_action,
17416 ..Default::default()
17417 };
17418 let handler = CallbackHandler::new(move |_| raw_result.clone());
17419 AgentBuilder::new()
17420 .system_prompt("Test HITL hooks.")
17421 .llm(Arc::new(mock_with_response("done")))
17422 .build()
17423 .unwrap()
17424 .with_hooks(hooks)
17425 .with_hitl(HITLEngine::new(config), Arc::new(handler))
17426 }
17427
17428 #[tokio::test]
17429 async fn approval_hooks_expose_direct_effective_decisions_after_raw_results() {
17430 let cases = vec![
17431 (ApprovalResult::Approved, "approved"),
17432 (
17433 ApprovalResult::Rejected {
17434 reason: Some("denied".to_string()),
17435 },
17436 "rejected",
17437 ),
17438 (
17439 ApprovalResult::Modified {
17440 changes: HashMap::from([("value".to_string(), serde_json::json!(2))]),
17441 },
17442 "modified",
17443 ),
17444 ];
17445
17446 for (raw_result, expected) in cases {
17447 let hooks = Arc::new(ApprovalRecordingHooks::new());
17448 let agent =
17449 agent_with_approval_result(raw_result, TimeoutAction::Reject, hooks.clone());
17450
17451 let result = agent.request_hitl_approval(approval_check()).await.unwrap();
17452
17453 assert_eq!(approval_result_name(&result), expected);
17454 assert_correlated_approval_events(&hooks.events(), expected, expected);
17455 }
17456 }
17457
17458 #[tokio::test]
17459 async fn approval_hooks_expose_timeout_policy_decisions() {
17460 for (timeout_action, expected) in [
17461 (TimeoutAction::Approve, "approved"),
17462 (TimeoutAction::Reject, "rejected"),
17463 ] {
17464 let hooks = Arc::new(ApprovalRecordingHooks::new());
17465 let agent =
17466 agent_with_approval_result(ApprovalResult::Timeout, timeout_action, hooks.clone());
17467
17468 let result = agent.request_hitl_approval(approval_check()).await.unwrap();
17469
17470 assert_eq!(approval_result_name(&result), expected);
17471 assert_correlated_approval_events(&hooks.events(), "timeout", expected);
17472 }
17473 }
17474
17475 #[tokio::test]
17476 async fn timeout_error_fires_correlated_resolved_error_before_returning() {
17477 let hooks = Arc::new(ApprovalRecordingHooks::new());
17478 let agent = agent_with_approval_result(
17479 ApprovalResult::Timeout,
17480 TimeoutAction::Error,
17481 hooks.clone(),
17482 );
17483
17484 let error = agent
17485 .request_hitl_approval(approval_check())
17486 .await
17487 .unwrap_err();
17488
17489 assert!(error.to_string().contains("HITL approval timeout"));
17490 assert_correlated_approval_events(&hooks.events(), "timeout", "error");
17491 }
17492
17493 #[tokio::test]
17495 async fn test_integration_yaml_to_chat_basic() {
17496 let mock = mock_with_response("Hello! How can I help you?");
17497 let agent = AgentBuilder::new()
17498 .system_prompt("You are a test assistant.")
17499 .llm(Arc::new(mock))
17500 .build()
17501 .unwrap();
17502
17503 let response = agent.chat("Hi").await.unwrap();
17504 assert!(!response.content.is_empty());
17505 assert_eq!(response.content, "Hello! How can I help you?");
17506 }
17507
17508 #[tokio::test]
17509 async fn stream_events_emit_one_authoritative_final_without_legacy_done() {
17510 let agent = AgentBuilder::new()
17511 .system_prompt("You are a test assistant.")
17512 .llm(Arc::new(mock_with_response(
17513 "Hello from the final response.",
17514 )))
17515 .build()
17516 .unwrap();
17517
17518 let mut stream = agent.chat_stream_events("Hi").await.unwrap();
17519 let mut final_responses = Vec::new();
17520 let mut legacy_done = 0;
17521 while let Some(event) = stream.next().await {
17522 match event {
17523 AgentStreamEvent::Chunk(StreamChunk::Done {}) => legacy_done += 1,
17524 AgentStreamEvent::Chunk(StreamChunk::Error { message }) => {
17525 panic!("unexpected stream error: {message}")
17526 }
17527 AgentStreamEvent::Final(response) => final_responses.push(response),
17528 AgentStreamEvent::Chunk(_) => {}
17529 }
17530 }
17531
17532 assert_eq!(legacy_done, 0);
17533 assert_eq!(final_responses.len(), 1);
17534 let response = final_responses.pop().unwrap();
17535 assert_eq!(response.content, "Hello from the final response.");
17536 assert!(
17537 response
17538 .metadata
17539 .as_ref()
17540 .is_some_and(|metadata| { metadata.contains_key("reasoning") })
17541 );
17542 }
17543
17544 #[tokio::test]
17545 async fn stream_final_content_includes_output_processing_after_provisional_chunks() {
17546 let yaml = r#"
17547name: ProcessedStreamAgent
17548system_prompt: "Answer directly."
17549process:
17550 output:
17551 - type: format
17552 config:
17553 template: "{{ response }} [finalized]"
17554streaming:
17555 enabled: true
17556"#;
17557 let agent = AgentBuilder::from_yaml(yaml)
17558 .unwrap()
17559 .llm(Arc::new(mock_with_response("provisional answer")))
17560 .auto_configure_features()
17561 .unwrap()
17562 .build()
17563 .unwrap();
17564
17565 let mut stream = agent.chat_stream_events("Hi").await.unwrap();
17566 let mut provisional = String::new();
17567 let mut final_content = None;
17568 while let Some(event) = stream.next().await {
17569 match event {
17570 AgentStreamEvent::Chunk(StreamChunk::Content { text }) => {
17571 provisional.push_str(&text)
17572 }
17573 AgentStreamEvent::Chunk(StreamChunk::Error { message }) => {
17574 panic!("unexpected stream error: {message}")
17575 }
17576 AgentStreamEvent::Final(response) => final_content = Some(response.content),
17577 AgentStreamEvent::Chunk(_) => {}
17578 }
17579 }
17580
17581 assert_eq!(provisional, "provisional answer");
17582 assert_eq!(
17583 final_content.as_deref(),
17584 Some("provisional answer [finalized]")
17585 );
17586 }
17587
17588 #[tokio::test]
17589 async fn stream_events_preserve_tool_progress_and_final_tool_calls() {
17590 let agent = AgentBuilder::new()
17591 .system_prompt("Use the echo tool once, then answer.")
17592 .llm(Arc::new(mock_with_responses(vec![
17593 r#"{"tool":"echo","arguments":{"message":"hello"}}"#,
17594 "Echo completed.",
17595 ])))
17596 .tool(Arc::new(ai_agents_tools::EchoTool::new()))
17597 .build()
17598 .unwrap();
17599
17600 let mut stream = agent.chat_stream_events("echo hello").await.unwrap();
17601 let mut starts = 0;
17602 let mut results = 0;
17603 let mut ends = 0;
17604 let mut final_response = None;
17605 while let Some(event) = stream.next().await {
17606 match event {
17607 AgentStreamEvent::Chunk(StreamChunk::ToolCallStart { name, .. }) => {
17608 assert_eq!(name, "echo");
17609 starts += 1;
17610 }
17611 AgentStreamEvent::Chunk(StreamChunk::ToolResult { name, success, .. }) => {
17612 assert_eq!(name, "echo");
17613 assert!(success);
17614 results += 1;
17615 }
17616 AgentStreamEvent::Chunk(StreamChunk::ToolCallEnd { .. }) => ends += 1,
17617 AgentStreamEvent::Chunk(StreamChunk::Error { message }) => {
17618 panic!("unexpected stream error: {message}")
17619 }
17620 AgentStreamEvent::Final(response) => final_response = Some(response),
17621 AgentStreamEvent::Chunk(_) => {}
17622 }
17623 }
17624
17625 assert_eq!((starts, results, ends), (1, 1, 1));
17626 let response = final_response.expect("tool stream must finalize");
17627 assert_eq!(response.content, "Echo completed.");
17628 assert_eq!(
17629 response.tool_calls.as_ref().map(|calls| calls
17630 .iter()
17631 .map(|call| call.name.as_str())
17632 .collect::<Vec<_>>()),
17633 Some(vec!["echo"])
17634 );
17635 }
17636
17637 #[tokio::test]
17638 async fn legacy_stream_still_emits_one_done_chunk() {
17639 let agent = AgentBuilder::new()
17640 .system_prompt("You are a test assistant.")
17641 .llm(Arc::new(mock_with_response(
17642 "Hello from the legacy stream.",
17643 )))
17644 .build()
17645 .unwrap();
17646
17647 let mut stream = agent.chat_stream("Hi").await.unwrap();
17648 let mut done = 0;
17649 while let Some(chunk) = stream.next().await {
17650 match chunk {
17651 StreamChunk::Done {} => done += 1,
17652 StreamChunk::Error { message } => panic!("unexpected stream error: {message}"),
17653 _ => {}
17654 }
17655 }
17656
17657 assert_eq!(done, 1);
17658 }
17659
17660 #[tokio::test]
17662 async fn test_integration_multi_turn_conversation() {
17663 let mock = mock_with_responses(vec![
17664 "Hello! I'm your assistant.",
17665 "The weather is sunny today.",
17666 "Goodbye!",
17667 ]);
17668 let agent = AgentBuilder::new()
17669 .system_prompt("You are helpful.")
17670 .llm(Arc::new(mock))
17671 .build()
17672 .unwrap();
17673
17674 let r1 = agent.chat("Hi").await.unwrap();
17675 assert_eq!(r1.content, "Hello! I'm your assistant.");
17676
17677 let r2 = agent.chat("What's the weather?").await.unwrap();
17678 assert_eq!(r2.content, "The weather is sunny today.");
17679
17680 let r3 = agent.chat("Bye").await.unwrap();
17681 assert_eq!(r3.content, "Goodbye!");
17682
17683 let messages = agent.memory.get_messages(None).await.unwrap();
17685 assert_eq!(messages.len(), 6);
17687 }
17688
17689 #[test]
17690 fn later_approval_preserves_modified_evidence() {
17691 let arguments = serde_json::json!({"dry_run": true});
17692 let mut record = Some(ToolApprovalRecord {
17693 status: ToolApprovalStatus::Modified,
17694 reason: None,
17695 modified_arguments: Some(arguments.clone()),
17696 });
17697
17698 merge_approved_record(&mut record);
17699
17700 let record = record.unwrap();
17701 assert!(matches!(record.status, ToolApprovalStatus::Modified));
17702 assert_eq!(record.modified_arguments, Some(arguments));
17703 }
17704
17705 #[test]
17706 fn approval_binding_rejects_replaced_tool_implementation() {
17707 let reviewed_tool: Arc<dyn ai_agents_core::Tool> = Arc::new(ContextEchoTool);
17708 let same_tool = Arc::clone(&reviewed_tool);
17709 let replacement_tool: Arc<dyn ai_agents_core::Tool> = Arc::new(ContextEchoTool);
17710 let arguments = serde_json::json!({"path": "."});
17711 let versions = ToolDecisionVersions {
17712 policy: 2,
17713 registry: 3,
17714 runtime_control: 4,
17715 state: Some(5),
17716 };
17717 let binding = ToolApprovalBinding {
17718 canonical_id: "context_echo".to_string(),
17719 arguments: arguments.clone(),
17720 confirmation_required: true,
17721 policy_version: versions.policy,
17722 runtime_control_version: versions.runtime_control,
17723 state_generation: versions.state,
17724 reviewed_tool,
17725 };
17726
17727 assert!(!binding.is_stale("context_echo", &arguments, true, versions, &same_tool,));
17728 assert!(binding.is_stale(
17729 "context_echo",
17730 &arguments,
17731 true,
17732 versions,
17733 &replacement_tool,
17734 ));
17735 }
17736
17737 #[tokio::test]
17738 async fn approved_mutation_to_dry_run_remains_executable() {
17739 use ai_agents_hitl::CallbackHandler;
17740
17741 let handler = CallbackHandler::new(|_| ApprovalResult::Modified {
17742 changes: HashMap::from([("dry_run".to_string(), serde_json::json!(true))]),
17743 });
17744 let agent = AgentBuilder::new()
17745 .system_prompt("Test safer approval modifications.")
17746 .llm(Arc::new(mock_with_response("done")))
17747 .tool(Arc::new(ai_agents_tools::FileWriteTool::new()))
17748 .hitl_engine(HITLEngine::new(ai_agents_hitl::HITLConfig::default()))
17749 .approval_handler(Arc::new(handler))
17750 .build()
17751 .unwrap();
17752
17753 let record = agent
17754 .invoke_tool(ToolExecutionRequest::new(
17755 "approved-dry-run",
17756 "file_write",
17757 serde_json::json!({
17758 "path": "./approval-dry-run.txt",
17759 "content": "not written"
17760 }),
17761 ToolCallSource::Manual,
17762 ))
17763 .await
17764 .unwrap();
17765
17766 assert!(record.executed);
17767 assert!(record.success);
17768 assert_eq!(record.executed_arguments["dry_run"], true);
17769 assert!(matches!(
17770 record.approval.as_ref().map(|approval| &approval.status),
17771 Some(ToolApprovalStatus::Modified)
17772 ));
17773 let output: Value = serde_json::from_str(&record.output).unwrap();
17774 assert_eq!(output["mutation_performed"], false);
17775 }
17776
17777 #[tokio::test]
17779 async fn shared_executor_approval_reaches_web_fetch_transport() {
17780 use ai_agents_hitl::{CallbackHandler, HITLConfig};
17781 use ai_agents_tools::{DomainPolicyConfig, ToolPolicyConfig};
17782
17783 let calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
17784 let tool = WebFetchTool::with_transport_and_resolver(
17785 Arc::new(RuntimeWebFetchTransport {
17786 calls: Arc::clone(&calls),
17787 }),
17788 Arc::new(RuntimeWebFetchResolver),
17789 );
17790 let mut security = ToolSecurityConfig {
17791 enabled: true,
17792 fail_closed: true,
17793 ..Default::default()
17794 };
17795 security.tools.insert(
17796 "web_fetch".to_string(),
17797 ToolPolicyConfig {
17798 domains: DomainPolicyConfig {
17799 requires_approval: vec!["approval.test".to_string()],
17800 ..Default::default()
17801 },
17802 allowed_schemes: vec!["https".to_string()],
17803 allowed_ports: vec![443],
17804 ..Default::default()
17805 },
17806 );
17807 let handler = CallbackHandler::new(|_| ApprovalResult::Approved);
17808 let agent = AgentBuilder::new()
17809 .system_prompt("Test approved web fetch execution.")
17810 .llm(Arc::new(mock_with_response("done")))
17811 .tool(Arc::new(tool))
17812 .tool_security(ToolSecurityEngine::new(security))
17813 .build()
17814 .unwrap()
17815 .with_hitl(HITLEngine::new(HITLConfig::default()), Arc::new(handler));
17816
17817 let record = agent
17818 .invoke_tool(ToolExecutionRequest::new(
17819 "approved-web-fetch",
17820 "web_fetch",
17821 serde_json::json!({
17822 "url": "https://approval.test/page",
17823 "cache_ttl_seconds": 0
17824 }),
17825 ToolCallSource::Manual,
17826 ))
17827 .await
17828 .unwrap();
17829
17830 assert!(record.success);
17831 assert!(
17832 record
17833 .approval
17834 .as_ref()
17835 .is_some_and(|approval| matches!(approval.status, ToolApprovalStatus::Approved))
17836 );
17837 assert_eq!(calls.load(Ordering::SeqCst), 1);
17838 }
17839
17840 #[tokio::test]
17841 async fn context_preserves_requested_and_canonical_identity() {
17842 let mock = mock_with_response("hello");
17843 let mut tools = ai_agents_tools::ToolRegistry::new();
17844 tools.register(Arc::new(ContextEchoTool)).unwrap();
17845
17846 let mut security = ToolSecurityConfig {
17847 enabled: true,
17848 fail_closed: true,
17849 ..Default::default()
17850 };
17851 let mut policy = ai_agents_tools::ToolPolicyConfig {
17852 read_paths: vec![".".to_string()],
17853 max_results: Some(7),
17854 ..Default::default()
17855 };
17856 policy
17857 .config
17858 .insert("backend".to_string(), serde_json::json!("memory"));
17859 security.tools.insert("context_echo".to_string(), policy);
17860
17861 let agent = AgentBuilder::new()
17862 .system_prompt("You are helpful.")
17863 .llm(Arc::new(mock))
17864 .tools(tools)
17865 .tool_security(ToolSecurityEngine::new(security))
17866 .build()
17867 .unwrap();
17868
17869 let record = agent
17870 .invoke_tool(ToolExecutionRequest::new(
17871 "ctx-call",
17872 "Context Echo",
17873 serde_json::json!({"path": ".", "max_results": 99}),
17874 ToolCallSource::Manual,
17875 ))
17876 .await
17877 .unwrap();
17878
17879 assert!(record.success);
17880 assert!(matches!(&record.source, ToolCallSource::Manual));
17881 assert_eq!(record.requested_name, "Context Echo");
17882 assert_eq!(record.canonical_id, "context_echo");
17883 assert_eq!(record.policy.outcome, PermissionOutcome::Allow);
17884 assert_eq!(record.executed_arguments["max_results"], 7);
17885 let output: Value = serde_json::from_str(&record.output).unwrap();
17886 assert_eq!(output["requested_name"], "Context Echo");
17887 assert_eq!(output["canonical_id"], "context_echo");
17888 assert_eq!(output["max_results"], 7);
17889 assert_eq!(output["custom_config"]["backend"], "memory");
17890 assert!(record.metadata.contains_key("effective_limits"));
17891 assert!(record.metadata.contains_key("policy_snapshot"));
17892 }
17893
17894 #[tokio::test]
17895 async fn test_runtime_control_cancels_active_tool_call() {
17896 let mock = mock_with_response("hello");
17897 let agent = Arc::new(
17898 AgentBuilder::new()
17899 .system_prompt("You are helpful.")
17900 .llm(Arc::new(mock))
17901 .tool(Arc::new(SlowTool))
17902 .build()
17903 .unwrap(),
17904 );
17905 let control = agent.runtime_control();
17906 let running_agent = Arc::clone(&agent);
17907 let handle = tokio::spawn(async move {
17908 running_agent
17909 .invoke_tool(ToolExecutionRequest::new(
17910 "slow-call",
17911 "slow",
17912 serde_json::json!({}),
17913 ToolCallSource::Manual,
17914 ))
17915 .await
17916 .unwrap()
17917 });
17918
17919 tokio::time::sleep(std::time::Duration::from_millis(100)).await;
17920 control.cancel_all();
17921 let record = handle.await.unwrap();
17922
17923 assert!(record.executed);
17924 assert!(record.cancelled);
17925 assert!(!record.success);
17926 assert!(record.cancellation_reason.is_some());
17927 }
17928
17929 #[tokio::test]
17931 async fn cancelled_tool_does_not_enter_fallback() {
17932 let fallback_calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
17933 let agent = Arc::new(
17934 AgentBuilder::new()
17935 .system_prompt("Test cancellation before fallback.")
17936 .llm(Arc::new(mock_with_response("done")))
17937 .tool(Arc::new(SlowTool))
17938 .tool(Arc::new(RecoveryTestTool {
17939 id: "fallback".to_string(),
17940 succeeds: true,
17941 calls: Arc::clone(&fallback_calls),
17942 max_output_chars: None,
17943 }))
17944 .recovery_manager(recovery_manager_with_fallbacks([(
17945 "slow".to_string(),
17946 "fallback".to_string(),
17947 )]))
17948 .build()
17949 .unwrap(),
17950 );
17951 let control = agent.runtime_control();
17952 let running_agent = Arc::clone(&agent);
17953 let handle = tokio::spawn(async move {
17954 running_agent
17955 .invoke_tool(ToolExecutionRequest::new(
17956 "cancelled-fallback-call",
17957 "slow",
17958 serde_json::json!({}),
17959 ToolCallSource::Manual,
17960 ))
17961 .await
17962 .unwrap()
17963 });
17964
17965 tokio::time::sleep(Duration::from_millis(100)).await;
17966 control.cancel_all();
17967 let record = handle.await.unwrap();
17968
17969 assert!(record.executed);
17970 assert!(record.cancelled);
17971 assert!(!record.success);
17972 assert_eq!(record.canonical_id, "slow");
17973 assert_eq!(fallback_calls.load(Ordering::SeqCst), 0);
17974 assert_eq!(agent.tool_call_history().len(), 1);
17975 }
17976
17977 #[tokio::test]
17978 async fn non_idempotent_tool_calls_are_not_retried() {
17979 use ai_agents_recovery::{ErrorRecoveryConfig, ToolRecoveryConfig, ToolRetryConfig};
17980
17981 let mock = mock_with_response("hello");
17982 let calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
17983 let agent = AgentBuilder::new()
17984 .system_prompt("You are helpful.")
17985 .llm(Arc::new(mock))
17986 .tool(Arc::new(FlakyWriteTool {
17987 calls: Arc::clone(&calls),
17988 }))
17989 .recovery_manager(RecoveryManager::new(ErrorRecoveryConfig {
17990 tools: ToolRecoveryConfig {
17991 default: ToolRetryConfig {
17992 max_retries: 2,
17993 ..Default::default()
17994 },
17995 ..Default::default()
17996 },
17997 ..Default::default()
17998 }))
17999 .build()
18000 .unwrap();
18001
18002 let record = agent
18003 .invoke_tool(ToolExecutionRequest::new(
18004 "flaky-call",
18005 "flaky_write",
18006 serde_json::json!({"path": "./tmp.txt"}),
18007 ToolCallSource::Manual,
18008 ))
18009 .await
18010 .unwrap();
18011
18012 assert!(!record.success);
18013 assert_eq!(calls.load(Ordering::SeqCst), 1);
18014 }
18015
18016 #[tokio::test]
18017 async fn safely_retryable_tool_receives_a_fresh_deadline_per_attempt() {
18018 use ai_agents_recovery::{ErrorRecoveryConfig, ToolRecoveryConfig, ToolRetryConfig};
18019
18020 let calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
18021 let deadlines = Arc::new(parking_lot::Mutex::new(Vec::new()));
18022 let remaining_ms = Arc::new(parking_lot::Mutex::new(Vec::new()));
18023 let agent = AgentBuilder::new()
18024 .system_prompt("Test retry deadlines.")
18025 .llm(Arc::new(mock_with_response("done")))
18026 .tool(Arc::new(RetryDeadlineTool {
18027 calls: Arc::clone(&calls),
18028 deadlines: Arc::clone(&deadlines),
18029 remaining_ms: Arc::clone(&remaining_ms),
18030 }))
18031 .recovery_manager(RecoveryManager::new(ErrorRecoveryConfig {
18032 tools: ToolRecoveryConfig {
18033 per_tool: HashMap::from([(
18034 "retry_deadline".to_string(),
18035 ToolRetryConfig {
18036 max_retries: 1,
18037 ..Default::default()
18038 },
18039 )]),
18040 ..Default::default()
18041 },
18042 ..Default::default()
18043 }))
18044 .build()
18045 .unwrap();
18046
18047 let record = agent
18048 .invoke_tool(ToolExecutionRequest::new(
18049 "retry-deadline-call",
18050 "retry_deadline",
18051 serde_json::json!({}),
18052 ToolCallSource::Manual,
18053 ))
18054 .await
18055 .unwrap();
18056
18057 assert!(record.executed);
18058 assert!(record.success);
18059 assert_eq!(calls.load(Ordering::SeqCst), 2);
18060 let deadlines = deadlines.lock();
18061 assert_eq!(deadlines.len(), 2);
18062 assert!(
18063 deadlines[1] > deadlines[0],
18064 "retry inherited the first invocation deadline"
18065 );
18066 let remaining_ms = remaining_ms.lock();
18067 assert_eq!(remaining_ms.len(), 2);
18068 assert!(
18069 remaining_ms
18070 .iter()
18071 .all(|remaining| (800..=1_000).contains(remaining))
18072 );
18073 }
18074
18075 #[tokio::test]
18077 async fn call_classification_timeout_controls_deadline_and_timer() {
18078 let calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
18079 let remaining_ms = Arc::new(parking_lot::Mutex::new(Vec::new()));
18080 let agent = AgentBuilder::new()
18081 .system_prompt("Test call-level timeout.")
18082 .llm(Arc::new(mock_with_response("done")))
18083 .tool(Arc::new(ClassifiedTimeoutTool {
18084 id: "classified_timeout",
18085 calls: Arc::clone(&calls),
18086 timeout_ms: 100,
18087 sleep_ms: 150,
18088 requires_approval: false,
18089 remaining_ms: Arc::clone(&remaining_ms),
18090 }))
18091 .build()
18092 .unwrap();
18093
18094 let started = Instant::now();
18095 let record = agent
18096 .invoke_tool(ToolExecutionRequest::new(
18097 "classified-timeout-call",
18098 "classified_timeout",
18099 serde_json::json!({}),
18100 ToolCallSource::Manual,
18101 ))
18102 .await
18103 .unwrap();
18104
18105 assert!(record.executed);
18106 assert!(record.timed_out);
18107 assert!(!record.success);
18108 assert_eq!(calls.load(Ordering::SeqCst), 1);
18109 assert!(started.elapsed() < Duration::from_secs(1));
18110 let remaining_ms = remaining_ms.lock();
18111 assert_eq!(remaining_ms.len(), 1);
18112 assert!((1..=100).contains(&remaining_ms[0]));
18113 }
18114
18115 #[tokio::test]
18117 async fn recovery_timeout_only_lowers_call_and_policy_timeouts() {
18118 use ai_agents_recovery::{ErrorRecoveryConfig, ToolRecoveryConfig, ToolRetryConfig};
18119
18120 let calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
18121 let remaining_ms = Arc::new(parking_lot::Mutex::new(Vec::new()));
18122 let agent = AgentBuilder::new()
18123 .system_prompt("Test recovery timeout.")
18124 .llm(Arc::new(mock_with_response("done")))
18125 .tool(Arc::new(ClassifiedTimeoutTool {
18126 id: "recovery_timeout",
18127 calls: Arc::clone(&calls),
18128 timeout_ms: 1_000,
18129 sleep_ms: 150,
18130 requires_approval: false,
18131 remaining_ms: Arc::clone(&remaining_ms),
18132 }))
18133 .recovery_manager(RecoveryManager::new(ErrorRecoveryConfig {
18134 tools: ToolRecoveryConfig {
18135 per_tool: HashMap::from([(
18136 "recovery_timeout".to_string(),
18137 ToolRetryConfig {
18138 timeout_ms: Some(100),
18139 ..Default::default()
18140 },
18141 )]),
18142 ..Default::default()
18143 },
18144 ..Default::default()
18145 }))
18146 .build()
18147 .unwrap();
18148
18149 let started = Instant::now();
18150 let record = agent
18151 .invoke_tool(ToolExecutionRequest::new(
18152 "recovery-timeout-call",
18153 "recovery_timeout",
18154 serde_json::json!({}),
18155 ToolCallSource::Manual,
18156 ))
18157 .await
18158 .unwrap();
18159
18160 assert!(record.executed);
18161 assert!(record.timed_out);
18162 assert!(!record.success);
18163 assert_eq!(calls.load(Ordering::SeqCst), 1);
18164 assert!(started.elapsed() < Duration::from_secs(1));
18165 assert_eq!(record.metadata["effective_limits"]["timeout_ms"], 100);
18166 let remaining_ms = remaining_ms.lock();
18167 assert_eq!(remaining_ms.len(), 1);
18168 assert!((1..=100).contains(&remaining_ms[0]));
18169 }
18170
18171 #[tokio::test]
18173 async fn recovery_default_timeout_controls_deadline_and_timer() {
18174 use ai_agents_recovery::{ErrorRecoveryConfig, ToolRecoveryConfig, ToolRetryConfig};
18175
18176 let calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
18177 let remaining_ms = Arc::new(parking_lot::Mutex::new(Vec::new()));
18178 let agent = AgentBuilder::new()
18179 .system_prompt("Test default recovery timeout.")
18180 .llm(Arc::new(mock_with_response("done")))
18181 .tool(Arc::new(ClassifiedTimeoutTool {
18182 id: "default_recovery_timeout",
18183 calls: Arc::clone(&calls),
18184 timeout_ms: 1_000,
18185 sleep_ms: 150,
18186 requires_approval: false,
18187 remaining_ms: Arc::clone(&remaining_ms),
18188 }))
18189 .recovery_manager(RecoveryManager::new(ErrorRecoveryConfig {
18190 tools: ToolRecoveryConfig {
18191 default: ToolRetryConfig {
18192 timeout_ms: Some(100),
18193 ..Default::default()
18194 },
18195 ..Default::default()
18196 },
18197 ..Default::default()
18198 }))
18199 .build()
18200 .unwrap();
18201
18202 let started = Instant::now();
18203 let record = agent
18204 .invoke_tool(ToolExecutionRequest::new(
18205 "default-recovery-timeout-call",
18206 "default_recovery_timeout",
18207 serde_json::json!({}),
18208 ToolCallSource::Manual,
18209 ))
18210 .await
18211 .unwrap();
18212
18213 assert!(record.executed);
18214 assert!(record.timed_out);
18215 assert!(!record.success);
18216 assert_eq!(calls.load(Ordering::SeqCst), 1);
18217 assert!(started.elapsed() < Duration::from_secs(1));
18218 assert_eq!(record.metadata["effective_limits"]["timeout_ms"], 100);
18219 let remaining_ms = remaining_ms.lock();
18220 assert_eq!(remaining_ms.len(), 1);
18221 assert!((1..=100).contains(&remaining_ms[0]));
18222 }
18223
18224 #[test]
18226 fn recovery_timeout_cannot_widen_security_baseline() {
18227 let security_engine = ToolSecurityEngine::new(ToolSecurityConfig {
18228 default_timeout_ms: 100,
18229 ..Default::default()
18230 });
18231 let safety = ToolSafetyMetadata::compute();
18232 let mut classification = ToolCallClassification::from_metadata(&safety);
18233 classification.timeout_ms = Some(500);
18234
18235 let (limits, timeout) = RuntimeAgent::effective_tool_limits(
18236 &security_engine,
18237 "recovery_cannot_widen",
18238 &safety,
18239 &classification,
18240 Some(1_000),
18241 )
18242 .unwrap();
18243
18244 assert_eq!(limits.timeout_ms, Some(100));
18245 assert_eq!(timeout.timer, Duration::from_millis(100));
18246 }
18247
18248 #[tokio::test]
18250 async fn invalid_call_timeout_stops_before_approval_or_tool_invocation() {
18251 let tool_calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
18252 let approval_calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
18253 let remaining_ms = Arc::new(parking_lot::Mutex::new(Vec::new()));
18254 let mut security = ToolSecurityConfig {
18255 enabled: true,
18256 ..Default::default()
18257 };
18258 security.tools.insert(
18259 "invalid_call_timeout".to_string(),
18260 ai_agents_tools::ToolPolicyConfig {
18261 require_confirmation: true,
18262 ..Default::default()
18263 },
18264 );
18265 let agent = AgentBuilder::new()
18266 .system_prompt("Test invalid call timeout.")
18267 .llm(Arc::new(mock_with_response("done")))
18268 .tool(Arc::new(ClassifiedTimeoutTool {
18269 id: "invalid_call_timeout",
18270 calls: Arc::clone(&tool_calls),
18271 timeout_ms: u64::MAX,
18272 sleep_ms: 0,
18273 requires_approval: false,
18274 remaining_ms,
18275 }))
18276 .tool_security(ToolSecurityEngine::new(security))
18277 .hitl_engine(HITLEngine::new(ai_agents_hitl::HITLConfig::default()))
18278 .approval_handler(Arc::new(CountingApprovalHandler {
18279 calls: Arc::clone(&approval_calls),
18280 }))
18281 .build()
18282 .unwrap();
18283
18284 let error = agent
18285 .invoke_tool(ToolExecutionRequest::new(
18286 "invalid-call-timeout",
18287 "invalid_call_timeout",
18288 serde_json::json!({}),
18289 ToolCallSource::Manual,
18290 ))
18291 .await
18292 .unwrap_err();
18293
18294 assert!(error.to_string().contains(
18295 "effective tool timeout_ms must be no greater than 3153600000000000 milliseconds"
18296 ));
18297 assert_eq!(approval_calls.load(Ordering::SeqCst), 0);
18298 assert_eq!(tool_calls.load(Ordering::SeqCst), 0);
18299 }
18300
18301 #[tokio::test]
18303 async fn invalid_modified_call_timeout_stops_before_lock_or_invocation() {
18304 use ai_agents_hitl::CallbackHandler;
18305
18306 let blocker_gate = PathMutationGate::new();
18307 let tool_calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
18308 let hooks = Arc::new(ToolLifecycleRecordingHooks::new());
18309 let handler = CallbackHandler::new(|_| ApprovalResult::Modified {
18310 changes: HashMap::from([("invalid_timeout".to_string(), Value::Bool(true))]),
18311 });
18312 let agent = Arc::new(
18313 AgentBuilder::new()
18314 .system_prompt("Test final call timeout validation.")
18315 .llm(Arc::new(mock_with_response("done")))
18316 .tool(Arc::new(BlockingPathMutationTool {
18317 id: "timeout_lock_blocker",
18318 path_fields: vec![ai_agents_core::PathPolicyBinding::write("path")],
18319 gate: blocker_gate.clone(),
18320 }))
18321 .tool(Arc::new(ApprovalModifiedTimeoutTool {
18322 calls: Arc::clone(&tool_calls),
18323 }))
18324 .hitl_engine(HITLEngine::new(ai_agents_hitl::HITLConfig::default()))
18325 .approval_handler(Arc::new(handler))
18326 .hooks(hooks.clone())
18327 .build()
18328 .unwrap(),
18329 );
18330 let blocking_agent = Arc::clone(&agent);
18331 let blocker = tokio::spawn(async move {
18332 blocking_agent
18333 .invoke_tool(ToolExecutionRequest::new(
18334 "timeout-lock-blocker",
18335 "timeout_lock_blocker",
18336 serde_json::json!({"path": "./shared-timeout.txt"}),
18337 ToolCallSource::Manual,
18338 ))
18339 .await
18340 .unwrap()
18341 });
18342 blocker_gate.wait_until_entered().await;
18343
18344 let record = tokio::time::timeout(
18345 Duration::from_millis(500),
18346 agent.invoke_tool(ToolExecutionRequest::new(
18347 "invalid-modified-timeout",
18348 "approval_modified_timeout",
18349 serde_json::json!({
18350 "path": "./shared-timeout.txt",
18351 "invalid_timeout": false
18352 }),
18353 ToolCallSource::Manual,
18354 )),
18355 )
18356 .await
18357 .expect("final timeout validation must not wait for the held path lock")
18358 .unwrap();
18359
18360 blocker_gate.release();
18361 assert!(blocker.await.unwrap().success);
18362 assert!(!record.executed);
18363 assert!(!record.success);
18364 assert_eq!(record.policy.outcome, PermissionOutcome::Deny);
18365 assert!(record.output.contains(
18366 "effective tool timeout_ms must be no greater than 3153600000000000 milliseconds"
18367 ));
18368 assert_eq!(tool_calls.load(Ordering::SeqCst), 0);
18369 let invalid_request_events = hooks
18370 .events()
18371 .into_iter()
18372 .filter(|event| event.contains("approval_modified_timeout") || event == "error")
18373 .collect::<Vec<_>>();
18374 assert_eq!(
18375 invalid_request_events,
18376 vec![
18377 "start:approval_modified_timeout",
18378 "complete:approval_modified_timeout:false",
18379 "record:approval_modified_timeout:false",
18380 "error"
18381 ]
18382 );
18383 }
18384
18385 #[tokio::test]
18386 async fn side_effecting_tools_are_serialized_per_resource() {
18387 let mock = mock_with_response("hello");
18388 let active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
18389 let max_active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
18390 let agent = Arc::new(
18391 AgentBuilder::new()
18392 .system_prompt("You are helpful.")
18393 .llm(Arc::new(mock))
18394 .tool(Arc::new(LockedWriteTool {
18395 active: Arc::clone(&active),
18396 max_active: Arc::clone(&max_active),
18397 }))
18398 .build()
18399 .unwrap(),
18400 );
18401
18402 let left = {
18403 let agent = Arc::clone(&agent);
18404 tokio::spawn(async move {
18405 agent
18406 .invoke_tool(ToolExecutionRequest::new(
18407 "lock-1",
18408 "locked_write",
18409 serde_json::json!({"path": "./same.txt"}),
18410 ToolCallSource::Manual,
18411 ))
18412 .await
18413 .unwrap()
18414 })
18415 };
18416 let right = {
18417 let agent = Arc::clone(&agent);
18418 tokio::spawn(async move {
18419 agent
18420 .invoke_tool(ToolExecutionRequest::new(
18421 "lock-2",
18422 "locked_write",
18423 serde_json::json!({"path": "./same.txt"}),
18424 ToolCallSource::Manual,
18425 ))
18426 .await
18427 .unwrap()
18428 })
18429 };
18430
18431 let left = left.await.unwrap();
18432 let right = right.await.unwrap();
18433 assert!(left.success);
18434 assert!(right.success);
18435 assert_eq!(max_active.load(Ordering::SeqCst), 1);
18436 }
18437
18438 #[tokio::test]
18439 async fn path_resources_use_shared_global_lock_and_cleanup() {
18440 let active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
18441 let max_active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
18442 let bindings = ai_agents_core::ToolPolicyBindings {
18443 path_fields: vec![
18444 ai_agents_core::PathPolicyBinding::read_write("source_path"),
18445 ai_agents_core::PathPolicyBinding::write("destination_path"),
18446 ],
18447 ..Default::default()
18448 };
18449 let classification = ai_agents_core::ToolCallClassification::from_metadata(
18450 &MultiResourceWriteTool {
18451 active: Arc::clone(&active),
18452 max_active: Arc::clone(&max_active),
18453 }
18454 .safety_metadata(),
18455 );
18456 let left_args = serde_json::json!({
18457 "source_path": "./a/../first.txt",
18458 "destination_path": "./second.txt"
18459 });
18460 let right_args = serde_json::json!({
18461 "source_path": "./second.txt",
18462 "destination_path": "./first.txt"
18463 });
18464 let left_keys = tool_resource_lock_keys(
18465 "multi_resource_write",
18466 &left_args,
18467 &bindings,
18468 &classification,
18469 );
18470 let right_keys = tool_resource_lock_keys(
18471 "multi_resource_write",
18472 &right_args,
18473 &bindings,
18474 &classification,
18475 );
18476 assert_eq!(left_keys, right_keys);
18477 assert_eq!(left_keys, vec!["path-mutation:global".to_string()]);
18478
18479 let locks = new_tool_resource_locks();
18480 let build_agent = || {
18481 AgentBuilder::new()
18482 .system_prompt("Test shared resource locks.")
18483 .llm(Arc::new(mock_with_response("done")))
18484 .tool(Arc::new(MultiResourceWriteTool {
18485 active: Arc::clone(&active),
18486 max_active: Arc::clone(&max_active),
18487 }))
18488 .build()
18489 .unwrap()
18490 .with_shared_resource_locks(Arc::clone(&locks))
18491 };
18492 let left_agent = Arc::new(build_agent());
18493 let right_agent = Arc::new(build_agent());
18494 let left = tokio::spawn(async move {
18495 left_agent
18496 .invoke_tool(ToolExecutionRequest::new(
18497 "multi-left",
18498 "multi_resource_write",
18499 left_args,
18500 ToolCallSource::Manual,
18501 ))
18502 .await
18503 .unwrap()
18504 });
18505 let right = tokio::spawn(async move {
18506 right_agent
18507 .invoke_tool(ToolExecutionRequest::new(
18508 "multi-right",
18509 "multi_resource_write",
18510 right_args,
18511 ToolCallSource::Manual,
18512 ))
18513 .await
18514 .unwrap()
18515 });
18516 let (left, right) = tokio::time::timeout(std::time::Duration::from_secs(2), async {
18517 tokio::join!(left, right)
18518 })
18519 .await
18520 .expect("reversed resource acquisition must not deadlock");
18521
18522 assert!(left.unwrap().success);
18523 assert!(right.unwrap().success);
18524 assert_eq!(max_active.load(Ordering::SeqCst), 1);
18525 assert!(locks.read().is_empty());
18526 }
18527
18528 #[tokio::test]
18529 async fn global_path_lock_serializes_copy_destination_with_file_write() {
18530 assert_path_mutation_pair_serialized(
18531 "copy_path",
18532 CopyPathTool::new().policy_bindings().path_fields,
18533 serde_json::json!({
18534 "source_path": "./source.txt",
18535 "destination_path": "./shared.txt"
18536 }),
18537 "file_write",
18538 FileWriteTool::new().policy_bindings().path_fields,
18539 serde_json::json!({"path": "./shared.txt"}),
18540 )
18541 .await;
18542 }
18543
18544 #[tokio::test]
18545 async fn parent_and_spawned_runtime_share_global_path_lock() {
18546 let workspace = MutationTestWorkspace::new();
18547 let destination = workspace.root.join("spawned.txt");
18548 let parent_gate = PathMutationGate::new();
18549 let parent = Arc::new(
18550 AgentBuilder::from_yaml(
18551 r#"
18552name: LockParent
18553system_prompt: parent
18554llm:
18555 default: default
18556tools:
18557 - parent_path_write
18558spawner:
18559 shared_llms: true
18560"#,
18561 )
18562 .unwrap()
18563 .llm(Arc::new(mock_with_response("done")))
18564 .auto_configure_spawner()
18565 .await
18566 .unwrap()
18567 .tool(Arc::new(BlockingPathMutationTool {
18568 id: "parent_path_write",
18569 path_fields: vec![ai_agents_core::PathPolicyBinding::write("path")],
18570 gate: parent_gate.clone(),
18571 }))
18572 .build()
18573 .unwrap(),
18574 );
18575
18576 let mut child_spec = crate::spec::AgentSpec {
18577 name: "LockChild".to_string(),
18578 system_prompt: "child".to_string(),
18579 tools: Some(vec![crate::spec::ToolEntry::Simple(
18580 "file_write".to_string(),
18581 )]),
18582 ..Default::default()
18583 };
18584 child_spec.tool_security.enabled = true;
18585 child_spec.tool_security.fail_closed = true;
18586 let file_write_policy = ai_agents_tools::ToolPolicyConfig {
18587 write_paths: vec![workspace.root.to_string_lossy().into_owned()],
18588 allow_without_confirmation: true,
18589 ..Default::default()
18590 };
18591 child_spec
18592 .tool_security
18593 .tools
18594 .insert("file_write".to_string(), file_write_policy);
18595 let spawned = parent
18596 .spawner()
18597 .unwrap()
18598 .spawn_from_spec(child_spec)
18599 .await
18600 .unwrap();
18601 assert!(Arc::ptr_eq(
18602 &parent.resource_locks,
18603 &spawned.agent.resource_locks
18604 ));
18605 assert!(!Arc::ptr_eq(
18606 &parent.runtime_control,
18607 &spawned.agent.runtime_control
18608 ));
18609
18610 let parent_call = {
18611 let parent = Arc::clone(&parent);
18612 let destination = destination.clone();
18613 tokio::spawn(async move {
18614 parent
18615 .invoke_tool(ToolExecutionRequest::new(
18616 "parent-lock-holder",
18617 "parent_path_write",
18618 serde_json::json!({"path": destination}),
18619 ToolCallSource::Manual,
18620 ))
18621 .await
18622 .unwrap()
18623 })
18624 };
18625 parent_gate.wait_until_entered().await;
18626
18627 let child_call = {
18628 let child = Arc::clone(&spawned.agent);
18629 let destination = destination.clone();
18630 tokio::spawn(async move {
18631 child
18632 .invoke_tool(ToolExecutionRequest::new(
18633 "spawned-file-write",
18634 "file_write",
18635 serde_json::json!({
18636 "path": destination,
18637 "content": "spawned",
18638 "dry_run": false
18639 }),
18640 ToolCallSource::Manual,
18641 ))
18642 .await
18643 .unwrap()
18644 })
18645 };
18646 wait_for_resource_lock_strong_count(&parent.resource_locks, 2).await;
18647 assert!(!child_call.is_finished());
18648
18649 parent_gate.release();
18650 let (parent_record, child_record) =
18651 tokio::time::timeout(std::time::Duration::from_secs(2), async {
18652 tokio::join!(parent_call, child_call)
18653 })
18654 .await
18655 .expect("parent and spawned path mutations did not finish");
18656 assert!(parent_record.unwrap().success);
18657 assert!(child_record.unwrap().success);
18658 assert_eq!(std::fs::read_to_string(destination).unwrap(), "spawned");
18659 assert!(parent.resource_locks.read().is_empty());
18660 }
18661
18662 #[tokio::test]
18663 async fn cancelled_global_path_lock_waiter_does_not_retain_weak_entry() {
18664 let locks = new_tool_resource_locks();
18665 let holder_gate = PathMutationGate::new();
18666 let waiter_gate = PathMutationGate::new();
18667 waiter_gate.release();
18668 let holder = Arc::new(
18669 AgentBuilder::new()
18670 .system_prompt("Hold the global path lock.")
18671 .llm(Arc::new(mock_with_response("done")))
18672 .tool(Arc::new(BlockingPathMutationTool {
18673 id: "holder_write",
18674 path_fields: vec![ai_agents_core::PathPolicyBinding::write("path")],
18675 gate: holder_gate.clone(),
18676 }))
18677 .build()
18678 .unwrap()
18679 .with_shared_resource_locks(Arc::clone(&locks)),
18680 );
18681 let waiter = Arc::new(
18682 AgentBuilder::new()
18683 .system_prompt("Wait for the global path lock.")
18684 .llm(Arc::new(mock_with_response("done")))
18685 .tool(Arc::new(BlockingPathMutationTool {
18686 id: "waiter_write",
18687 path_fields: vec![ai_agents_core::PathPolicyBinding::write("path")],
18688 gate: waiter_gate.clone(),
18689 }))
18690 .build()
18691 .unwrap()
18692 .with_shared_resource_locks(Arc::clone(&locks)),
18693 );
18694
18695 let holder_call = {
18696 let holder = Arc::clone(&holder);
18697 tokio::spawn(async move {
18698 holder
18699 .invoke_tool(ToolExecutionRequest::new(
18700 "holder-call",
18701 "holder_write",
18702 serde_json::json!({"path": "./shared.txt"}),
18703 ToolCallSource::Manual,
18704 ))
18705 .await
18706 .unwrap()
18707 })
18708 };
18709 holder_gate.wait_until_entered().await;
18710
18711 let waiter_call = {
18712 let waiter = Arc::clone(&waiter);
18713 tokio::spawn(async move {
18714 waiter
18715 .invoke_tool(ToolExecutionRequest::new(
18716 "waiter-call",
18717 "waiter_write",
18718 serde_json::json!({"path": "./shared.txt"}),
18719 ToolCallSource::Manual,
18720 ))
18721 .await
18722 .unwrap()
18723 })
18724 };
18725 wait_for_resource_lock_strong_count(&locks, 2).await;
18726 waiter.runtime_control().cancel_all();
18727
18728 let waiter_record = tokio::time::timeout(std::time::Duration::from_secs(2), waiter_call)
18729 .await
18730 .expect("cancelled lock waiter did not finish")
18731 .unwrap();
18732 assert!(!waiter_record.success);
18733 assert!(!waiter_record.executed);
18734 assert!(waiter_record.cancelled);
18735 assert_eq!(
18736 waiter_record.cancellation_reason.as_deref(),
18737 Some("runtime control cancellation")
18738 );
18739 assert!(!waiter_gate.entered.load(Ordering::SeqCst));
18740 assert_eq!(
18741 locks
18742 .read()
18743 .get("path-mutation:global")
18744 .map_or(0, |lock| lock.strong_count()),
18745 1
18746 );
18747
18748 holder_gate.release();
18749 let holder_record = tokio::time::timeout(std::time::Duration::from_secs(2), holder_call)
18750 .await
18751 .expect("lock holder did not finish")
18752 .unwrap();
18753 assert!(holder_record.success);
18754 assert!(locks.read().is_empty());
18755 }
18756
18757 #[tokio::test]
18758 async fn path_mutation_policy_and_approval_denials_do_not_invoke_tools() {
18759 for denial in [MutationDenial::Policy, MutationDenial::Approval] {
18760 let tools: [Arc<dyn Tool>; 3] = [
18761 Arc::new(CopyPathTool::new()),
18762 Arc::new(MovePathTool::new()),
18763 Arc::new(DeletePathTool::new()),
18764 ];
18765 for tool in tools {
18766 assert_path_mutation_denied(tool, denial).await;
18767 }
18768 }
18769 }
18770
18771 #[tokio::test]
18772 async fn policy_denial_keeps_executor_hook_lifecycle_and_record_authority() {
18773 let workspace = MutationTestWorkspace::new();
18774 let target = workspace.root.join("denied.txt");
18775 let hooks = Arc::new(ToolLifecycleRecordingHooks::new());
18776 let agent = AgentBuilder::new()
18777 .system_prompt("Test denied tool hooks.")
18778 .llm(Arc::new(mock_with_response("done")))
18779 .tool(Arc::new(FileWriteTool::new()))
18780 .tool_security(ToolSecurityEngine::new(mutation_denial_security_config(
18781 "file_write",
18782 &workspace.root,
18783 MutationDenial::Policy,
18784 )))
18785 .hooks(hooks.clone())
18786 .build()
18787 .unwrap();
18788
18789 let record = agent
18790 .invoke_tool(ToolExecutionRequest::new(
18791 "denied-hook-call",
18792 "file_write",
18793 serde_json::json!({
18794 "path": target.to_string_lossy(),
18795 "content": "blocked"
18796 }),
18797 ToolCallSource::Manual,
18798 ))
18799 .await
18800 .unwrap();
18801
18802 assert!(!record.executed);
18803 assert!(!record.success);
18804 assert_eq!(record.policy.outcome, PermissionOutcome::Deny);
18805 assert_eq!(
18806 hooks.events(),
18807 vec![
18808 "start:file_write",
18809 "complete:file_write:false",
18810 "record:file_write:false",
18811 "error"
18812 ]
18813 );
18814 assert!(!target.exists());
18815 }
18816
18817 #[tokio::test]
18818 async fn approval_argument_changes_are_rechecked_against_final_scope() {
18819 let active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
18820 let max_active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
18821 let entered = Arc::new(tokio::sync::Barrier::new(2));
18822 let release = Arc::new(tokio::sync::Notify::new());
18823 let handler = Arc::new(BlockingApprovalHandler {
18824 entered: Arc::clone(&entered),
18825 release: Arc::clone(&release),
18826 result: ApprovalResult::Modified {
18827 changes: HashMap::from([(
18828 "path".to_string(),
18829 Value::String("./after-approval.txt".to_string()),
18830 )]),
18831 },
18832 });
18833 let agent = Arc::new(
18834 AgentBuilder::new()
18835 .system_prompt("Test final scope validation.")
18836 .llm(Arc::new(mock_with_response("done")))
18837 .tool(Arc::new(LockedWriteTool {
18838 active: Arc::clone(&active),
18839 max_active: Arc::clone(&max_active),
18840 }))
18841 .tool_security(ToolSecurityEngine::new(approval_security_config(true)))
18842 .hitl_engine(HITLEngine::new(ai_agents_hitl::HITLConfig::default()))
18843 .approval_handler(handler)
18844 .build()
18845 .unwrap(),
18846 );
18847 let control = agent.runtime_control();
18848 let running = Arc::clone(&agent);
18849 let call = tokio::spawn(async move {
18850 running
18851 .invoke_tool(ToolExecutionRequest::new(
18852 "approval-scope",
18853 "locked_write",
18854 serde_json::json!({"path": "./before-approval.txt"}),
18855 ToolCallSource::Manual,
18856 ))
18857 .await
18858 .unwrap()
18859 });
18860 entered.wait().await;
18861 let expected_version = control.set_tool_scope(Vec::new());
18862 release.notify_one();
18863 let record = call.await.unwrap();
18864
18865 assert!(!record.executed);
18866 assert!(!record.success);
18867 assert_eq!(record.runtime_config_version, expected_version);
18868 assert_eq!(record.executed_arguments["path"], "./after-approval.txt");
18869 assert_eq!(max_active.load(Ordering::SeqCst), 0);
18870 assert_eq!(
18871 record.metadata["runtime_scope_snapshot"],
18872 serde_json::json!([])
18873 );
18874 }
18875
18876 #[tokio::test]
18877 async fn approval_is_rechecked_against_final_policy_snapshot() {
18878 let active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
18879 let max_active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
18880 let entered = Arc::new(tokio::sync::Barrier::new(2));
18881 let release = Arc::new(tokio::sync::Notify::new());
18882 let handler = Arc::new(BlockingApprovalHandler {
18883 entered: Arc::clone(&entered),
18884 release: Arc::clone(&release),
18885 result: ApprovalResult::Approved,
18886 });
18887 let agent = Arc::new(
18888 AgentBuilder::new()
18889 .system_prompt("Test final policy validation.")
18890 .llm(Arc::new(mock_with_response("done")))
18891 .tool(Arc::new(LockedWriteTool {
18892 active: Arc::clone(&active),
18893 max_active: Arc::clone(&max_active),
18894 }))
18895 .tool_security(ToolSecurityEngine::new(approval_security_config(true)))
18896 .hitl_engine(HITLEngine::new(ai_agents_hitl::HITLConfig::default()))
18897 .approval_handler(handler)
18898 .build()
18899 .unwrap(),
18900 );
18901 let control = agent.runtime_control();
18902 let running = Arc::clone(&agent);
18903 let call = tokio::spawn(async move {
18904 running
18905 .invoke_tool(ToolExecutionRequest::new(
18906 "approval-policy",
18907 "locked_write",
18908 serde_json::json!({"path": "./policy.txt"}),
18909 ToolCallSource::Manual,
18910 ))
18911 .await
18912 .unwrap()
18913 });
18914 entered.wait().await;
18915 let expected_version = control.set_tool_security(approval_security_config(false));
18916 release.notify_one();
18917 let record = call.await.unwrap();
18918
18919 assert!(!record.executed);
18920 assert!(!record.success);
18921 assert_eq!(record.runtime_config_version, expected_version);
18922 assert_eq!(record.policy.outcome, PermissionOutcome::Unavailable);
18923 assert_eq!(max_active.load(Ordering::SeqCst), 0);
18924 assert!(record.metadata.contains_key("policy_snapshot"));
18925 }
18926
18927 #[test]
18928 fn invalid_live_policy_does_not_replace_snapshot_or_generation() {
18929 let agent = AgentBuilder::new()
18930 .system_prompt("Test runtime policy validation.")
18931 .llm(Arc::new(mock_with_response("done")))
18932 .build()
18933 .unwrap();
18934 let control = agent.runtime_control();
18935 let mut valid = ToolSecurityConfig::default();
18936 valid.tools.insert(
18937 "web_search".to_string(),
18938 ai_agents_tools::ToolPolicyConfig {
18939 max_results: Some(5),
18940 ..Default::default()
18941 },
18942 );
18943 let generation = control.try_set_tool_security(valid).unwrap();
18944
18945 let mut invalid = ToolSecurityConfig::default();
18946 invalid.tools.insert(
18947 "web_search".to_string(),
18948 ai_agents_tools::ToolPolicyConfig {
18949 max_results: Some(0),
18950 ..Default::default()
18951 },
18952 );
18953 let error = control.try_set_tool_security(invalid).unwrap_err();
18954
18955 assert!(
18956 error
18957 .to_string()
18958 .contains("max_results must be greater than 0")
18959 );
18960 assert_eq!(control.version(), generation);
18961 assert_eq!(
18962 control
18963 .state
18964 .tool_security_override
18965 .read()
18966 .as_ref()
18967 .unwrap()
18968 .config()
18969 .tools["web_search"]
18970 .max_results,
18971 Some(5)
18972 );
18973 }
18974
18975 #[test]
18977 fn invalid_timeout_config_stops_before_approval_or_tool_invocation() {
18978 let tool_calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
18979 let approval_calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
18980 let spec = crate::spec::AgentSpec {
18981 tool_security: ToolSecurityConfig {
18982 enabled: true,
18983 default_timeout_ms: u64::MAX,
18984 ..Default::default()
18985 },
18986 ..Default::default()
18987 };
18988
18989 let result = AgentBuilder::from_spec(spec)
18990 .llm(Arc::new(mock_with_response("done")))
18991 .tool(Arc::new(FlakyWriteTool {
18992 calls: Arc::clone(&tool_calls),
18993 }))
18994 .hitl_engine(HITLEngine::new(ai_agents_hitl::HITLConfig::default()))
18995 .approval_handler(Arc::new(CountingApprovalHandler {
18996 calls: Arc::clone(&approval_calls),
18997 }))
18998 .build();
18999
19000 assert!(result.is_err());
19001 assert_eq!(approval_calls.load(Ordering::SeqCst), 0);
19002 assert_eq!(tool_calls.load(Ordering::SeqCst), 0);
19003 }
19004
19005 #[test]
19007 fn invalid_recovery_timeout_config_stops_before_approval_or_tool_invocation() {
19008 use ai_agents_recovery::{ErrorRecoveryConfig, ToolRecoveryConfig, ToolRetryConfig};
19009
19010 let tool_calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19011 let approval_calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19012 let spec = crate::spec::AgentSpec {
19013 error_recovery: ErrorRecoveryConfig {
19014 tools: ToolRecoveryConfig {
19015 default: ToolRetryConfig {
19016 timeout_ms: Some(u64::MAX),
19017 ..Default::default()
19018 },
19019 ..Default::default()
19020 },
19021 ..Default::default()
19022 },
19023 ..Default::default()
19024 };
19025
19026 let result = AgentBuilder::from_spec(spec)
19027 .llm(Arc::new(mock_with_response("done")))
19028 .tool(Arc::new(FlakyWriteTool {
19029 calls: Arc::clone(&tool_calls),
19030 }))
19031 .hitl_engine(HITLEngine::new(ai_agents_hitl::HITLConfig::default()))
19032 .approval_handler(Arc::new(CountingApprovalHandler {
19033 calls: Arc::clone(&approval_calls),
19034 }))
19035 .build();
19036
19037 assert!(result.is_err());
19038 assert_eq!(approval_calls.load(Ordering::SeqCst), 0);
19039 assert_eq!(tool_calls.load(Ordering::SeqCst), 0);
19040 }
19041
19042 #[test]
19044 fn invalid_timeout_policy_does_not_replace_snapshot_or_generation() {
19045 let agent = AgentBuilder::new()
19046 .system_prompt("Test runtime timeout policy validation.")
19047 .llm(Arc::new(mock_with_response("done")))
19048 .build()
19049 .unwrap();
19050 let control = agent.runtime_control();
19051 let valid = ToolSecurityConfig {
19052 default_timeout_ms: 5_000,
19053 ..Default::default()
19054 };
19055 let generation = control.try_set_tool_security(valid).unwrap();
19056
19057 let invalid = ToolSecurityConfig {
19058 default_timeout_ms: MAX_TOOL_TIMEOUT_MS + 1,
19059 ..Default::default()
19060 };
19061 let error = control.try_set_tool_security(invalid).unwrap_err();
19062
19063 assert!(error.to_string().contains(&format!(
19064 "tool_security.default_timeout_ms must be no greater than {MAX_TOOL_TIMEOUT_MS} milliseconds"
19065 )));
19066 assert_eq!(control.version(), generation);
19067 assert_eq!(
19068 control
19069 .state
19070 .tool_security_override
19071 .read()
19072 .as_ref()
19073 .unwrap()
19074 .config()
19075 .default_timeout_ms,
19076 5_000
19077 );
19078 }
19079
19080 #[test]
19082 fn runtime_tool_timeout_conversion_enforces_the_stable_boundary() {
19083 let timeout = RuntimeAgent::validated_tool_timeout(MAX_TOOL_TIMEOUT_MS).unwrap();
19084 assert_eq!(timeout.timer, Duration::from_millis(MAX_TOOL_TIMEOUT_MS));
19085 assert_eq!(
19086 timeout.deadline_delta,
19087 chrono::Duration::milliseconds(MAX_TOOL_TIMEOUT_MS as i64)
19088 );
19089
19090 for timeout_ms in [MAX_TOOL_TIMEOUT_MS + 1, u64::MAX] {
19091 let error = RuntimeAgent::validated_tool_timeout(timeout_ms).unwrap_err();
19092 assert!(error.to_string().contains(&format!(
19093 "effective tool timeout_ms must be no greater than {MAX_TOOL_TIMEOUT_MS} milliseconds"
19094 )));
19095 }
19096 }
19097
19098 #[tokio::test]
19099 async fn persistent_override_preserves_rate_history_within_generation() {
19100 let calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19101 let agent = AgentBuilder::new()
19102 .system_prompt("Test persistent policy overrides.")
19103 .llm(Arc::new(mock_with_response("done")))
19104 .tool(Arc::new(RecoveryTestTool {
19105 id: "limited_override".to_string(),
19106 succeeds: true,
19107 calls: Arc::clone(&calls),
19108 max_output_chars: None,
19109 }))
19110 .build()
19111 .unwrap();
19112 let mut security = ToolSecurityConfig {
19113 enabled: true,
19114 fail_closed: true,
19115 ..Default::default()
19116 };
19117 let policy = ai_agents_tools::ToolPolicyConfig {
19118 write_paths: vec![".".to_string()],
19119 rate_limit: Some(1),
19120 ..Default::default()
19121 };
19122 security
19123 .tools
19124 .insert("limited_override".to_string(), policy);
19125 let generation = agent.runtime_control().set_tool_security(security);
19126
19127 let first = agent
19128 .invoke_tool(ToolExecutionRequest::new(
19129 "limited-first",
19130 "limited_override",
19131 serde_json::json!({"path": "./limited.txt"}),
19132 ToolCallSource::Manual,
19133 ))
19134 .await
19135 .unwrap();
19136 let second = agent
19137 .invoke_tool(ToolExecutionRequest::new(
19138 "limited-second",
19139 "limited_override",
19140 serde_json::json!({"path": "./limited.txt"}),
19141 ToolCallSource::Manual,
19142 ))
19143 .await
19144 .unwrap();
19145
19146 assert!(first.success);
19147 assert_eq!(first.policy_version, generation);
19148 assert!(!second.executed);
19149 assert!(second.output.contains("Rate limit exceeded"));
19150 assert_eq!(second.policy_version, generation);
19151 assert_eq!(calls.load(Ordering::SeqCst), 1);
19152 }
19153
19154 #[tokio::test]
19155 async fn concurrent_rate_admission_consumes_capacity_atomically() {
19156 let calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19157 let tool = Arc::new(RecoveryTestTool {
19158 id: "atomic_rate".to_string(),
19159 succeeds: true,
19160 calls: Arc::clone(&calls),
19161 max_output_chars: None,
19162 });
19163 let arguments = serde_json::json!({"path": "./atomic-rate.txt"});
19164 let bindings = tool.policy_bindings();
19165 let classification = tool.classify_call(&arguments);
19166 let resource_keys =
19167 tool_resource_lock_keys(tool.id(), &arguments, &bindings, &classification);
19168 let mut security = ToolSecurityConfig {
19169 enabled: true,
19170 fail_closed: true,
19171 ..Default::default()
19172 };
19173 let policy = ai_agents_tools::ToolPolicyConfig {
19174 write_paths: vec![".".to_string()],
19175 rate_limit: Some(1),
19176 ..Default::default()
19177 };
19178 security.tools.insert(tool.id().to_string(), policy);
19179 let agent = Arc::new(
19180 AgentBuilder::new()
19181 .system_prompt("Test atomic rate admission.")
19182 .llm(Arc::new(mock_with_response("done")))
19183 .tool(tool)
19184 .tool_security(ToolSecurityEngine::new(security))
19185 .build()
19186 .unwrap(),
19187 );
19188 let held = agent
19189 .acquire_tool_resource_locks(&resource_keys)
19190 .await
19191 .unwrap();
19192 let left = {
19193 let agent = Arc::clone(&agent);
19194 let arguments = arguments.clone();
19195 tokio::spawn(async move {
19196 agent
19197 .invoke_tool(ToolExecutionRequest::new(
19198 "atomic-rate-left",
19199 "atomic_rate",
19200 arguments,
19201 ToolCallSource::Manual,
19202 ))
19203 .await
19204 .unwrap()
19205 })
19206 };
19207 let right = {
19208 let agent = Arc::clone(&agent);
19209 tokio::spawn(async move {
19210 agent
19211 .invoke_tool(ToolExecutionRequest::new(
19212 "atomic-rate-right",
19213 "atomic_rate",
19214 arguments,
19215 ToolCallSource::Manual,
19216 ))
19217 .await
19218 .unwrap()
19219 })
19220 };
19221 tokio::time::sleep(std::time::Duration::from_millis(25)).await;
19222 drop(held);
19223 let (left, right) = tokio::join!(left, right);
19224 let records = [left.unwrap(), right.unwrap()];
19225
19226 assert_eq!(records.iter().filter(|record| record.success).count(), 1);
19227 assert_eq!(records.iter().filter(|record| record.executed).count(), 1);
19228 assert!(
19229 records.iter().any(|record| {
19230 !record.executed && record.output.contains("Rate limit exceeded")
19231 })
19232 );
19233 assert_eq!(calls.load(Ordering::SeqCst), 1);
19234 }
19235
19236 #[tokio::test]
19237 async fn changed_policy_generation_invalidates_pending_approval() {
19238 let active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19239 let max_active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19240 let entered = Arc::new(tokio::sync::Barrier::new(2));
19241 let release = Arc::new(tokio::sync::Notify::new());
19242 let handler = Arc::new(BlockingApprovalHandler {
19243 entered: Arc::clone(&entered),
19244 release: Arc::clone(&release),
19245 result: ApprovalResult::Approved,
19246 });
19247 let agent = Arc::new(
19248 AgentBuilder::new()
19249 .system_prompt("Test stale approval denial.")
19250 .llm(Arc::new(mock_with_response("done")))
19251 .tool(Arc::new(LockedWriteTool {
19252 active: Arc::clone(&active),
19253 max_active: Arc::clone(&max_active),
19254 }))
19255 .tool_security(ToolSecurityEngine::new(approval_security_config(true)))
19256 .hitl_engine(HITLEngine::new(ai_agents_hitl::HITLConfig::default()))
19257 .approval_handler(handler)
19258 .build()
19259 .unwrap(),
19260 );
19261 let running = Arc::clone(&agent);
19262 let call = tokio::spawn(async move {
19263 running
19264 .invoke_tool(ToolExecutionRequest::new(
19265 "stale-approval",
19266 "locked_write",
19267 serde_json::json!({"path": "./stale.txt"}),
19268 ToolCallSource::Manual,
19269 ))
19270 .await
19271 .unwrap()
19272 });
19273 entered.wait().await;
19274 let generation = agent
19275 .runtime_control()
19276 .set_tool_security(approval_security_config(true));
19277 release.notify_one();
19278 let record = call.await.unwrap();
19279
19280 assert!(!record.executed);
19281 assert!(record.output.contains("Approval became stale"));
19282 assert_eq!(record.policy_version, generation);
19283 assert_eq!(max_active.load(Ordering::SeqCst), 0);
19284 }
19285
19286 #[tokio::test]
19287 async fn final_policy_reapplies_argument_caps_after_approval_changes() {
19288 use ai_agents_hitl::CallbackHandler;
19289
19290 let mut security = ToolSecurityConfig {
19291 enabled: true,
19292 fail_closed: true,
19293 ..Default::default()
19294 };
19295 let policy = ai_agents_tools::ToolPolicyConfig {
19296 read_paths: vec![".".to_string()],
19297 max_results: Some(5),
19298 require_confirmation: true,
19299 ..Default::default()
19300 };
19301 security.tools.insert("context_echo".to_string(), policy);
19302 let handler = CallbackHandler::new(|_| ApprovalResult::Modified {
19303 changes: HashMap::from([("max_results".to_string(), serde_json::json!(99))]),
19304 });
19305 let agent = AgentBuilder::new()
19306 .system_prompt("Test final argument caps.")
19307 .llm(Arc::new(mock_with_response("done")))
19308 .tool(Arc::new(ContextEchoTool))
19309 .tool_security(ToolSecurityEngine::new(security))
19310 .hitl_engine(HITLEngine::new(ai_agents_hitl::HITLConfig::default()))
19311 .approval_handler(Arc::new(handler))
19312 .build()
19313 .unwrap();
19314
19315 let record = agent
19316 .invoke_tool(ToolExecutionRequest::new(
19317 "final-cap",
19318 "context_echo",
19319 serde_json::json!({"path": ".", "max_results": 1}),
19320 ToolCallSource::Manual,
19321 ))
19322 .await
19323 .unwrap();
19324
19325 assert!(record.success);
19326 assert_eq!(record.executed_arguments["max_results"], 5);
19327 assert_eq!(
19328 record.approval.unwrap().modified_arguments.unwrap()["max_results"],
19329 5
19330 );
19331 }
19332
19333 #[tokio::test]
19334 async fn no_binding_writes_use_canonical_fallback_lock() {
19335 let active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19336 let max_active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19337 let agent = Arc::new(
19338 AgentBuilder::new()
19339 .system_prompt("Test fallback resource locks.")
19340 .llm(Arc::new(mock_with_response("done")))
19341 .tool(Arc::new(NoBindingWriteTool {
19342 active: Arc::clone(&active),
19343 max_active: Arc::clone(&max_active),
19344 }))
19345 .build()
19346 .unwrap(),
19347 );
19348 let left = {
19349 let agent = Arc::clone(&agent);
19350 tokio::spawn(async move {
19351 agent
19352 .invoke_tool(ToolExecutionRequest::new(
19353 "no-binding-left",
19354 "no_binding_write",
19355 serde_json::json!({}),
19356 ToolCallSource::Manual,
19357 ))
19358 .await
19359 .unwrap()
19360 })
19361 };
19362 let right = {
19363 let agent = Arc::clone(&agent);
19364 tokio::spawn(async move {
19365 agent
19366 .invoke_tool(ToolExecutionRequest::new(
19367 "no-binding-right",
19368 "no_binding_write",
19369 serde_json::json!({}),
19370 ToolCallSource::Manual,
19371 ))
19372 .await
19373 .unwrap()
19374 })
19375 };
19376 let (left, right) = tokio::join!(left, right);
19377
19378 assert!(left.unwrap().success);
19379 assert!(right.unwrap().success);
19380 assert_eq!(max_active.load(Ordering::SeqCst), 1);
19381 }
19382
19383 #[tokio::test]
19384 async fn parent_and_child_paths_share_a_resource_lock() {
19385 let active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19386 let max_active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19387 let agent = Arc::new(
19388 AgentBuilder::new()
19389 .system_prompt("Test parent child resource locks.")
19390 .llm(Arc::new(mock_with_response("done")))
19391 .tool(Arc::new(LockedWriteTool {
19392 active: Arc::clone(&active),
19393 max_active: Arc::clone(&max_active),
19394 }))
19395 .build()
19396 .unwrap(),
19397 );
19398 let parent = format!("./lock-parent-{}", uuid::Uuid::new_v4());
19399 let child = format!("{}/child.txt", parent);
19400 let left = {
19401 let agent = Arc::clone(&agent);
19402 tokio::spawn(async move {
19403 agent
19404 .invoke_tool(ToolExecutionRequest::new(
19405 "parent-lock",
19406 "locked_write",
19407 serde_json::json!({"path": parent}),
19408 ToolCallSource::Manual,
19409 ))
19410 .await
19411 .unwrap()
19412 })
19413 };
19414 let right = {
19415 let agent = Arc::clone(&agent);
19416 tokio::spawn(async move {
19417 agent
19418 .invoke_tool(ToolExecutionRequest::new(
19419 "child-lock",
19420 "locked_write",
19421 serde_json::json!({"path": child}),
19422 ToolCallSource::Manual,
19423 ))
19424 .await
19425 .unwrap()
19426 })
19427 };
19428 let (left, right) = tokio::join!(left, right);
19429
19430 assert!(left.unwrap().success);
19431 assert!(right.unwrap().success);
19432 assert_eq!(max_active.load(Ordering::SeqCst), 1);
19433 }
19434
19435 #[tokio::test]
19436 async fn tool_hooks_can_reenter_after_resource_guards_are_dropped() {
19437 let calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19438 let hooks = Arc::new(ReentrantToolHooks {
19439 agent: parking_lot::Mutex::new(None),
19440 invoked: AtomicBool::new(false),
19441 nested_success: AtomicBool::new(false),
19442 });
19443 let agent = Arc::new(
19444 AgentBuilder::new()
19445 .system_prompt("Test hook reentrancy.")
19446 .llm(Arc::new(mock_with_response("done")))
19447 .tool(Arc::new(RecoveryTestTool {
19448 id: "reentrant_write".to_string(),
19449 succeeds: true,
19450 calls: Arc::clone(&calls),
19451 max_output_chars: None,
19452 }))
19453 .hooks(hooks.clone())
19454 .build()
19455 .unwrap(),
19456 );
19457 *hooks.agent.lock() = Some(Arc::downgrade(&agent));
19458 let record = tokio::time::timeout(
19459 std::time::Duration::from_secs(2),
19460 agent.invoke_tool(ToolExecutionRequest::new(
19461 "outer-hook-call",
19462 "reentrant_write",
19463 serde_json::json!({"path": "./hook.txt"}),
19464 ToolCallSource::Manual,
19465 )),
19466 )
19467 .await
19468 .expect("tool completion hook must not retain resource guards")
19469 .unwrap();
19470
19471 assert!(record.success);
19472 assert!(hooks.nested_success.load(Ordering::SeqCst));
19473 assert_eq!(calls.load(Ordering::SeqCst), 2);
19474 }
19475
19476 #[tokio::test]
19478 async fn fallback_finalizes_original_record_before_shared_execution() {
19479 let primary_calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19480 let fallback_calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19481 let hooks = Arc::new(ToolLifecycleRecordingHooks::new());
19482 let agent = AgentBuilder::new()
19483 .system_prompt("Test fallback execution.")
19484 .llm(Arc::new(mock_with_response("done")))
19485 .tool(Arc::new(RecoveryTestTool {
19486 id: "primary".to_string(),
19487 succeeds: false,
19488 calls: Arc::clone(&primary_calls),
19489 max_output_chars: None,
19490 }))
19491 .tool(Arc::new(RecoveryTestTool {
19492 id: "fallback".to_string(),
19493 succeeds: true,
19494 calls: Arc::clone(&fallback_calls),
19495 max_output_chars: None,
19496 }))
19497 .recovery_manager(recovery_manager_with_fallbacks([(
19498 "primary".to_string(),
19499 "fallback".to_string(),
19500 )]))
19501 .hooks(hooks.clone())
19502 .build()
19503 .unwrap();
19504 let record = tokio::time::timeout(
19505 std::time::Duration::from_secs(2),
19506 agent.invoke_tool(ToolExecutionRequest::new(
19507 "fallback-call",
19508 "primary",
19509 serde_json::json!({"path": "./shared.txt"}),
19510 ToolCallSource::Manual,
19511 )),
19512 )
19513 .await
19514 .expect("fallback must not retain the primary resource guard")
19515 .unwrap();
19516
19517 assert_eq!(
19518 hooks.events(),
19519 vec![
19520 "start:primary",
19521 "complete:primary:false",
19522 "record:primary:true",
19523 "error",
19524 "start:fallback",
19525 "complete:fallback:true",
19526 "record:fallback:true",
19527 ]
19528 );
19529 let records = hooks.records();
19530 assert_eq!(records.len(), 2);
19531 let original = &records[0];
19532 assert_eq!(original.canonical_id, "primary");
19533 assert!(matches!(original.source, ToolCallSource::Manual));
19534 assert!(original.executed);
19535 assert!(!original.success);
19536
19537 let fallback = &records[1];
19538 assert_eq!(fallback.canonical_id, "fallback");
19539 assert_eq!(fallback.call_id, "fallback-call");
19540 assert!(matches!(
19541 &fallback.source,
19542 ToolCallSource::Fallback { original_tool } if original_tool == "primary"
19543 ));
19544 assert!(fallback.executed);
19545 assert!(fallback.success);
19546 assert_eq!(record.canonical_id, fallback.canonical_id);
19547 assert_eq!(record.output, fallback.output);
19548
19549 let history = agent.tool_call_history();
19550 assert_eq!(
19551 history
19552 .iter()
19553 .map(|entry| entry.tool_id.as_str())
19554 .collect::<Vec<_>>(),
19555 vec!["primary", "fallback"]
19556 );
19557 assert_eq!(history[0].result.get("success"), Some(&Value::Bool(false)));
19558 assert_eq!(primary_calls.load(Ordering::SeqCst), 1);
19559 assert_eq!(fallback_calls.load(Ordering::SeqCst), 1);
19560 }
19561
19562 #[tokio::test]
19564 async fn self_fallback_cycle_is_denied_before_reinvocation() {
19565 let calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19566 let hooks = Arc::new(ToolLifecycleRecordingHooks::new());
19567 let agent = AgentBuilder::new()
19568 .system_prompt("Test self-fallback cycle admission.")
19569 .llm(Arc::new(mock_with_response("done")))
19570 .tool(Arc::new(RecoveryTestTool {
19571 id: "primary".to_string(),
19572 succeeds: false,
19573 calls: Arc::clone(&calls),
19574 max_output_chars: None,
19575 }))
19576 .recovery_manager(recovery_manager_with_fallbacks([(
19577 "primary".to_string(),
19578 "primary".to_string(),
19579 )]))
19580 .hooks(hooks.clone())
19581 .build()
19582 .unwrap();
19583
19584 let record = tokio::time::timeout(
19585 std::time::Duration::from_secs(2),
19586 agent.invoke_tool(ToolExecutionRequest::new(
19587 "self-fallback-call",
19588 "primary",
19589 serde_json::json!({"path": "./shared.txt"}),
19590 ToolCallSource::Manual,
19591 )),
19592 )
19593 .await
19594 .expect("self fallback must terminate without recursive execution")
19595 .unwrap();
19596
19597 assert_eq!(calls.load(Ordering::SeqCst), 1);
19598 assert_eq!(record.canonical_id, "primary");
19599 assert!(!record.executed);
19600 assert!(!record.success);
19601 assert_eq!(record.policy.outcome, PermissionOutcome::Deny);
19602 assert!(record.output.contains("fallback cycle"));
19603 assert!(matches!(
19604 record.source,
19605 ToolCallSource::Fallback { ref original_tool } if original_tool == "primary"
19606 ));
19607 assert_eq!(
19608 record.metadata.get("fallback_chain"),
19609 Some(&serde_json::json!(["primary"]))
19610 );
19611 assert_eq!(
19612 hooks.events(),
19613 vec![
19614 "start:primary",
19615 "complete:primary:false",
19616 "record:primary:true",
19617 "error",
19618 "complete:primary:false",
19619 "record:primary:false",
19620 "error",
19621 ]
19622 );
19623 assert_eq!(agent.tool_call_history().len(), 2);
19624 }
19625
19626 #[tokio::test]
19628 async fn alias_mediated_fallback_cycle_is_denied_canonically() {
19629 let primary_calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19630 let secondary_calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19631 let hooks = Arc::new(ToolLifecycleRecordingHooks::new());
19632 let agent = AgentBuilder::new()
19633 .system_prompt("Test canonical fallback cycle admission.")
19634 .llm(Arc::new(mock_with_response("done")))
19635 .tool(Arc::new(RecoveryTestTool {
19636 id: "primary".to_string(),
19637 succeeds: false,
19638 calls: Arc::clone(&primary_calls),
19639 max_output_chars: None,
19640 }))
19641 .tool(Arc::new(RecoveryTestTool {
19642 id: "secondary".to_string(),
19643 succeeds: false,
19644 calls: Arc::clone(&secondary_calls),
19645 max_output_chars: None,
19646 }))
19647 .recovery_manager(recovery_manager_with_fallbacks([
19648 ("primary".to_string(), "secondary".to_string()),
19649 ("secondary".to_string(), "primary alias".to_string()),
19650 ]))
19651 .hooks(hooks.clone())
19652 .build()
19653 .unwrap();
19654 agent.tools.set_tool_aliases(
19655 "primary",
19656 ToolAliases::new().with_name("en", "primary alias"),
19657 );
19658
19659 let record = tokio::time::timeout(
19660 std::time::Duration::from_secs(2),
19661 agent.invoke_tool(ToolExecutionRequest::new(
19662 "alias-fallback-call",
19663 "primary",
19664 serde_json::json!({"path": "./shared.txt"}),
19665 ToolCallSource::Manual,
19666 )),
19667 )
19668 .await
19669 .expect("alias-mediated fallback cycle must terminate")
19670 .unwrap();
19671
19672 assert_eq!(primary_calls.load(Ordering::SeqCst), 1);
19673 assert_eq!(secondary_calls.load(Ordering::SeqCst), 1);
19674 assert_eq!(record.requested_name, "primary alias");
19675 assert_eq!(record.canonical_id, "primary");
19676 assert!(!record.executed);
19677 assert!(record.output.contains("fallback cycle"));
19678 assert_eq!(
19679 record.metadata.get("fallback_chain"),
19680 Some(&serde_json::json!(["primary", "secondary"]))
19681 );
19682 assert_eq!(hooks.records().len(), 3);
19683 assert_eq!(agent.tool_call_history().len(), 3);
19684 }
19685
19686 #[tokio::test]
19688 async fn final_canonical_drift_cannot_bypass_fallback_ancestry() {
19689 let primary_calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19690 let secondary_calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19691 let provider = Arc::new(DriftingFallbackProvider {
19692 refreshed: AtomicBool::new(false),
19693 primary_calls: Arc::clone(&primary_calls),
19694 secondary_calls: Arc::clone(&secondary_calls),
19695 });
19696 let registry = ToolRegistry::new();
19697 registry.register_provider(provider).await.unwrap();
19698 let lifecycle = Arc::new(ToolLifecycleRecordingHooks::new());
19699 let hooks = Arc::new(RefreshFallbackProviderHooks {
19700 agent: parking_lot::Mutex::new(None),
19701 lifecycle: Arc::clone(&lifecycle),
19702 });
19703 let agent = Arc::new(
19704 AgentBuilder::new()
19705 .system_prompt("Test final canonical fallback admission.")
19706 .llm(Arc::new(mock_with_response("done")))
19707 .tools(registry)
19708 .recovery_manager(recovery_manager_with_fallbacks([(
19709 "primary".to_string(),
19710 "fallback alias".to_string(),
19711 )]))
19712 .hooks(hooks.clone())
19713 .build()
19714 .unwrap(),
19715 );
19716 *hooks.agent.lock() = Some(Arc::downgrade(&agent));
19717
19718 let record = agent
19719 .invoke_tool(ToolExecutionRequest::new(
19720 "drifting-fallback-call",
19721 "primary",
19722 serde_json::json!({"path": "./shared.txt"}),
19723 ToolCallSource::Manual,
19724 ))
19725 .await
19726 .unwrap();
19727
19728 assert_eq!(primary_calls.load(Ordering::SeqCst), 1);
19729 assert_eq!(secondary_calls.load(Ordering::SeqCst), 0);
19730 assert_eq!(record.canonical_id, "secondary");
19731 assert!(!record.executed);
19732 assert_eq!(record.policy.outcome, PermissionOutcome::Deny);
19733 assert!(record.output.contains("fallback cycle"));
19734 assert_eq!(
19735 record.metadata.get("fallback_chain"),
19736 Some(&serde_json::json!(["primary", "secondary"]))
19737 );
19738 assert_eq!(
19739 record.metadata.get("final_resolved_canonical_id"),
19740 Some(&serde_json::json!("primary"))
19741 );
19742 assert_eq!(
19743 lifecycle.events(),
19744 vec![
19745 "start:primary",
19746 "complete:primary:false",
19747 "record:primary:true",
19748 "error",
19749 "start:secondary",
19750 "complete:secondary:false",
19751 "record:secondary:false",
19752 "error",
19753 ]
19754 );
19755 let records = lifecycle.records();
19756 assert_eq!(records.len(), 2);
19757 assert_eq!(records[1].canonical_id, "secondary");
19758 assert_eq!(
19759 records[1].metadata.get("final_resolved_canonical_id"),
19760 Some(&serde_json::json!("primary"))
19761 );
19762 let history = agent.tool_call_history();
19763 assert_eq!(
19764 history
19765 .iter()
19766 .map(|entry| entry.tool_id.as_str())
19767 .collect::<Vec<_>>(),
19768 vec!["primary", "secondary"]
19769 );
19770 }
19771
19772 #[tokio::test]
19774 async fn acyclic_fallback_chain_is_denied_after_the_hop_limit() {
19775 let tool_count = MAX_TOOL_FALLBACK_HOPS + 2;
19776 let calls = (0..tool_count)
19777 .map(|_| Arc::new(std::sync::atomic::AtomicUsize::new(0)))
19778 .collect::<Vec<_>>();
19779 let mut builder = AgentBuilder::new()
19780 .system_prompt("Test bounded acyclic fallback admission.")
19781 .llm(Arc::new(mock_with_response("done")));
19782 for (index, counter) in calls.iter().enumerate() {
19783 builder = builder.tool(Arc::new(RecoveryTestTool {
19784 id: format!("fallback_{index}"),
19785 succeeds: false,
19786 calls: Arc::clone(counter),
19787 max_output_chars: None,
19788 }));
19789 }
19790 let fallbacks = (0..tool_count - 1).map(|index| {
19791 (
19792 format!("fallback_{index}"),
19793 format!("fallback_{}", index + 1),
19794 )
19795 });
19796 let agent = builder
19797 .recovery_manager(recovery_manager_with_fallbacks(fallbacks))
19798 .build()
19799 .unwrap();
19800
19801 let record = tokio::time::timeout(
19802 std::time::Duration::from_secs(2),
19803 agent.invoke_tool(ToolExecutionRequest::new(
19804 "bounded-fallback-call",
19805 "fallback_0",
19806 serde_json::json!({"path": "./shared.txt"}),
19807 ToolCallSource::Manual,
19808 )),
19809 )
19810 .await
19811 .expect("bounded fallback chain must terminate")
19812 .unwrap();
19813
19814 for counter in calls.iter().take(MAX_TOOL_FALLBACK_HOPS + 1) {
19815 assert_eq!(counter.load(Ordering::SeqCst), 1);
19816 }
19817 assert_eq!(calls[MAX_TOOL_FALLBACK_HOPS + 1].load(Ordering::SeqCst), 0);
19818 assert_eq!(
19819 record.canonical_id,
19820 format!("fallback_{}", MAX_TOOL_FALLBACK_HOPS + 1)
19821 );
19822 assert!(!record.executed);
19823 assert!(record.output.contains("maximum of 16 hops"));
19824 assert_eq!(agent.tool_call_history().len(), tool_count);
19825 }
19826
19827 #[tokio::test]
19828 async fn diagnostics_without_provider_records_unavailable_without_execution() {
19829 let mock = mock_with_response("hello");
19830 let yaml = r#"
19831name: DiagnosticsNoProviderAgent
19832system_prompt: "Review diagnostics."
19833tools: [diagnostics]
19834"#;
19835 let agent = AgentBuilder::from_yaml(yaml)
19836 .unwrap()
19837 .llm(Arc::new(mock))
19838 .auto_configure_features()
19839 .unwrap()
19840 .build()
19841 .unwrap();
19842
19843 let record = agent
19844 .invoke_tool(ToolExecutionRequest::new(
19845 "diagnostics-call",
19846 "diagnostics",
19847 serde_json::json!({}),
19848 ToolCallSource::Manual,
19849 ))
19850 .await
19851 .unwrap();
19852
19853 assert!(!record.executed);
19854 assert!(!record.success);
19855 assert_eq!(record.policy.outcome, PermissionOutcome::Unavailable);
19856 }
19857
19858 #[tokio::test]
19859 async fn web_search_without_provider_records_unavailable_without_execution() {
19860 let mock = mock_with_response("hello");
19861 let yaml = r#"
19862name: WebSearchNoProviderAgent
19863system_prompt: "You search the web."
19864tools: [web_search]
19865"#;
19866 let agent = AgentBuilder::from_yaml(yaml)
19867 .unwrap()
19868 .llm(Arc::new(mock))
19869 .auto_configure_features()
19870 .unwrap()
19871 .build()
19872 .unwrap();
19873
19874 let record = agent
19875 .invoke_tool(ToolExecutionRequest::new(
19876 "web-search-call",
19877 "web_search",
19878 serde_json::json!({"query": "rust async"}),
19879 ToolCallSource::Manual,
19880 ))
19881 .await
19882 .unwrap();
19883
19884 assert!(!record.executed);
19885 assert!(!record.success);
19886 assert_eq!(record.policy.outcome, PermissionOutcome::Unavailable);
19887 }
19888
19889 #[tokio::test]
19890 async fn unavailable_host_tool_does_not_request_approval() {
19891 let approvals = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19892 let handler = Arc::new(CountingApprovalHandler {
19893 calls: Arc::clone(&approvals),
19894 });
19895 let mut security = ToolSecurityConfig {
19896 enabled: true,
19897 fail_closed: true,
19898 ..Default::default()
19899 };
19900 security.tools.insert(
19901 "web_search".to_string(),
19902 ai_agents_tools::ToolPolicyConfig {
19903 enabled: true,
19904 require_confirmation: true,
19905 ..Default::default()
19906 },
19907 );
19908 let yaml = r#"
19909name: UnavailableApprovalAgent
19910system_prompt: "Search only with approval."
19911tools: [web_search]
19912"#;
19913 let agent = AgentBuilder::from_yaml(yaml)
19914 .unwrap()
19915 .llm(Arc::new(mock_with_response("done")))
19916 .auto_configure_features()
19917 .unwrap()
19918 .tool_security(ToolSecurityEngine::new(security))
19919 .hitl_engine(HITLEngine::new(ai_agents_hitl::HITLConfig::default()))
19920 .approval_handler(handler)
19921 .build()
19922 .unwrap();
19923
19924 let record = agent
19925 .invoke_tool(ToolExecutionRequest::new(
19926 "unavailable-before-approval",
19927 "web_search",
19928 serde_json::json!({"query": "rust async"}),
19929 ToolCallSource::Manual,
19930 ))
19931 .await
19932 .unwrap();
19933
19934 assert_eq!(approvals.load(Ordering::SeqCst), 0);
19935 assert!(!record.executed);
19936 assert!(!record.success);
19937 assert_eq!(record.policy.outcome, PermissionOutcome::Unavailable);
19938 assert!(
19939 record
19940 .approval
19941 .as_ref()
19942 .is_some_and(|approval| matches!(approval.status, ToolApprovalStatus::Unavailable))
19943 );
19944 }
19945
19946 #[tokio::test]
19947 async fn test_spawner_section_does_not_grant_core_tools_when_top_level_tools_omitted() {
19948 let mock = mock_with_response("hello");
19949 let yaml = r#"
19950name: SpawnerNoGrantAgent
19951system_prompt: "You manage agents."
19952spawner:
19953 max_agents: 2
19954"#;
19955 let agent = AgentBuilder::from_yaml(yaml)
19956 .unwrap()
19957 .llm(Arc::new(mock))
19958 .auto_configure_features()
19959 .unwrap()
19960 .auto_configure_spawner()
19961 .await
19962 .unwrap()
19963 .build()
19964 .unwrap();
19965
19966 let available = agent.get_available_tool_ids().await.unwrap();
19967 assert!(available.is_empty());
19968 }
19969
19970 #[tokio::test]
19971 async fn test_spawner_section_does_not_grant_core_tools_when_top_level_tools_empty() {
19972 let mock = mock_with_response("hello");
19973 let yaml = r#"
19974name: EmptySpawnerNoGrantAgent
19975system_prompt: "You manage agents."
19976tools: []
19977spawner:
19978 max_agents: 2
19979"#;
19980 let agent = AgentBuilder::from_yaml(yaml)
19981 .unwrap()
19982 .llm(Arc::new(mock))
19983 .auto_configure_features()
19984 .unwrap()
19985 .auto_configure_spawner()
19986 .await
19987 .unwrap()
19988 .build()
19989 .unwrap();
19990
19991 let available = agent.get_available_tool_ids().await.unwrap();
19992 assert!(available.is_empty());
19993 }
19994
19995 #[tokio::test]
19996 async fn test_management_tools_flag_grants_core_tools_when_top_level_tools_empty() {
19997 let mock = mock_with_response("hello");
19998 let yaml = r#"
19999name: ManagementGrantAgent
20000system_prompt: "You manage agents."
20001tools: []
20002spawner:
20003 management_tools: true
20004"#;
20005 let agent = AgentBuilder::from_yaml(yaml)
20006 .unwrap()
20007 .llm(Arc::new(mock))
20008 .auto_configure_features()
20009 .unwrap()
20010 .auto_configure_spawner()
20011 .await
20012 .unwrap()
20013 .build()
20014 .unwrap();
20015
20016 let available = agent.get_available_tool_ids().await.unwrap();
20017 assert_eq!(available.len(), 4);
20018 assert!(available.contains(&"spawn_agent".to_string()));
20019 assert!(available.contains(&"send_agent_message".to_string()));
20020 assert!(available.contains(&"list_agents".to_string()));
20021 assert!(available.contains(&"remove_agent".to_string()));
20022 }
20023
20024 #[tokio::test]
20025 async fn test_management_tools_flag_grants_core_tools_when_top_level_tools_omitted() {
20026 let mock = mock_with_response("hello");
20027 let yaml = r#"
20028name: ManagementOmittedToolsGrantAgent
20029system_prompt: "You manage agents."
20030spawner:
20031 management_tools: true
20032"#;
20033 let agent = AgentBuilder::from_yaml(yaml)
20034 .unwrap()
20035 .llm(Arc::new(mock))
20036 .auto_configure_features()
20037 .unwrap()
20038 .auto_configure_spawner()
20039 .await
20040 .unwrap()
20041 .build()
20042 .unwrap();
20043
20044 let available = agent.get_available_tool_ids().await.unwrap();
20045 assert_eq!(available.len(), 4);
20046 assert!(available.contains(&"spawn_agent".to_string()));
20047 assert!(available.contains(&"send_agent_message".to_string()));
20048 assert!(available.contains(&"list_agents".to_string()));
20049 assert!(available.contains(&"remove_agent".to_string()));
20050 }
20051
20052 #[tokio::test]
20053 async fn test_management_tools_selected_grants_only_selected_tools() {
20054 let mock = mock_with_response("hello");
20055 let yaml = r#"
20056name: ManagementSelectedGrantAgent
20057system_prompt: "You manage agents."
20058tools: []
20059spawner:
20060 management_tools:
20061 - spawn_agent
20062 - send_agent_message
20063 - list_agents
20064"#;
20065 let agent = AgentBuilder::from_yaml(yaml)
20066 .unwrap()
20067 .llm(Arc::new(mock))
20068 .auto_configure_features()
20069 .unwrap()
20070 .auto_configure_spawner()
20071 .await
20072 .unwrap()
20073 .build()
20074 .unwrap();
20075
20076 let available = agent.get_available_tool_ids().await.unwrap();
20077 assert_eq!(available.len(), 3);
20078 assert!(available.contains(&"spawn_agent".to_string()));
20079 assert!(available.contains(&"send_agent_message".to_string()));
20080 assert!(available.contains(&"list_agents".to_string()));
20081 assert!(!available.contains(&"remove_agent".to_string()));
20082 }
20083
20084 #[tokio::test]
20085 async fn test_orchestration_tools_flag_grants_tools_when_top_level_tools_empty() {
20086 let mock = mock_with_response("hello");
20087 let yaml = r#"
20088name: OrchestrationGrantAgent
20089system_prompt: "You coordinate agents."
20090llms:
20091 default:
20092 provider: openai
20093 model: gpt-4
20094 router:
20095 provider: openai
20096 model: gpt-4
20097llm:
20098 default: default
20099 router: router
20100tools: []
20101spawner:
20102 orchestration_tools: true
20103"#;
20104 let agent = AgentBuilder::from_yaml(yaml)
20105 .unwrap()
20106 .llm(Arc::new(mock))
20107 .auto_configure_features()
20108 .unwrap()
20109 .auto_configure_spawner()
20110 .await
20111 .unwrap()
20112 .build()
20113 .unwrap();
20114
20115 let available = agent.get_available_tool_ids().await.unwrap();
20116 assert_eq!(available.len(), 5);
20117 assert!(available.contains(&"route_to_agent".to_string()));
20118 assert!(available.contains(&"pipeline_process".to_string()));
20119 assert!(available.contains(&"concurrent_ask".to_string()));
20120 assert!(available.contains(&"group_discussion".to_string()));
20121 assert!(available.contains(&"handoff_conversation".to_string()));
20122 }
20123
20124 #[tokio::test]
20125 async fn test_persona_evolve_flag_grants_tool_when_top_level_tools_empty() {
20126 let mock = mock_with_response("hello");
20127 let yaml = r#"
20128name: PersonaGrantAgent
20129system_prompt: "You can evolve persona."
20130llm:
20131 provider: openai
20132 model: gpt-4
20133tools: []
20134persona:
20135 identity:
20136 name: "Guide"
20137 role: "Helper"
20138 evolution:
20139 enabled: true
20140 allow_llm_evolve: true
20141 mutable_fields:
20142 - traits.personality
20143"#;
20144 let agent = AgentBuilder::from_yaml(yaml)
20145 .unwrap()
20146 .llm(Arc::new(mock))
20147 .build()
20148 .unwrap();
20149
20150 let available = agent.get_available_tool_ids().await.unwrap();
20151 assert_eq!(available, vec!["persona_evolve".to_string()]);
20152 }
20153
20154 #[tokio::test]
20155 async fn test_persona_evolve_flag_grants_tool_when_top_level_tools_omitted() {
20156 let mock = mock_with_response("hello");
20157 let yaml = r#"
20158name: PersonaOmittedToolsGrantAgent
20159system_prompt: "You can evolve persona."
20160llm:
20161 provider: openai
20162 model: gpt-4
20163persona:
20164 identity:
20165 name: "Guide"
20166 role: "Helper"
20167 evolution:
20168 enabled: true
20169 allow_llm_evolve: true
20170 mutable_fields:
20171 - traits.personality
20172"#;
20173 let agent = AgentBuilder::from_yaml(yaml)
20174 .unwrap()
20175 .llm(Arc::new(mock))
20176 .build()
20177 .unwrap();
20178
20179 let available = agent.get_available_tool_ids().await.unwrap();
20180 assert_eq!(available, vec!["persona_evolve".to_string()]);
20181 }
20182
20183 #[tokio::test]
20184 async fn test_omitted_yaml_tools_exposes_no_tools() {
20185 let mock = mock_with_response("hello");
20186 let yaml = r#"
20187name: NoToolsAgent
20188system_prompt: "You are helpful."
20189"#;
20190 let agent = AgentBuilder::from_yaml(yaml)
20191 .unwrap()
20192 .llm(Arc::new(mock))
20193 .auto_configure_features()
20194 .unwrap()
20195 .build()
20196 .unwrap();
20197
20198 let available = agent.get_available_tool_ids().await.unwrap();
20199 assert!(available.is_empty());
20200 }
20201
20202 #[tokio::test]
20203 async fn runtime_scope_cannot_widen_omitted_or_empty_yaml_grants() {
20204 for tools in ["", "tools: []"] {
20205 let yaml = format!(
20206 r#"
20207name: RuntimeScopeNoGrantAgent
20208system_prompt: "No ordinary tools are granted."
20209{tools}
20210"#
20211 );
20212 let agent = AgentBuilder::from_yaml(&yaml)
20213 .unwrap()
20214 .llm(Arc::new(mock_with_response("done")))
20215 .auto_configure_features()
20216 .unwrap()
20217 .build()
20218 .unwrap();
20219
20220 agent
20221 .runtime_control()
20222 .set_tool_scope(vec!["calculator".to_string()]);
20223
20224 assert!(agent.get_available_tool_ids().await.unwrap().is_empty());
20225 }
20226 }
20227
20228 #[tokio::test]
20229 async fn runtime_scope_widening_attempt_keeps_only_declared_tools() {
20230 let yaml = r#"
20231name: RuntimeScopeWideningAgent
20232system_prompt: "Runtime scope cannot add authority."
20233tools: [calculator]
20234"#;
20235 let agent = AgentBuilder::from_yaml(yaml)
20236 .unwrap()
20237 .llm(Arc::new(mock_with_response("done")))
20238 .auto_configure_features()
20239 .unwrap()
20240 .build()
20241 .unwrap();
20242
20243 agent
20244 .runtime_control()
20245 .set_tool_scope(vec!["calculator".to_string(), "datetime".to_string()]);
20246
20247 assert_eq!(
20248 agent.get_available_tool_ids().await.unwrap(),
20249 vec!["calculator".to_string()]
20250 );
20251 }
20252
20253 #[tokio::test]
20254 async fn runtime_scope_is_canonical_unique_ordered_and_clear_restores_declared_grant() {
20255 let yaml = r#"
20256name: RuntimeScopeIntersectionAgent
20257system_prompt: "Use only declared tools."
20258tools: [calculator, datetime]
20259"#;
20260 let agent = AgentBuilder::from_yaml(yaml)
20261 .unwrap()
20262 .llm(Arc::new(mock_with_response("done")))
20263 .auto_configure_features()
20264 .unwrap()
20265 .build()
20266 .unwrap();
20267 let mut aliases = ai_agents_tools::ToolAliases::default();
20268 aliases
20269 .names
20270 .insert("en".to_string(), "calculate_alias".to_string());
20271 agent.tools.set_tool_aliases("calculator", aliases);
20272 let control = agent.runtime_control();
20273
20274 control.set_tool_scope(vec![
20275 "datetime".to_string(),
20276 "calculate_alias".to_string(),
20277 "calculator".to_string(),
20278 "unknown".to_string(),
20279 "datetime".to_string(),
20280 ]);
20281 assert_eq!(
20282 agent.get_available_tool_ids().await.unwrap(),
20283 vec!["calculator".to_string(), "datetime".to_string()]
20284 );
20285
20286 control.set_tool_scope(vec!["datetime".to_string()]);
20287 assert_eq!(
20288 agent.get_available_tool_ids().await.unwrap(),
20289 vec!["datetime".to_string()]
20290 );
20291
20292 control.clear_tool_scope_override();
20293 assert_eq!(
20294 agent.get_available_tool_ids().await.unwrap(),
20295 vec!["calculator".to_string(), "datetime".to_string()]
20296 );
20297 }
20298
20299 #[tokio::test]
20300 async fn runtime_scope_preserves_programmatic_registration_as_declared_grant() {
20301 let agent = AgentBuilder::new()
20302 .system_prompt("Use registered tools.")
20303 .llm(Arc::new(mock_with_response("done")))
20304 .tool(Arc::new(ContextEchoTool))
20305 .tool(Arc::new(SlowTool))
20306 .build()
20307 .unwrap();
20308
20309 agent.runtime_control().set_tool_scope(vec![
20310 "Context Echo".to_string(),
20311 "context_echo".to_string(),
20312 "unknown".to_string(),
20313 ]);
20314
20315 assert_eq!(
20316 agent.get_available_tool_ids().await.unwrap(),
20317 vec!["context_echo".to_string()]
20318 );
20319 }
20320
20321 #[tokio::test]
20322 async fn nested_state_scopes_intersect_every_ancestor_with_aliases() {
20323 let yaml = r#"
20324name: NestedStateScopeAgent
20325system_prompt: "Honor every state scope."
20326tools: [calculator, datetime, echo]
20327states:
20328 initial: root
20329 states:
20330 root:
20331 tools: [calculate_alias, datetime]
20332 initial: middle
20333 states:
20334 middle:
20335 initial: leaf
20336 states:
20337 leaf:
20338 tools: [datetime_alias, echo]
20339"#;
20340 let agent = AgentBuilder::from_yaml(yaml)
20341 .unwrap()
20342 .llm(Arc::new(mock_with_response("done")))
20343 .auto_configure_features()
20344 .unwrap()
20345 .build()
20346 .unwrap();
20347 let mut calculator_aliases = ai_agents_tools::ToolAliases::default();
20348 calculator_aliases
20349 .names
20350 .insert("en".to_string(), "calculate_alias".to_string());
20351 agent
20352 .tools
20353 .set_tool_aliases("calculator", calculator_aliases);
20354 let mut datetime_aliases = ai_agents_tools::ToolAliases::default();
20355 datetime_aliases
20356 .names
20357 .insert("en".to_string(), "datetime_alias".to_string());
20358 agent.tools.set_tool_aliases("datetime", datetime_aliases);
20359 agent.runtime_control().set_tool_scope(vec![
20360 "unknown".to_string(),
20361 "datetime_alias".to_string(),
20362 "calculate_alias".to_string(),
20363 "datetime".to_string(),
20364 ]);
20365
20366 assert_eq!(agent.current_state().as_deref(), Some("root.middle.leaf"));
20367 assert_eq!(
20368 agent.get_available_tool_ids().await.unwrap(),
20369 vec!["datetime".to_string()]
20370 );
20371 }
20372
20373 #[tokio::test]
20374 async fn ancestor_empty_state_scope_denies_omitted_descendants() {
20375 let yaml = r#"
20376name: NestedEmptyStateScopeAgent
20377system_prompt: "An empty ancestor scope denies all tools."
20378tools: [calculator]
20379states:
20380 initial: root
20381 states:
20382 root:
20383 tools: []
20384 initial: middle
20385 states:
20386 middle:
20387 initial: leaf
20388 states:
20389 leaf: {}
20390"#;
20391 let agent = AgentBuilder::from_yaml(yaml)
20392 .unwrap()
20393 .llm(Arc::new(mock_with_response("done")))
20394 .auto_configure_features()
20395 .unwrap()
20396 .build()
20397 .unwrap();
20398
20399 assert!(agent.get_available_tool_ids().await.unwrap().is_empty());
20400 }
20401
20402 #[tokio::test]
20403 async fn state_change_during_approval_invalidates_the_reviewed_authority() {
20404 let active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
20405 let max_active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
20406 let entered = Arc::new(tokio::sync::Barrier::new(2));
20407 let release = Arc::new(tokio::sync::Notify::new());
20408 let handler = Arc::new(BlockingApprovalHandler {
20409 entered: Arc::clone(&entered),
20410 release: Arc::clone(&release),
20411 result: ApprovalResult::Approved,
20412 });
20413 let yaml = r#"
20414name: ApprovalStateGenerationAgent
20415system_prompt: "State authority may change during approval."
20416tools: [locked_write]
20417states:
20418 initial: first
20419 states:
20420 first:
20421 tools: [locked_write]
20422 second:
20423 tools: [locked_write]
20424"#;
20425 let agent = Arc::new(
20426 AgentBuilder::from_yaml(yaml)
20427 .unwrap()
20428 .llm(Arc::new(mock_with_response("done")))
20429 .tool(Arc::new(LockedWriteTool {
20430 active: Arc::clone(&active),
20431 max_active: Arc::clone(&max_active),
20432 }))
20433 .tool_security(ToolSecurityEngine::new(approval_security_config(true)))
20434 .hitl_engine(HITLEngine::new(ai_agents_hitl::HITLConfig::default()))
20435 .approval_handler(handler)
20436 .build()
20437 .unwrap(),
20438 );
20439 let running = Arc::clone(&agent);
20440 let call = tokio::spawn(async move {
20441 running
20442 .invoke_tool(ToolExecutionRequest::new(
20443 "approval-state-generation",
20444 "locked_write",
20445 serde_json::json!({"path": "./state-generation.txt"}),
20446 ToolCallSource::Manual,
20447 ))
20448 .await
20449 .unwrap()
20450 });
20451
20452 entered.wait().await;
20453 agent.transition_to("second").await.unwrap();
20454 release.notify_one();
20455 let record = call.await.unwrap();
20456
20457 assert!(!record.executed);
20458 assert!(record.output.contains("Approval became stale"));
20459 assert_eq!(max_active.load(Ordering::SeqCst), 0);
20460 }
20461
20462 #[tokio::test]
20463 async fn state_change_while_waiting_for_resource_lock_fails_final_admission() {
20464 let holder_gate = PathMutationGate::new();
20465 let waiter_gate = PathMutationGate::new();
20466 let yaml = r#"
20467name: LockedStateGenerationAgent
20468system_prompt: "State authority must remain stable through admission."
20469tools: [state_lock_holder, state_lock_waiter]
20470states:
20471 initial: first
20472 states:
20473 first:
20474 tools: [state_lock_holder, state_lock_waiter]
20475 second:
20476 tools: [state_lock_holder, state_lock_waiter]
20477"#;
20478 let agent = Arc::new(
20479 AgentBuilder::from_yaml(yaml)
20480 .unwrap()
20481 .llm(Arc::new(mock_with_response("done")))
20482 .tool(Arc::new(BlockingPathMutationTool {
20483 id: "state_lock_holder",
20484 path_fields: vec![ai_agents_core::PathPolicyBinding::write("path")],
20485 gate: holder_gate.clone(),
20486 }))
20487 .tool(Arc::new(BlockingPathMutationTool {
20488 id: "state_lock_waiter",
20489 path_fields: vec![ai_agents_core::PathPolicyBinding::write("path")],
20490 gate: waiter_gate.clone(),
20491 }))
20492 .build()
20493 .unwrap(),
20494 );
20495 let holder_call = {
20496 let agent = Arc::clone(&agent);
20497 tokio::spawn(async move {
20498 agent
20499 .invoke_tool(ToolExecutionRequest::new(
20500 "state-lock-holder",
20501 "state_lock_holder",
20502 serde_json::json!({"path": "./shared-state-path.txt"}),
20503 ToolCallSource::Manual,
20504 ))
20505 .await
20506 .unwrap()
20507 })
20508 };
20509 holder_gate.wait_until_entered().await;
20510 let waiter_call = {
20511 let agent = Arc::clone(&agent);
20512 tokio::spawn(async move {
20513 agent
20514 .invoke_tool(ToolExecutionRequest::new(
20515 "state-lock-waiter",
20516 "state_lock_waiter",
20517 serde_json::json!({"path": "./shared-state-path.txt"}),
20518 ToolCallSource::Manual,
20519 ))
20520 .await
20521 .unwrap()
20522 })
20523 };
20524
20525 wait_for_resource_lock_strong_count(&agent.resource_locks, 2).await;
20526 agent.transition_to("second").await.unwrap();
20527 holder_gate.release();
20528 let holder_record = holder_call.await.unwrap();
20529 let waiter_record = waiter_call.await.unwrap();
20530
20531 assert!(holder_record.success);
20532 assert!(!waiter_record.executed);
20533 assert!(
20534 waiter_record
20535 .output
20536 .contains("state scope changed before admission")
20537 );
20538 assert!(!waiter_gate.entered.load(Ordering::SeqCst));
20539 }
20540
20541 #[tokio::test]
20542 async fn test_state_tools_cannot_widen_top_level_grant() {
20543 let mock = mock_with_response("hello");
20544 let yaml = r#"
20545name: NarrowToolsAgent
20546system_prompt: "You are helpful."
20547tools:
20548 - calculator
20549states:
20550 initial: current
20551 states:
20552 current:
20553 tools: [datetime]
20554"#;
20555 let agent = AgentBuilder::from_yaml(yaml)
20556 .unwrap()
20557 .llm(Arc::new(mock))
20558 .auto_configure_features()
20559 .unwrap()
20560 .build()
20561 .unwrap();
20562
20563 let available = agent.get_available_tool_ids().await.unwrap();
20564 assert!(available.is_empty());
20565 }
20566
20567 #[tokio::test]
20569 async fn test_integration_tool_execution() {
20570 let mock = mock_with_responses(vec![
20572 r#"I'll calculate that for you.
20574{"tool": "calculator", "arguments": {"expression": "2+2"}}"#,
20575 "The answer is 4.",
20577 ]);
20578 let observed = mock.clone();
20579 let mut tools = ai_agents_tools::ToolRegistry::new();
20580 tools
20581 .register(Arc::new(ai_agents_tools::CalculatorTool))
20582 .unwrap();
20583
20584 let agent = AgentBuilder::new()
20585 .system_prompt("You are a calculator assistant.")
20586 .llm(Arc::new(mock))
20587 .tools(tools)
20588 .build()
20589 .unwrap();
20590
20591 let response = agent.chat("What is 2+2?").await.unwrap();
20592
20593 assert_eq!(response.content, "The answer is 4.");
20594 assert_eq!(response.tool_calls.as_ref().map(Vec::len), Some(1));
20595 assert_eq!(
20596 observed.call_count(),
20597 2,
20598 "tool result must trigger a second LLM call"
20599 );
20600 let history = agent.tool_call_history();
20601 assert_eq!(history.len(), 1);
20602 assert_eq!(history[0].tool_id, "calculator");
20603 assert_eq!(
20604 history[0].result.get("result"),
20605 Some(&serde_json::json!(4.0)),
20606 "{:?}",
20607 history[0].result
20608 );
20609 }
20610
20611 #[test]
20614 fn legacy_tool_call_marker_is_plain_text() {
20615 let agent = AgentBuilder::new()
20616 .system_prompt("x")
20617 .llm(Arc::new(mock_with_response("x")))
20618 .build()
20619 .unwrap();
20620 let parsed = agent
20621 .parse_tool_calls(
20622 r#"[TOOL_CALL: {"name": "calculator", "arguments": {"expression": "2+2"}}]"#,
20623 )
20624 .unwrap();
20625 assert!(parsed.is_none());
20626 }
20627
20628 #[tokio::test]
20629 async fn test_tool_hitl_rejection_finalizes_blocking_turn() {
20630 let responses = Arc::new(std::sync::atomic::AtomicUsize::new(0));
20631 let hooks = Arc::new(ResponseCountingHooks {
20632 responses: Arc::clone(&responses),
20633 });
20634 let mock = mock_with_response(r#"{"tool":"echo","arguments":{"message":"hello"}}"#);
20635 let yaml = r#"
20636name: ToolRejectAgent
20637system_prompt: "You use tools when requested."
20638tools:
20639 - echo
20640hitl:
20641 tools:
20642 echo:
20643 require_approval: true
20644 approval_message: "Approve echo?"
20645"#;
20646 let agent = AgentBuilder::from_yaml(yaml)
20647 .unwrap()
20648 .llm(Arc::new(mock))
20649 .auto_configure_features()
20650 .unwrap()
20651 .hooks(hooks)
20652 .build()
20653 .unwrap();
20654
20655 let response = agent.chat("echo hello").await.unwrap();
20656
20657 assert!(
20658 response.content.contains("Operation cancelled"),
20659 "unexpected response: {}",
20660 response.content
20661 );
20662 assert_eq!(responses.load(Ordering::SeqCst), 1);
20663 let messages = agent.memory.get_messages(None).await.unwrap();
20664 assert_eq!(messages.len(), 3);
20665 assert_eq!(messages[0].content, "echo hello");
20666 assert!(messages[1].content.contains("\"tool\":\"echo\""));
20667 assert!(messages[2].content.contains("rejected by the approver"));
20668 }
20669
20670 #[tokio::test]
20671 async fn test_tool_hitl_rejection_finalizes_streaming_turn() {
20672 use futures::StreamExt;
20673
20674 let responses = Arc::new(std::sync::atomic::AtomicUsize::new(0));
20675 let hooks = Arc::new(ResponseCountingHooks {
20676 responses: Arc::clone(&responses),
20677 });
20678 let mock = mock_with_response(r#"{"tool":"echo","arguments":{"message":"hello"}}"#);
20679 let yaml = r#"
20680name: ToolRejectStreamingAgent
20681system_prompt: "You use tools when requested."
20682tools:
20683 - echo
20684streaming:
20685 enabled: true
20686hitl:
20687 tools:
20688 echo:
20689 require_approval: true
20690 approval_message: "Approve echo?"
20691"#;
20692 let agent = AgentBuilder::from_yaml(yaml)
20693 .unwrap()
20694 .llm(Arc::new(mock))
20695 .auto_configure_features()
20696 .unwrap()
20697 .hooks(hooks)
20698 .build()
20699 .unwrap();
20700
20701 let mut stream = agent.chat_stream("echo hello").await.unwrap();
20702 let mut terminal_error = String::new();
20703 let mut done = false;
20704 while let Some(chunk) = stream.next().await {
20705 match chunk {
20706 StreamChunk::Error { message } => terminal_error = message,
20707 StreamChunk::Done {} => {
20708 done = true;
20709 break;
20710 }
20711 _ => {}
20712 }
20713 }
20714
20715 assert!(done);
20716 assert!(
20717 terminal_error.contains("Operation cancelled"),
20718 "unexpected terminal error: {}",
20719 terminal_error
20720 );
20721 assert_eq!(responses.load(Ordering::SeqCst), 1);
20722 let messages = agent.memory.get_messages(None).await.unwrap();
20723 assert_eq!(messages.len(), 3);
20724 assert_eq!(messages[0].content, "echo hello");
20725 assert!(messages[1].content.contains("\"tool\":\"echo\""));
20726 assert!(messages[2].content.contains("rejected by the approver"));
20727 }
20728
20729 #[tokio::test]
20730 async fn tool_hitl_rejection_preserves_legacy_error_but_finalizes_event_stream() {
20731 let mock = mock_with_response(r#"{"tool":"echo","arguments":{"message":"hello"}}"#);
20732 let yaml = r#"
20733name: ToolRejectEventAgent
20734system_prompt: "You use tools when requested."
20735tools:
20736 - echo
20737streaming:
20738 enabled: true
20739hitl:
20740 tools:
20741 echo:
20742 require_approval: true
20743 approval_message: "Approve echo?"
20744"#;
20745 let agent = AgentBuilder::from_yaml(yaml)
20746 .unwrap()
20747 .llm(Arc::new(mock))
20748 .auto_configure_features()
20749 .unwrap()
20750 .build()
20751 .unwrap();
20752
20753 let mut stream = agent.chat_stream_events("echo hello").await.unwrap();
20754 let mut error_seen = false;
20755 let mut final_response = None;
20756 while let Some(event) = stream.next().await {
20757 match event {
20758 AgentStreamEvent::Chunk(StreamChunk::Error { .. }) => error_seen = true,
20759 AgentStreamEvent::Final(response) => final_response = Some(response),
20760 AgentStreamEvent::Chunk(_) => {}
20761 }
20762 }
20763
20764 assert!(!error_seen);
20765 assert!(
20766 final_response
20767 .is_some_and(|response| { response.content.contains("Operation cancelled") })
20768 );
20769 }
20770
20771 #[tokio::test]
20772 async fn test_pre_response_guard_transition_skips_old_state_llm() {
20773 let mock = mock_with_response("Billing state response");
20774 let call_counter = mock.clone();
20775 let yaml = r#"
20776name: OptimizedStateAgent
20777system_prompt: "You route before answering."
20778runtime:
20779 optimization:
20780 enabled: true
20781 pre_response_deterministic_transitions: true
20782states:
20783 initial: greeting
20784 states:
20785 greeting:
20786 prompt: "Old state prompt that should be skipped."
20787 transitions:
20788 - to: billing
20789 guard:
20790 context:
20791 topic:
20792 eq: billing
20793 timing: pre_response
20794 billing:
20795 prompt: "Answer from the billing state."
20796"#;
20797 let agent = AgentBuilder::from_yaml(yaml)
20798 .unwrap()
20799 .llm(Arc::new(mock))
20800 .build()
20801 .unwrap();
20802 agent
20803 .set_context("topic", serde_json::json!("billing"))
20804 .unwrap();
20805
20806 let response = agent.chat("I need billing help").await.unwrap();
20807
20808 assert_eq!(agent.current_state().as_deref(), Some("billing"));
20809 assert_eq!(response.content, "Billing state response");
20810 assert_eq!(call_counter.call_count(), 1);
20811 assert_eq!(agent.actor_facts().len(), 0);
20812 }
20813
20814 #[tokio::test]
20815 async fn test_set_context_supports_dotted_paths_for_pre_response_guards() {
20816 let mock = mock_with_response("Billing state response");
20817 let call_counter = mock.clone();
20818 let yaml = r#"
20819name: OptimizedStateAgent
20820system_prompt: "You route before answering."
20821runtime:
20822 optimization:
20823 enabled: true
20824 pre_response_deterministic_transitions: true
20825context:
20826 request:
20827 type: runtime
20828 default:
20829 topic: general
20830states:
20831 initial: greeting
20832 states:
20833 greeting:
20834 prompt: "Old state prompt that should be skipped."
20835 transitions:
20836 - to: billing
20837 guard:
20838 context:
20839 request.topic:
20840 eq: billing
20841 timing: pre_response
20842 billing:
20843 prompt: "Answer from the billing state."
20844"#;
20845 let agent = AgentBuilder::from_yaml(yaml)
20846 .unwrap()
20847 .llm(Arc::new(mock))
20848 .build()
20849 .unwrap();
20850 agent
20851 .set_context("request.topic", serde_json::json!("billing"))
20852 .unwrap();
20853
20854 let response = agent.chat("I need billing help").await.unwrap();
20855
20856 assert_eq!(agent.current_state().as_deref(), Some("billing"));
20857 assert_eq!(response.content, "Billing state response");
20858 assert_eq!(call_counter.call_count(), 1);
20859 assert_eq!(
20860 agent.get_context().get("request"),
20861 Some(&serde_json::json!({"topic": "billing"}))
20862 );
20863 }
20864
20865 #[tokio::test]
20866 async fn test_pre_response_rejection_does_not_commit_staged_context_or_user() {
20867 let mock = mock_with_response("billing");
20868 let yaml = r#"
20869name: OptimizedStateAgent
20870system_prompt: "You route before answering."
20871runtime:
20872 optimization:
20873 enabled: true
20874 pre_response_deterministic_transitions: true
20875hitl:
20876 states:
20877 billing:
20878 on_enter: require_approval
20879 approval_message: "Approve billing route?"
20880states:
20881 initial: greeting
20882 states:
20883 greeting:
20884 prompt: "Old state prompt."
20885 extract:
20886 - key: topic
20887 description: "Support topic"
20888 transitions:
20889 - to: billing
20890 guard:
20891 context:
20892 topic:
20893 eq: billing
20894 timing: pre_response
20895 run_extractors: true
20896 billing:
20897 prompt: "Billing state."
20898"#;
20899 let agent = AgentBuilder::from_yaml(yaml)
20900 .unwrap()
20901 .llm(Arc::new(mock))
20902 .build()
20903 .unwrap();
20904
20905 let response = agent
20906 .try_pre_response_transition("billing please")
20907 .await
20908 .unwrap();
20909
20910 assert!(response.is_none());
20911 assert_eq!(agent.current_state().as_deref(), Some("greeting"));
20912 assert!(!agent.get_context().contains_key("topic"));
20913 assert_eq!(agent.memory.get_messages(None).await.unwrap().len(), 0);
20914 }
20915
20916 #[tokio::test]
20917 async fn test_pre_response_extractor_commits_context_on_winning_path() {
20918 let mock = mock_with_responses(vec!["billing", "Billing response"]);
20919 let yaml = r#"
20920name: OptimizedStateAgent
20921system_prompt: "You route before answering."
20922runtime:
20923 optimization:
20924 enabled: true
20925 pre_response_deterministic_transitions: true
20926states:
20927 initial: greeting
20928 states:
20929 greeting:
20930 prompt: "Old state prompt."
20931 extract:
20932 - key: topic
20933 description: "Support topic"
20934 transitions:
20935 - to: billing
20936 guard:
20937 context:
20938 topic:
20939 eq: billing
20940 timing: pre_response
20941 run_extractors: true
20942 billing:
20943 prompt: "Billing state."
20944"#;
20945 let agent = AgentBuilder::from_yaml(yaml)
20946 .unwrap()
20947 .llm(Arc::new(mock))
20948 .build()
20949 .unwrap();
20950
20951 let response = agent.chat("billing please").await.unwrap();
20952
20953 assert_eq!(agent.current_state().as_deref(), Some("billing"));
20954 assert_eq!(response.content, "Billing response");
20955 assert_eq!(
20956 agent.get_context().get("topic"),
20957 Some(&serde_json::json!("billing"))
20958 );
20959 }
20960
20961 #[tokio::test]
20962 async fn test_pre_response_extractor_miss_does_not_mutate_context() {
20963 let mock = mock_with_response("__NONE__");
20964 let yaml = r#"
20965name: OptimizedStateAgent
20966system_prompt: "You route before answering."
20967runtime:
20968 optimization:
20969 enabled: true
20970 pre_response_deterministic_transitions: true
20971states:
20972 initial: greeting
20973 states:
20974 greeting:
20975 prompt: "Old state prompt."
20976 extract:
20977 - key: topic
20978 description: "Support topic"
20979 transitions:
20980 - to: billing
20981 guard:
20982 context:
20983 topic:
20984 eq: billing
20985 timing: pre_response
20986 run_extractors: true
20987 billing:
20988 prompt: "Billing state."
20989"#;
20990 let agent = AgentBuilder::from_yaml(yaml)
20991 .unwrap()
20992 .llm(Arc::new(mock))
20993 .build()
20994 .unwrap();
20995
20996 let response = agent.try_pre_response_transition("hello").await.unwrap();
20997
20998 assert!(response.is_none());
20999 assert_eq!(agent.current_state().as_deref(), Some("greeting"));
21000 assert!(!agent.get_context().contains_key("topic"));
21001 }
21002
21003 #[tokio::test]
21004 async fn test_default_guard_transition_stays_post_response() {
21005 let mock = mock_with_responses(vec!["Greeting response", "Billing response"]);
21006 let call_counter = mock.clone();
21007 let yaml = r#"
21008name: TimingAgent
21009system_prompt: "You route carefully."
21010runtime:
21011 optimization:
21012 enabled: true
21013 pre_response_deterministic_transitions: true
21014states:
21015 initial: greeting
21016 states:
21017 greeting:
21018 prompt: "Old state prompt."
21019 transitions:
21020 - to: billing
21021 guard:
21022 context:
21023 topic:
21024 eq: billing
21025 billing:
21026 prompt: "Billing state."
21027"#;
21028 let agent = AgentBuilder::from_yaml(yaml)
21029 .unwrap()
21030 .llm(Arc::new(mock))
21031 .build()
21032 .unwrap();
21033 agent
21034 .set_context("topic", serde_json::json!("billing"))
21035 .unwrap();
21036
21037 let response = agent.chat("billing please").await.unwrap();
21038
21039 assert_eq!(agent.current_state().as_deref(), Some("billing"));
21040 assert_eq!(response.content, "Billing response");
21041 assert_eq!(call_counter.call_count(), 2);
21042 }
21043
21044 #[tokio::test]
21045 async fn test_explicit_post_response_guard_transition_stays_post_response() {
21046 let mock = mock_with_responses(vec!["Greeting response", "Billing response"]);
21047 let call_counter = mock.clone();
21048 let yaml = r#"
21049name: TimingAgent
21050system_prompt: "You route carefully."
21051runtime:
21052 optimization:
21053 enabled: true
21054 pre_response_deterministic_transitions: true
21055states:
21056 initial: greeting
21057 states:
21058 greeting:
21059 prompt: "Old state prompt."
21060 transitions:
21061 - to: billing
21062 guard:
21063 context:
21064 topic:
21065 eq: billing
21066 timing: post_response
21067 billing:
21068 prompt: "Billing state."
21069"#;
21070 let agent = AgentBuilder::from_yaml(yaml)
21071 .unwrap()
21072 .llm(Arc::new(mock))
21073 .build()
21074 .unwrap();
21075 agent
21076 .set_context("topic", serde_json::json!("billing"))
21077 .unwrap();
21078
21079 let response = agent.chat("billing please").await.unwrap();
21080
21081 assert_eq!(agent.current_state().as_deref(), Some("billing"));
21082 assert_eq!(response.content, "Billing response");
21083 assert_eq!(call_counter.call_count(), 2);
21084 }
21085
21086 #[tokio::test]
21087 async fn test_pre_response_extractors_are_transition_scoped() {
21088 let mock = mock_with_responses(vec!["billing", "Billing response"]);
21089 let yaml = r#"
21090name: ScopedExtractorAgent
21091system_prompt: "You route carefully."
21092runtime:
21093 optimization:
21094 enabled: true
21095 pre_response_deterministic_transitions: true
21096states:
21097 initial: greeting
21098 states:
21099 greeting:
21100 prompt: "Old state prompt."
21101 extract:
21102 - key: topic
21103 description: "Support topic"
21104 transitions:
21105 - to: wrong
21106 guard:
21107 context:
21108 topic:
21109 eq: billing
21110 timing: pre_response
21111 - to: billing
21112 guard:
21113 context:
21114 topic:
21115 eq: billing
21116 timing: pre_response
21117 run_extractors: true
21118 wrong:
21119 prompt: "Wrong state."
21120 billing:
21121 prompt: "Billing state."
21122"#;
21123 let agent = AgentBuilder::from_yaml(yaml)
21124 .unwrap()
21125 .llm(Arc::new(mock))
21126 .build()
21127 .unwrap();
21128
21129 let response = agent.chat("billing please").await.unwrap();
21130
21131 assert_eq!(agent.current_state().as_deref(), Some("billing"));
21132 assert_eq!(response.content, "Billing response");
21133 }
21134
21135 #[tokio::test]
21136 async fn test_pre_response_resolved_intent_routes_early() {
21137 let mock = mock_with_response("Billing response");
21138 let yaml = r#"
21139name: IntentAgent
21140system_prompt: "You route carefully."
21141runtime:
21142 optimization:
21143 enabled: true
21144 pre_response_deterministic_transitions: true
21145states:
21146 initial: greeting
21147 states:
21148 greeting:
21149 prompt: "Old state prompt."
21150 transitions:
21151 - to: billing
21152 intent: billing
21153 timing: pre_response
21154 billing:
21155 prompt: "Billing state."
21156"#;
21157 let agent = AgentBuilder::from_yaml(yaml)
21158 .unwrap()
21159 .llm(Arc::new(mock))
21160 .build()
21161 .unwrap();
21162 agent
21163 .set_context("resolved_intent", serde_json::json!("billing"))
21164 .unwrap();
21165
21166 let response = agent
21167 .try_pre_response_transition("I need billing help")
21168 .await
21169 .unwrap()
21170 .unwrap();
21171
21172 assert_eq!(agent.current_state().as_deref(), Some("billing"));
21173 assert_eq!(response.content, "Billing response");
21174 }
21175
21176 #[tokio::test]
21177 async fn test_background_overflow_error_surfaces() {
21178 let mut config = RuntimeConfig::default();
21179 config.optimization.enabled = true;
21180 config.optimization.post_turn.max_background_tasks = 1;
21181 config.optimization.post_turn.on_background_overflow = BackgroundOverflowPolicy::Error;
21182 let policy = crate::optimization::MaintenanceTaskPolicy {
21183 mode: MaintenanceMode::Background,
21184 await_before_next_turn: AwaitBeforeNextTurn::Always,
21185 };
21186 let agent = AgentBuilder::new()
21187 .system_prompt("You are helpful.")
21188 .llm(Arc::new(mock_with_response("ok")))
21189 .build()
21190 .unwrap()
21191 .with_runtime_config(config);
21192 agent
21193 .background_maintenance
21194 .spawn(None, async { std::future::pending::<Result<()>>().await })
21195 .unwrap();
21196
21197 let result = agent
21198 .spawn_or_handle_background(None, async { Ok(()) }, "facts", &policy)
21199 .await;
21200
21201 assert!(result.is_err());
21202 }
21203
21204 #[tokio::test]
21205 async fn test_speculative_reasoning_low_cap_uses_serial_reasoning() {
21206 let default_mock = mock_with_response("Plain draft response");
21207 let router_mock = mock_with_response("cot");
21208 let router_counter = router_mock.clone();
21209 let yaml = r#"
21210name: ReasoningReservationAgent
21211system_prompt: "You answer plainly unless reasoning wins."
21212llm:
21213 default: default
21214 router: router
21215observability:
21216 enabled: true
21217 export:
21218 write_raw_events: true
21219reasoning:
21220 mode: auto
21221 judge_llm: router
21222runtime:
21223 optimization:
21224 enabled: true
21225 max_speculative_llm_calls_per_turn: 1
21226 speculative_reasoning_auto: true
21227 max_parallel_runtime_tasks: 2
21228"#;
21229 let agent = AgentBuilder::from_yaml(yaml)
21230 .unwrap()
21231 .llm_alias("default", Arc::new(default_mock))
21232 .llm_alias("router", Arc::new(router_mock))
21233 .build()
21234 .unwrap();
21235
21236 let response = agent.chat("hello").await.unwrap();
21237
21238 assert_eq!(response.content, "Plain draft response");
21239 assert_eq!(router_counter.call_count(), 1);
21240 let events = agent.observability().unwrap().raw_events();
21241 assert!(!events.iter().any(|event| {
21242 event.dimensions.get("commit_behavior") == Some(&"reasoning_decision".to_string())
21243 }));
21244 }
21245
21246 #[tokio::test]
21247 async fn test_forced_reasoning_skips_plain_speculative_draft() {
21248 let mock = mock_with_response("Reasoned response");
21249 let yaml = r#"
21250name: ForcedReasoningAgent
21251system_prompt: "You reason before answering."
21252observability:
21253 enabled: true
21254 export:
21255 write_raw_events: true
21256reasoning:
21257 mode: cot
21258runtime:
21259 optimization:
21260 enabled: true
21261 max_speculative_llm_calls_per_turn: 2
21262 speculative_state_transitions: true
21263 max_parallel_runtime_tasks: 2
21264states:
21265 initial: triage
21266 states:
21267 triage:
21268 prompt: "Answer from triage."
21269 transitions:
21270 - to: billing
21271 guard:
21272 context:
21273 route:
21274 eq: billing
21275 timing: parallel
21276 billing:
21277 prompt: "Billing state."
21278"#;
21279 let agent = AgentBuilder::from_yaml(yaml)
21280 .unwrap()
21281 .llm(Arc::new(mock))
21282 .build()
21283 .unwrap();
21284
21285 let response = agent.chat("hello").await.unwrap();
21286
21287 assert_eq!(response.content, "Reasoned response");
21288 let events = agent.observability().unwrap().raw_events();
21289 assert!(
21290 !events
21291 .iter()
21292 .any(|event| event.dimensions.contains_key("branch_status"))
21293 );
21294 }
21295
21296 #[tokio::test]
21297 async fn test_speculative_skill_low_cap_uses_serial_skill_route() {
21298 let default_mock = mock_with_response("Skill committed response");
21299 let router_mock = mock_with_response("helper");
21300 let router_counter = router_mock.clone();
21301 let yaml = r#"
21302name: SkillReservationAgent
21303system_prompt: "Use skills when they match."
21304llm:
21305 default: default
21306 router: router
21307observability:
21308 enabled: true
21309 export:
21310 write_raw_events: true
21311runtime:
21312 optimization:
21313 enabled: true
21314 max_speculative_llm_calls_per_turn: 1
21315 speculative_skill_routing: true
21316 max_parallel_runtime_tasks: 2
21317skills:
21318 - id: helper
21319 description: "Answer helper requests"
21320 trigger: "User asks for helper"
21321 steps:
21322 - prompt: "Answer the helper request: {{ user_input }}"
21323"#;
21324 let agent = AgentBuilder::from_yaml(yaml)
21325 .unwrap()
21326 .llm_alias("default", Arc::new(default_mock))
21327 .llm_alias("router", Arc::new(router_mock))
21328 .build()
21329 .unwrap();
21330
21331 let response = agent.chat("please use helper").await.unwrap();
21332
21333 assert_eq!(response.content, "Skill committed response");
21334 assert_eq!(router_counter.call_count(), 1);
21335 let events = agent.observability().unwrap().raw_events();
21336 assert!(
21337 !events
21338 .iter()
21339 .any(|event| event.dimensions.contains_key("branch_status"))
21340 );
21341 }
21342
21343 #[tokio::test]
21344 async fn test_parallel_transition_low_cap_allows_deterministic_route() {
21345 let mock = mock_with_response("unused");
21346 let call_counter = mock.clone();
21347 let yaml = r#"
21348name: ParallelTransitionLowCapAgent
21349system_prompt: "Route before stale responses when safe."
21350runtime:
21351 optimization:
21352 enabled: true
21353 max_speculative_llm_calls_per_turn: 1
21354 speculative_state_transitions: true
21355 max_parallel_runtime_tasks: 2
21356states:
21357 initial: triage
21358 states:
21359 triage:
21360 prompt: "Triage state."
21361 transitions:
21362 - to: billing
21363 guard:
21364 context:
21365 route:
21366 eq: billing
21367 timing: parallel
21368 billing:
21369 prompt: "Billing state."
21370"#;
21371 let agent = AgentBuilder::from_yaml(yaml)
21372 .unwrap()
21373 .llm(Arc::new(mock))
21374 .build()
21375 .unwrap();
21376 agent
21377 .set_context("route", serde_json::json!("billing"))
21378 .unwrap();
21379 agent.update_active_turn_context("billing help", HashMap::new());
21380 assert!(
21381 agent.reserve_active_speculative_llm_call(
21382 RuntimeOptimizationKind::ParallelStateTransition
21383 )
21384 );
21385
21386 let selection = agent
21387 .select_parallel_transition_candidate("billing help")
21388 .await
21389 .unwrap();
21390 agent.end_root_turn();
21391
21392 match selection {
21393 ParallelTransitionSelection::Candidate(candidate) => {
21394 assert_eq!(candidate.target(), "billing");
21395 }
21396 ParallelTransitionSelection::NoMatch => panic!("deterministic route did not match"),
21397 ParallelTransitionSelection::ReservationExhausted => {
21398 panic!("deterministic route consumed LLM budget")
21399 }
21400 }
21401 assert_eq!(call_counter.call_count(), 0);
21402 }
21403
21404 #[tokio::test]
21405 async fn speculative_transition_drops_loser_before_state_actions() {
21406 let lock = Arc::new(tokio::sync::Mutex::new(()));
21407 let first_started = Arc::new(tokio::sync::Notify::new());
21408 let first_dropped = Arc::new(AtomicBool::new(false));
21409 let committed_after_drop = Arc::new(AtomicBool::new(false));
21410 let default = Arc::new(FirstCallLockingProvider {
21411 lock,
21412 first_started: Arc::clone(&first_started),
21413 first_dropped: Arc::clone(&first_dropped),
21414 committed_after_drop: Arc::clone(&committed_after_drop),
21415 calls: AtomicU64::new(0),
21416 });
21417 let router = Arc::new(RoutingAfterProviderStart {
21418 provider_started: first_started,
21419 });
21420 let yaml = r#"
21421name: SpeculativeCancellationAgent
21422system_prompt: "Route before committed work."
21423llm:
21424 default: default
21425 router: router
21426runtime:
21427 optimization:
21428 enabled: true
21429 max_speculative_llm_calls_per_turn: 2
21430 speculative_state_transitions: true
21431 max_parallel_runtime_tasks: 2
21432states:
21433 initial: triage
21434 states:
21435 triage:
21436 prompt: "Triage state."
21437 transitions:
21438 - to: technical
21439 when: "The request needs technical support"
21440 timing: parallel
21441 technical:
21442 prompt: "Technical state."
21443 on_enter:
21444 - prompt: "Prepare technical context."
21445 llm: default
21446 store_as: preparation
21447"#;
21448 let agent = AgentBuilder::from_yaml(yaml)
21449 .unwrap()
21450 .llm_alias("default", default)
21451 .llm_alias("router", router)
21452 .build()
21453 .unwrap();
21454
21455 let response = tokio::time::timeout(
21456 std::time::Duration::from_secs(2),
21457 agent.chat("I cannot log in because of AUTH-17."),
21458 )
21459 .await
21460 .expect("committed work must not wait on the losing provider future")
21461 .unwrap();
21462
21463 assert_eq!(response.content, "Committed technical response.");
21464 assert_eq!(agent.current_state().as_deref(), Some("technical"));
21465 assert!(first_dropped.load(Ordering::SeqCst));
21466 assert!(committed_after_drop.load(Ordering::SeqCst));
21467 }
21468
21469 #[tokio::test]
21470 async fn buffered_transition_drops_stale_stream_before_redispatch() {
21471 use futures::StreamExt;
21472
21473 let lock = Arc::new(tokio::sync::Mutex::new(()));
21474 let stream_started = Arc::new(tokio::sync::Notify::new());
21475 let stream_dropped = Arc::new(AtomicBool::new(false));
21476 let committed_after_drop = Arc::new(AtomicBool::new(false));
21477 let default = Arc::new(BufferedLockingProvider {
21478 lock,
21479 stream_started: Arc::clone(&stream_started),
21480 stream_dropped: Arc::clone(&stream_dropped),
21481 committed_after_drop: Arc::clone(&committed_after_drop),
21482 });
21483 let router = Arc::new(RoutingAfterProviderStart {
21484 provider_started: stream_started,
21485 });
21486 let yaml = r#"
21487name: BufferedCancellationAgent
21488system_prompt: "Hide stale streamed output."
21489llm:
21490 default: default
21491 router: router
21492streaming:
21493 enabled: true
21494 buffer_size: 8
21495runtime:
21496 optimization:
21497 enabled: true
21498 max_speculative_llm_calls_per_turn: 2
21499 speculative_state_transitions: true
21500 streaming_policy: buffer_until_routing_done
21501 max_parallel_runtime_tasks: 2
21502states:
21503 initial: triage
21504 states:
21505 triage:
21506 prompt: "Triage state."
21507 transitions:
21508 - to: technical
21509 when: "The request needs technical support"
21510 timing: parallel
21511 technical:
21512 prompt: "Technical state."
21513"#;
21514 let agent = AgentBuilder::from_yaml(yaml)
21515 .unwrap()
21516 .llm_alias("default", default)
21517 .llm_alias("router", router)
21518 .build()
21519 .unwrap();
21520
21521 let content = tokio::time::timeout(std::time::Duration::from_secs(2), async {
21522 let mut stream = agent
21523 .chat_stream("AUTH-17 needs technical help.")
21524 .await
21525 .unwrap();
21526 let mut content = String::new();
21527 while let Some(chunk) = stream.next().await {
21528 match chunk {
21529 StreamChunk::Content { text } => content.push_str(&text),
21530 StreamChunk::Done {} => break,
21531 StreamChunk::Error { message } => panic!("unexpected stream error: {message}"),
21532 _ => {}
21533 }
21534 }
21535 content
21536 })
21537 .await
21538 .expect("redispatch must not wait on the stale streaming future");
21539
21540 assert_eq!(content, "Committed technical response.");
21541 assert_eq!(agent.current_state().as_deref(), Some("technical"));
21542 assert!(stream_dropped.load(Ordering::SeqCst));
21543 assert!(committed_after_drop.load(Ordering::SeqCst));
21544 }
21545
21546 #[tokio::test]
21547 async fn buffered_transition_drops_established_stream_before_redispatch() {
21548 use futures::StreamExt;
21549
21550 let stream_started = Arc::new(tokio::sync::Notify::new());
21551 let stream_dropped = Arc::new(AtomicBool::new(false));
21552 let stream_dropped_notify = Arc::new(tokio::sync::Notify::new());
21553 let committed_after_drop = Arc::new(AtomicBool::new(false));
21554 let default = Arc::new(EstablishedStreamProvider {
21555 stream_started: Arc::clone(&stream_started),
21556 stream_dropped: Arc::clone(&stream_dropped),
21557 stream_dropped_notify,
21558 committed_after_drop: Arc::clone(&committed_after_drop),
21559 });
21560 let router = Arc::new(RoutingAfterProviderStart {
21561 provider_started: stream_started,
21562 });
21563 let yaml = r#"
21564name: EstablishedStreamCancellationAgent
21565system_prompt: "Hide stale streamed output."
21566llm:
21567 default: default
21568 router: router
21569streaming:
21570 enabled: true
21571 buffer_size: 8
21572runtime:
21573 optimization:
21574 enabled: true
21575 max_speculative_llm_calls_per_turn: 2
21576 speculative_state_transitions: true
21577 streaming_policy: buffer_until_routing_done
21578 max_parallel_runtime_tasks: 2
21579states:
21580 initial: triage
21581 states:
21582 triage:
21583 prompt: "Triage state."
21584 transitions:
21585 - to: technical
21586 when: "The request needs technical support"
21587 timing: parallel
21588 technical:
21589 prompt: "Technical state."
21590"#;
21591 let agent = AgentBuilder::from_yaml(yaml)
21592 .unwrap()
21593 .llm_alias("default", default)
21594 .llm_alias("router", router)
21595 .build()
21596 .unwrap();
21597
21598 let content = tokio::time::timeout(std::time::Duration::from_secs(2), async {
21599 let mut stream = agent
21600 .chat_stream("AUTH-17 needs technical help.")
21601 .await
21602 .unwrap();
21603 let mut content = String::new();
21604 while let Some(chunk) = stream.next().await {
21605 match chunk {
21606 StreamChunk::Content { text } => content.push_str(&text),
21607 StreamChunk::Done {} => break,
21608 StreamChunk::Error { message } => panic!("unexpected stream error: {message}"),
21609 _ => {}
21610 }
21611 }
21612 content
21613 })
21614 .await
21615 .expect("redispatch must wait for the established stale stream to be dropped");
21616
21617 assert_eq!(content, "Committed technical response.");
21618 assert_eq!(agent.current_state().as_deref(), Some("technical"));
21619 assert!(stream_dropped.load(Ordering::SeqCst));
21620 assert!(committed_after_drop.load(Ordering::SeqCst));
21621 }
21622
21623 #[tokio::test]
21624 async fn test_buffered_streaming_transition_reservation_falls_back() {
21625 use futures::StreamExt;
21626
21627 let mock = mock_with_responses(vec![
21628 "Serial streaming response",
21629 "Serial streaming response",
21630 ]);
21631 let router_mock = mock_with_response("1");
21632 let router_counter = router_mock.clone();
21633 let yaml = r#"
21634name: BufferedReservationFallbackAgent
21635system_prompt: "Stream normally if speculative routing cannot be evaluated."
21636llm:
21637 default: default
21638 router: router
21639observability:
21640 enabled: true
21641 export:
21642 write_raw_events: true
21643streaming:
21644 enabled: true
21645 buffer_size: 8
21646runtime:
21647 optimization:
21648 enabled: true
21649 max_speculative_llm_calls_per_turn: 1
21650 speculative_state_transitions: true
21651 streaming_policy: buffer_until_routing_done
21652 max_parallel_runtime_tasks: 2
21653states:
21654 initial: triage
21655 states:
21656 triage:
21657 prompt: "Triage state."
21658 transitions:
21659 - to: billing
21660 guard:
21661 context:
21662 route:
21663 eq: billing
21664 when: "User asks about billing"
21665 timing: parallel
21666 billing:
21667 prompt: "Billing state."
21668"#;
21669 let agent = AgentBuilder::from_yaml(yaml)
21670 .unwrap()
21671 .llm_alias("default", Arc::new(mock))
21672 .llm_alias("router", Arc::new(router_mock))
21673 .build()
21674 .unwrap();
21675
21676 let mut stream = agent.chat_stream("hello").await.unwrap();
21677 let mut content = String::new();
21678 let mut error = None;
21679 while let Some(chunk) = stream.next().await {
21680 match chunk {
21681 StreamChunk::Content { text } => content.push_str(&text),
21682 StreamChunk::Error { message } => error = Some(message),
21683 StreamChunk::Done {} => break,
21684 _ => {}
21685 }
21686 }
21687
21688 assert_eq!(error, None);
21689 assert_eq!(content, "Serial streaming response");
21690 assert_eq!(router_counter.call_count(), 0);
21691 let events = agent.observability().unwrap().raw_events();
21692 assert!(events.iter().any(|event| {
21693 event.dimensions.get("branch_status") == Some(&"cancelled".to_string())
21694 && event.dimensions.get("commit_behavior")
21695 == Some(&"transition_decision".to_string())
21696 }));
21697 }
21698
21699 #[tokio::test]
21700 async fn test_blocking_error_cleanup_resets_root_turn_for_next_chat() {
21701 let mut mock = mock_with_response("Recovered response");
21702 mock.set_error("boom");
21703 let mut handle = mock.clone();
21704 let agent = AgentBuilder::new()
21705 .system_prompt("You are helpful.")
21706 .llm(Arc::new(mock))
21707 .build()
21708 .unwrap();
21709
21710 assert!(agent.chat("first").await.is_err());
21711 handle.clear_error();
21712 let response = agent.chat("second").await.unwrap();
21713
21714 assert_eq!(response.content, "Recovered response");
21715 let messages = agent.memory.get_messages(None).await.unwrap();
21716 let user_count = messages
21717 .iter()
21718 .filter(|message| message.role == ai_agents_core::Role::User)
21719 .count();
21720 assert_eq!(user_count, 2);
21721 }
21722
21723 #[tokio::test]
21724 async fn test_streaming_error_cleanup_resets_root_turn_for_next_chat() {
21725 use futures::StreamExt;
21726
21727 let mut mock = mock_with_response("Recovered response");
21728 mock.set_error("stream boom");
21729 let mut handle = mock.clone();
21730 let agent = AgentBuilder::new()
21731 .system_prompt("You are helpful.")
21732 .llm(Arc::new(mock))
21733 .build()
21734 .unwrap();
21735
21736 let mut stream = agent.chat_stream("first").await.unwrap();
21737 let mut saw_error = false;
21738 while let Some(chunk) = stream.next().await {
21739 if matches!(chunk, StreamChunk::Error { .. }) {
21740 saw_error = true;
21741 }
21742 }
21743 assert!(saw_error);
21744
21745 handle.clear_error();
21746 let response = agent.chat("second").await.unwrap();
21747
21748 assert_eq!(response.content, "Recovered response");
21749 let messages = agent.memory.get_messages(None).await.unwrap();
21750 let user_count = messages
21751 .iter()
21752 .filter(|message| message.role == ai_agents_core::Role::User)
21753 .count();
21754 assert_eq!(user_count, 2);
21755 }
21756
21757 #[tokio::test]
21758 async fn test_buffered_streaming_route_miss_releases_buffer_limit() {
21759 use futures::StreamExt;
21760
21761 let mut mock = mock_with_response("one two three");
21762 mock.set_latency(10);
21763 let yaml = r#"
21764name: BufferedMissAgent
21765system_prompt: "You stream safely."
21766llm:
21767 default: default
21768streaming:
21769 enabled: true
21770 buffer_size: 1
21771runtime:
21772 optimization:
21773 enabled: true
21774 max_speculative_llm_calls_per_turn: 2
21775 speculative_state_transitions: true
21776 streaming_policy: buffer_until_routing_done
21777 max_parallel_runtime_tasks: 2
21778states:
21779 initial: triage
21780 states:
21781 triage:
21782 prompt: "Answer from triage."
21783 transitions:
21784 - to: billing
21785 guard:
21786 context:
21787 route:
21788 eq: billing
21789 timing: parallel
21790 billing:
21791 prompt: "Billing state."
21792"#;
21793 let agent = AgentBuilder::from_yaml(yaml)
21794 .unwrap()
21795 .llm_alias("default", Arc::new(mock))
21796 .build()
21797 .unwrap();
21798
21799 let mut stream = agent.chat_stream("hello").await.unwrap();
21800 let mut content = String::new();
21801 let mut error = None;
21802 while let Some(chunk) = stream.next().await {
21803 match chunk {
21804 StreamChunk::Content { text } => content.push_str(&text),
21805 StreamChunk::Error { message } => error = Some(message),
21806 StreamChunk::Done {} => break,
21807 _ => {}
21808 }
21809 }
21810
21811 assert_eq!(error, None);
21812 assert_eq!(content, "one two three");
21813 }
21814
21815 #[tokio::test]
21816 async fn test_buffered_streaming_main_failure_finalizes_branch() {
21817 use futures::StreamExt;
21818
21819 let mock = mock_with_response("one two");
21820 let mut router_mock = mock_with_response("0");
21821 router_mock.set_latency(50);
21822 let yaml = r#"
21823name: BufferedFailureAgent
21824system_prompt: "You stream safely."
21825llm:
21826 default: default
21827 router: router
21828observability:
21829 enabled: true
21830 export:
21831 write_raw_events: true
21832streaming:
21833 enabled: true
21834 buffer_size: 1
21835runtime:
21836 optimization:
21837 enabled: true
21838 max_speculative_llm_calls_per_turn: 2
21839 speculative_state_transitions: true
21840 streaming_policy: buffer_until_routing_done
21841 max_parallel_runtime_tasks: 2
21842states:
21843 initial: triage
21844 states:
21845 triage:
21846 prompt: "Ask for the category."
21847 transitions:
21848 - to: billing
21849 when: "User asks about billing"
21850 timing: parallel
21851 billing:
21852 prompt: "Billing state."
21853"#;
21854 let agent = AgentBuilder::from_yaml(yaml)
21855 .unwrap()
21856 .llm_alias("default", Arc::new(mock))
21857 .llm_alias("router", Arc::new(router_mock))
21858 .build()
21859 .unwrap();
21860
21861 let mut stream = agent.chat_stream("hello").await.unwrap();
21862 let mut error = String::new();
21863 while let Some(chunk) = stream.next().await {
21864 if let StreamChunk::Error { message } = chunk {
21865 error = message;
21866 }
21867 }
21868
21869 assert!(
21870 error.contains("stream buffer filled"),
21871 "unexpected stream error: {}",
21872 error
21873 );
21874 let events = agent.observability().unwrap().raw_events();
21875 assert!(events.iter().any(|event| {
21876 event.dimensions.get("branch_status") == Some(&"failed".to_string())
21877 && event.dimensions.get("commit_behavior") == Some(&"final_response".to_string())
21878 && event.dimensions.get("optimization")
21879 == Some(&"buffered_streaming_routing".to_string())
21880 }));
21881 }
21882
21883 #[tokio::test]
21884 async fn test_streaming_preflight_does_not_emit_old_state_content() {
21885 use futures::StreamExt;
21886
21887 let mock = mock_with_response("Billing streamed response");
21888 let yaml = r#"
21889name: StreamingOptimizedAgent
21890system_prompt: "You route before streaming."
21891runtime:
21892 optimization:
21893 enabled: true
21894 pre_response_deterministic_transitions: true
21895streaming:
21896 enabled: true
21897states:
21898 initial: greeting
21899 states:
21900 greeting:
21901 prompt: "OLD_STATE_SENTINEL"
21902 transitions:
21903 - to: billing
21904 guard:
21905 context:
21906 topic:
21907 eq: billing
21908 timing: pre_response
21909 billing:
21910 prompt: "Billing state."
21911"#;
21912 let agent = AgentBuilder::from_yaml(yaml)
21913 .unwrap()
21914 .llm(Arc::new(mock))
21915 .build()
21916 .unwrap();
21917 agent
21918 .set_context("topic", serde_json::json!("billing"))
21919 .unwrap();
21920
21921 let mut stream = agent.chat_stream("billing please").await.unwrap();
21922 let mut content = String::new();
21923 while let Some(chunk) = stream.next().await {
21924 match chunk {
21925 StreamChunk::Content { text } => content.push_str(&text),
21926 StreamChunk::Error { message } => panic!("stream error: {}", message),
21927 StreamChunk::Done {} => break,
21928 _ => {}
21929 }
21930 }
21931
21932 assert_eq!(agent.current_state().as_deref(), Some("billing"));
21933 assert!(content.contains("Billing streamed response"));
21934 assert!(!content.contains("OLD_STATE_SENTINEL"));
21935 }
21936
21937 #[tokio::test]
21939 async fn test_integration_state_machine_basic() {
21940 let yaml = r#"
21941name: StateAgent
21942system_prompt: "You are a support agent."
21943states:
21944 initial: greeting
21945 states:
21946 greeting:
21947 prompt: "Welcome the user warmly."
21948 transitions:
21949 - to: support
21950 when: "User needs help"
21951 auto: true
21952 support:
21953 prompt: "Help solve the user's problem."
21954"#;
21955 let mock = mock_with_responses(vec![
21956 "Welcome! How can I help?", "1", "I'll help you with that.", ]);
21960 let builder = AgentBuilder::from_yaml(yaml).unwrap();
21961 let agent = builder.llm(Arc::new(mock)).build().unwrap();
21962
21963 assert_eq!(agent.current_state(), Some("greeting".to_string()));
21964 let _ = agent.chat("I need help").await.unwrap();
21965 }
21968
21969 #[tokio::test]
21971 async fn test_integration_state_on_enter_set_context() {
21972 let yaml = r#"
21973name: ActionAgent
21974system_prompt: "You are helpful."
21975states:
21976 initial: step1
21977 states:
21978 step1:
21979 prompt: "Step 1"
21980 on_exit:
21981 - set_context:
21982 step1_exited: true
21983 transitions:
21984 - to: step2
21985 when: "always"
21986 auto: true
21987 step2:
21988 prompt: "Step 2"
21989 on_enter:
21990 - set_context:
21991 step2_entered: true
21992"#;
21993 let mock = mock_with_responses(vec![
21995 "Processing step 1.",
21996 "0", ]);
21998 let builder = AgentBuilder::from_yaml(yaml).unwrap();
21999 let agent = builder.llm(Arc::new(mock)).build().unwrap();
22000
22001 assert_eq!(agent.current_state(), Some("step1".to_string()));
22002
22003 agent.transition_to("step2").await.unwrap();
22005
22006 assert_eq!(agent.current_state(), Some("step2".to_string()));
22007
22008 let ctx = agent.get_context();
22010 assert_eq!(ctx.get("step1_exited"), Some(&serde_json::json!(true)));
22011 assert_eq!(ctx.get("step2_entered"), Some(&serde_json::json!(true)));
22012 }
22013
22014 #[tokio::test]
22015 async fn state_action_tool_preserves_source_in_stored_record() {
22016 let yaml = r#"
22017name: StateActionToolAgent
22018system_prompt: "You are helpful."
22019tools:
22020 - context_echo
22021states:
22022 initial: idle
22023 states:
22024 idle:
22025 prompt: "Idle"
22026 active:
22027 prompt: "Active"
22028 on_enter:
22029 - set_context:
22030 action_started: true
22031 - tool: context_echo
22032 args: {}
22033"#;
22034 let agent = AgentBuilder::from_yaml(yaml)
22035 .unwrap()
22036 .llm(Arc::new(mock_with_response("unused")))
22037 .tool(Arc::new(ContextEchoTool))
22038 .build()
22039 .unwrap();
22040
22041 agent.transition_to("active").await.unwrap();
22042
22043 let record: ToolExecutionRecord = serde_json::from_value(
22044 agent
22045 .get_context()
22046 .get("last_tool_record")
22047 .cloned()
22048 .expect("successful state action must store its execution record"),
22049 )
22050 .unwrap();
22051 assert!(record.executed);
22052 assert!(record.success);
22053 assert_eq!(record.canonical_id, "context_echo");
22054 assert!(matches!(
22055 &record.source,
22056 ToolCallSource::StateAction {
22057 state: Some(state),
22058 action_index: 1,
22059 } if state == "active"
22060 ));
22061 }
22062
22063 #[tokio::test]
22064 async fn test_ordinary_transition_uses_on_enter_then_on_reenter() {
22065 let yaml = r#"
22066name: OrdinaryLifecycleAgent
22067system_prompt: "You are helpful."
22068states:
22069 initial: intake
22070 regenerate_on_transition: false
22071 states:
22072 intake:
22073 prompt: "Intake"
22074 transitions:
22075 - to: drafting
22076 guard:
22077 context:
22078 route:
22079 eq: drafting
22080 drafting:
22081 prompt: "Drafting"
22082 on_enter:
22083 - set_context:
22084 draft_version: 1
22085 on_reenter:
22086 - set_context:
22087 draft_version: 2
22088 transitions:
22089 - to: review
22090 guard:
22091 context:
22092 route:
22093 eq: review
22094 review:
22095 prompt: "Review"
22096 on_enter:
22097 - set_context:
22098 review_entry: first
22099 transitions:
22100 - to: drafting
22101 guard:
22102 context:
22103 route:
22104 eq: drafting
22105"#;
22106 let agent = AgentBuilder::from_yaml(yaml)
22107 .unwrap()
22108 .llm(Arc::new(mock_with_responses(vec![
22109 "Intake response",
22110 "Draft response",
22111 "Review response",
22112 ])))
22113 .build()
22114 .unwrap();
22115
22116 agent
22117 .set_context("route", serde_json::json!("drafting"))
22118 .unwrap();
22119 agent.chat("Start a draft").await.unwrap();
22120 assert_eq!(agent.current_state().as_deref(), Some("drafting"));
22121 assert_eq!(
22122 agent.get_context().get("draft_version"),
22123 Some(&serde_json::json!(1))
22124 );
22125
22126 agent
22127 .set_context("route", serde_json::json!("review"))
22128 .unwrap();
22129 agent.chat("Review this").await.unwrap();
22130 assert_eq!(agent.current_state().as_deref(), Some("review"));
22131 assert_eq!(
22132 agent.get_context().get("review_entry"),
22133 Some(&serde_json::json!("first"))
22134 );
22135
22136 agent
22137 .set_context("route", serde_json::json!("drafting"))
22138 .unwrap();
22139 agent.chat("Revise this").await.unwrap();
22140 assert_eq!(agent.current_state().as_deref(), Some("drafting"));
22141 assert_eq!(
22142 agent.get_context().get("draft_version"),
22143 Some(&serde_json::json!(2))
22144 );
22145 }
22146
22147 #[tokio::test]
22148 async fn test_manual_transition_uses_on_enter_then_on_reenter() {
22149 let yaml = r#"
22150name: ManualLifecycleAgent
22151system_prompt: "You are helpful."
22152states:
22153 initial: intake
22154 states:
22155 intake:
22156 prompt: "Intake"
22157 drafting:
22158 prompt: "Drafting"
22159 on_enter:
22160 - set_context:
22161 draft_version: 1
22162 on_reenter:
22163 - set_context:
22164 draft_version: 2
22165 review:
22166 prompt: "Review"
22167"#;
22168 let agent = AgentBuilder::from_yaml(yaml)
22169 .unwrap()
22170 .llm(Arc::new(mock_with_response("unused")))
22171 .build()
22172 .unwrap();
22173
22174 assert!(!agent.get_context().contains_key("draft_version"));
22175 agent.transition_to("drafting").await.unwrap();
22176 assert_eq!(agent.current_state().as_deref(), Some("drafting"));
22177 assert_eq!(
22178 agent.get_context().get("draft_version"),
22179 Some(&serde_json::json!(1))
22180 );
22181
22182 agent.transition_to("review").await.unwrap();
22183 agent.transition_to("drafting").await.unwrap();
22184 assert_eq!(agent.current_state().as_deref(), Some("drafting"));
22185 assert_eq!(
22186 agent.get_context().get("draft_version"),
22187 Some(&serde_json::json!(2))
22188 );
22189 }
22190
22191 #[tokio::test]
22192 async fn test_timeout_transition_uses_on_enter_then_on_reenter() {
22193 let yaml = r#"
22194name: TimeoutLifecycleAgent
22195system_prompt: "You are helpful."
22196states:
22197 initial: intake
22198 regenerate_on_transition: false
22199 states:
22200 intake:
22201 prompt: "Intake"
22202 max_turns: 1
22203 timeout_to: drafting
22204 drafting:
22205 prompt: "Drafting"
22206 max_turns: 1
22207 timeout_to: review
22208 on_enter:
22209 - set_context:
22210 draft_version: 1
22211 on_reenter:
22212 - set_context:
22213 draft_version: 2
22214 review:
22215 prompt: "Review"
22216 max_turns: 1
22217 timeout_to: drafting
22218 on_enter:
22219 - set_context:
22220 review_entry: first
22221"#;
22222 let agent = AgentBuilder::from_yaml(yaml)
22223 .unwrap()
22224 .llm(Arc::new(mock_with_responses(vec![
22225 "Intake",
22226 "First draft",
22227 "Review",
22228 "Revised draft",
22229 ])))
22230 .build()
22231 .unwrap();
22232
22233 agent.chat("First turn").await.unwrap();
22234 assert_eq!(agent.current_state().as_deref(), Some("intake"));
22235 assert!(!agent.get_context().contains_key("draft_version"));
22236
22237 agent.chat("Second turn").await.unwrap();
22238 assert_eq!(agent.current_state().as_deref(), Some("drafting"));
22239 assert_eq!(
22240 agent.get_context().get("draft_version"),
22241 Some(&serde_json::json!(1))
22242 );
22243
22244 agent.chat("Third turn").await.unwrap();
22245 assert_eq!(agent.current_state().as_deref(), Some("review"));
22246 assert_eq!(
22247 agent.get_context().get("review_entry"),
22248 Some(&serde_json::json!("first"))
22249 );
22250
22251 agent.chat("Fourth turn").await.unwrap();
22252 assert_eq!(agent.current_state().as_deref(), Some("drafting"));
22253 assert_eq!(
22254 agent.get_context().get("draft_version"),
22255 Some(&serde_json::json!(2))
22256 );
22257 }
22258
22259 #[tokio::test]
22261 async fn test_integration_process_normalize() {
22262 let yaml = r#"
22263name: ProcessAgent
22264system_prompt: "You are helpful."
22265process:
22266 input:
22267 - type: normalize
22268 config:
22269 trim: true
22270 collapse_whitespace: true
22271"#;
22272 let mock = mock_with_response("Got your message.");
22273 let builder = AgentBuilder::from_yaml(yaml).unwrap();
22274 let agent = builder.llm(Arc::new(mock.clone())).build().unwrap();
22275
22276 let _ = agent.chat(" hello world ").await.unwrap();
22277
22278 let history = mock.call_history();
22280 assert!(!history.is_empty());
22281 let last_call = history.last().unwrap();
22283 let user_msg = last_call
22284 .messages
22285 .iter()
22286 .find(|m| m.role == ai_agents_core::Role::User)
22287 .unwrap();
22288 assert_eq!(user_msg.content, "hello world");
22289 }
22290
22291 #[tokio::test]
22295 async fn test_integration_memory_compression() {
22296 let yaml = r#"
22297name: MemoryAgent
22298system_prompt: "You are helpful."
22299memory:
22300 type: compacting
22301 max_messages: 100
22302 compress_threshold: 5
22303 max_recent_messages: 3
22304 summarize_batch_size: 2
22305"#;
22306 let responses: Vec<&str> = (0..8).map(|_| "Response from assistant.").collect();
22308 let mock = mock_with_responses(responses);
22309 let builder = AgentBuilder::from_yaml(yaml).unwrap();
22310 let agent = builder.llm(Arc::new(mock)).build().unwrap();
22311
22312 for i in 0..6 {
22314 let _ = agent.chat(&format!("Message {}", i)).await.unwrap();
22315 }
22316
22317 let messages = agent.memory.get_messages(None).await.unwrap();
22320 assert!(messages.len() <= 12); }
22324
22325 #[tokio::test]
22327 async fn test_integration_multi_llm_registry() {
22328 let mut mock_default = MockLLMProvider::new("default");
22329 mock_default.set_response("Default LLM response.");
22330 let mut mock_router = MockLLMProvider::new("router");
22331 mock_router.set_response("Router response.");
22332
22333 let agent = AgentBuilder::new()
22334 .system_prompt("You are helpful.")
22335 .llm_alias("default", Arc::new(mock_default))
22336 .llm_alias("router", Arc::new(mock_router))
22337 .build()
22338 .unwrap();
22339
22340 let response = agent.chat("Hello").await.unwrap();
22341 assert_eq!(response.content, "Default LLM response.");
22342 }
22343
22344 #[tokio::test]
22346 async fn test_integration_agent_reset() {
22347 let mock = mock_with_responses(vec!["Hello!", "Hello again!"]);
22348 let agent = AgentBuilder::new()
22349 .system_prompt("You are helpful.")
22350 .llm(Arc::new(mock))
22351 .build()
22352 .unwrap();
22353
22354 let _ = agent.chat("Hi").await.unwrap();
22355 let messages = agent.memory.get_messages(None).await.unwrap();
22356 assert_eq!(messages.len(), 2); agent.reset().await.unwrap();
22359 let messages = agent.memory.get_messages(None).await.unwrap();
22360 assert_eq!(messages.len(), 0);
22361 }
22362
22363 #[tokio::test]
22365 async fn test_integration_process_validate_reject() {
22366 use ai_agents_process::{ProcessConfig, ProcessProcessor};
22367
22368 let validate_config = ai_agents_process::ValidateStage {
22369 id: Some("length_check".to_string()),
22370 condition: None,
22371 config: ai_agents_process::ValidateConfig {
22372 rules: vec![ai_agents_process::ValidationRule::MinLength {
22373 min_length: 10,
22374 on_fail: ai_agents_process::ValidationAction {
22375 action: ai_agents_process::ValidationActionType::Reject,
22376 message: None,
22377 },
22378 }],
22379 ..Default::default()
22380 },
22381 };
22382 let process_config = ProcessConfig {
22383 input: vec![ai_agents_process::ProcessStage::Validate(validate_config)],
22384 ..Default::default()
22385 };
22386 let processor = ProcessProcessor::new(process_config);
22387
22388 let mock = mock_with_response("Should not reach here.");
22389 let agent = AgentBuilder::new()
22390 .system_prompt("You are helpful.")
22391 .llm(Arc::new(mock))
22392 .process_processor(processor)
22393 .build()
22394 .unwrap();
22395
22396 let response = agent.chat("Hi").await.unwrap();
22397 assert!(
22399 response.content.contains("rejected")
22400 || response.content.contains("Input rejected")
22401 || response.content.contains("too short")
22402 || response.content.contains("Too short")
22403 || response.content.len() < 50, "Expected rejection response, got: {}",
22405 response.content
22406 );
22407 }
22408
22409 #[tokio::test]
22411 async fn test_llm_fallback_on_failure() {
22412 use ai_agents_recovery::{ErrorRecoveryConfig, LLMFailureAction, LLMRecoveryConfig};
22413
22414 let mut primary = MockLLMProvider::new("primary");
22415 primary.set_error("Primary LLM is unavailable");
22416
22417 let mut fallback = MockLLMProvider::new("fallback");
22418 fallback.set_response("Fallback response works!");
22419
22420 let agent = AgentBuilder::new()
22421 .system_prompt("You are helpful.")
22422 .llm_alias("default", Arc::new(primary))
22423 .llm_alias("backup", Arc::new(fallback))
22424 .recovery_manager(RecoveryManager::new(ErrorRecoveryConfig {
22425 llm: LLMRecoveryConfig {
22426 on_failure: LLMFailureAction::FallbackLlm {
22427 fallback_llm: "backup".to_string(),
22428 },
22429 ..Default::default()
22430 },
22431 ..Default::default()
22432 }))
22433 .build()
22434 .unwrap();
22435
22436 let response = agent.chat("Hello").await.unwrap();
22437 assert!(
22438 response.content.contains("Fallback response"),
22439 "Expected fallback response, got: {}",
22440 response.content
22441 );
22442 }
22443
22444 #[tokio::test]
22446 async fn test_llm_fallback_response_static_message() {
22447 use ai_agents_recovery::{ErrorRecoveryConfig, LLMFailureAction, LLMRecoveryConfig};
22448
22449 let mut primary = MockLLMProvider::new("primary");
22450 primary.set_error("Primary LLM is unavailable");
22451
22452 let agent = AgentBuilder::new()
22453 .system_prompt("You are helpful.")
22454 .llm(Arc::new(primary))
22455 .recovery_manager(RecoveryManager::new(ErrorRecoveryConfig {
22456 llm: LLMRecoveryConfig {
22457 on_failure: LLMFailureAction::FallbackResponse {
22458 message: "I am temporarily unavailable. Please try again later."
22459 .to_string(),
22460 },
22461 ..Default::default()
22462 },
22463 ..Default::default()
22464 }))
22465 .build()
22466 .unwrap();
22467
22468 let response = agent.chat("Hello").await.unwrap();
22469 assert!(
22470 response.content.contains("temporarily unavailable"),
22471 "Expected static fallback message, got: {}",
22472 response.content
22473 );
22474 }
22475
22476 #[tokio::test]
22479 async fn test_tool_failure_skip() {
22480 use ai_agents_recovery::{
22481 ErrorRecoveryConfig, ToolFailureAction, ToolRecoveryConfig, ToolRetryConfig,
22482 };
22483
22484 let mock = mock_with_responses(vec![
22485 r#"{"tool": "calculator", "arguments": {"expression": "not a number +"}}"#,
22486 "The calculation was skipped, but I can still help you.",
22487 ]);
22488 let observed = mock.clone();
22489 let mut tools = ai_agents_tools::ToolRegistry::new();
22490 tools
22491 .register(Arc::new(ai_agents_tools::CalculatorTool))
22492 .unwrap();
22493
22494 let agent = AgentBuilder::new()
22495 .system_prompt("You are helpful.")
22496 .llm(Arc::new(mock))
22497 .tools(tools)
22498 .recovery_manager(RecoveryManager::new(ErrorRecoveryConfig {
22499 tools: ToolRecoveryConfig {
22500 default: ToolRetryConfig {
22501 max_retries: 0,
22502 timeout_ms: None,
22503 on_failure: ToolFailureAction::Skip,
22504 },
22505 ..Default::default()
22506 },
22507 ..Default::default()
22508 }))
22509 .build()
22510 .unwrap();
22511
22512 let response = agent.chat("Compute this").await.unwrap();
22513
22514 assert_eq!(
22515 response.content,
22516 "The calculation was skipped, but I can still help you."
22517 );
22518 assert_eq!(observed.call_count(), 2);
22519 let history = agent.tool_call_history();
22521 assert_eq!(history.len(), 1);
22522 assert_eq!(history[0].tool_id, "calculator");
22523 assert_eq!(
22524 history[0].result.get("skipped"),
22525 Some(&serde_json::json!(true)),
22526 "{:?}",
22527 history[0].result
22528 );
22529 }
22530
22531 #[tokio::test]
22533 async fn test_unregistered_tool_call_records_unavailable_and_continues() {
22534 let mock = mock_with_responses(vec![
22535 r#"{"tool": "nonexistent_tool", "arguments": {}}"#,
22536 "The tool was unavailable, but I can still help you.",
22537 ]);
22538 let observed = mock.clone();
22539
22540 let agent = AgentBuilder::new()
22541 .system_prompt("You are helpful.")
22542 .llm(Arc::new(mock))
22543 .build()
22544 .unwrap();
22545
22546 let response = agent.chat("Use the nonexistent tool").await.unwrap();
22547
22548 assert_eq!(
22549 response.content,
22550 "The tool was unavailable, but I can still help you."
22551 );
22552 assert_eq!(observed.call_count(), 2);
22553 let history = agent.tool_call_history();
22554 assert_eq!(history.len(), 1);
22555 assert_eq!(history[0].tool_id, "nonexistent_tool");
22556 assert_eq!(
22557 history[0].result.pointer("/error/kind"),
22558 Some(&serde_json::json!("tool_unavailable")),
22559 "{:?}",
22560 history[0].result
22561 );
22562 }
22563
22564 fn fallback_llm_recovery(fallback_llm: &str) -> RecoveryManager {
22569 use ai_agents_recovery::{ErrorRecoveryConfig, LLMFailureAction, LLMRecoveryConfig};
22570 RecoveryManager::new(ErrorRecoveryConfig {
22571 llm: LLMRecoveryConfig {
22572 on_failure: LLMFailureAction::FallbackLlm {
22573 fallback_llm: fallback_llm.to_string(),
22574 },
22575 ..Default::default()
22576 },
22577 ..Default::default()
22578 })
22579 }
22580
22581 #[tokio::test]
22582 async fn test_stream_llm_fallback_on_open_failure() {
22583 let mut primary = MockLLMProvider::new("primary");
22584 primary.set_error("Primary LLM is unavailable");
22585 let mut fallback = MockLLMProvider::new("fallback");
22586 fallback.set_response("Fallback response works!");
22587
22588 let agent = AgentBuilder::new()
22589 .system_prompt("You are helpful.")
22590 .llm_alias("default", Arc::new(primary))
22591 .llm_alias("backup", Arc::new(fallback))
22592 .recovery_manager(fallback_llm_recovery("backup"))
22593 .build()
22594 .unwrap();
22595
22596 let (content, chunks, final_response) = collect_stream_events(&agent, "Hello").await;
22597 assert!(
22598 !chunks.iter().any(StreamChunk::is_error),
22599 "fallback must not surface as a stream error: {chunks:?}"
22600 );
22601 let final_response = final_response.expect("Final must be emitted after fallback");
22602 assert!(content.contains("Fallback response"));
22603 assert!(final_response.content.contains("Fallback response"));
22604 }
22605
22606 #[tokio::test]
22607 async fn test_stream_llm_fallback_response_static_message() {
22608 use ai_agents_recovery::{ErrorRecoveryConfig, LLMFailureAction, LLMRecoveryConfig};
22609
22610 let mut primary = MockLLMProvider::new("primary");
22611 primary.set_error("Primary LLM is unavailable");
22612
22613 let agent = AgentBuilder::new()
22614 .system_prompt("You are helpful.")
22615 .llm(Arc::new(primary))
22616 .recovery_manager(RecoveryManager::new(ErrorRecoveryConfig {
22617 llm: LLMRecoveryConfig {
22618 on_failure: LLMFailureAction::FallbackResponse {
22619 message: "Service is temporarily unavailable.".to_string(),
22620 },
22621 ..Default::default()
22622 },
22623 ..Default::default()
22624 }))
22625 .build()
22626 .unwrap();
22627
22628 let (content, chunks, final_response) = collect_stream_events(&agent, "Hello").await;
22629 assert!(!chunks.iter().any(StreamChunk::is_error));
22630 let content_chunks = chunks.iter().filter(|c| c.is_content()).count();
22631 assert_eq!(content_chunks, 1, "static fallback is one content chunk");
22632 assert_eq!(content, "Service is temporarily unavailable.");
22633 assert_eq!(
22634 final_response.expect("Final").content,
22635 "Service is temporarily unavailable."
22636 );
22637 }
22638
22639 struct FailOnceStreamProvider {
22641 remaining_failures: Arc<std::sync::atomic::AtomicUsize>,
22642 open_attempts: Arc<std::sync::atomic::AtomicUsize>,
22643 }
22644
22645 #[async_trait]
22646 impl LLMProvider for FailOnceStreamProvider {
22647 async fn complete(
22648 &self,
22649 _messages: &[ChatMessage],
22650 _config: Option<&LLMConfig>,
22651 ) -> std::result::Result<LLMResponse, LLMError> {
22652 Ok(LLMResponse::new("blocking path", FinishReason::Stop))
22653 }
22654
22655 async fn complete_stream(
22656 &self,
22657 _messages: &[ChatMessage],
22658 _config: Option<&LLMConfig>,
22659 ) -> std::result::Result<
22660 Box<dyn Stream<Item = std::result::Result<LLMChunk, LLMError>> + Unpin + Send>,
22661 LLMError,
22662 > {
22663 self.open_attempts.fetch_add(1, Ordering::SeqCst);
22664 if self
22665 .remaining_failures
22666 .fetch_update(Ordering::SeqCst, Ordering::SeqCst, |n| n.checked_sub(1))
22667 .is_ok()
22668 {
22669 return Err(LLMError::Network("connection reset".to_string()));
22670 }
22671 Ok(Box::new(futures::stream::iter(vec![Ok(LLMChunk::new(
22672 "Recovered after retry",
22673 true,
22674 ))])))
22675 }
22676
22677 fn provider_name(&self) -> &str {
22678 "fail-once-stream"
22679 }
22680
22681 fn supports(&self, feature: LLMFeature) -> bool {
22682 matches!(feature, LLMFeature::Streaming)
22683 }
22684 }
22685
22686 #[tokio::test]
22687 async fn test_stream_llm_retry_then_success() {
22688 use ai_agents_recovery::{BackoffConfig, ErrorRecoveryConfig, RetryConfig};
22689
22690 let open_attempts = Arc::new(std::sync::atomic::AtomicUsize::new(0));
22691 let provider = FailOnceStreamProvider {
22692 remaining_failures: Arc::new(std::sync::atomic::AtomicUsize::new(1)),
22693 open_attempts: Arc::clone(&open_attempts),
22694 };
22695
22696 let agent = AgentBuilder::new()
22697 .system_prompt("You are helpful.")
22698 .llm(Arc::new(provider))
22699 .recovery_manager(RecoveryManager::new(ErrorRecoveryConfig {
22700 default: RetryConfig {
22701 max_retries: 1,
22702 backoff: BackoffConfig {
22703 initial_ms: 1,
22704 max_ms: 1,
22705 ..Default::default()
22706 },
22707 ..Default::default()
22708 },
22709 ..Default::default()
22710 }))
22711 .build()
22712 .unwrap();
22713
22714 let (content, chunks, final_response) = collect_stream_events(&agent, "Hello").await;
22715 assert!(!chunks.iter().any(StreamChunk::is_error), "{chunks:?}");
22716 assert_eq!(open_attempts.load(Ordering::SeqCst), 2);
22717 assert_eq!(content, "Recovered after retry");
22718 assert_eq!(
22719 final_response.expect("Final").content,
22720 "Recovered after retry"
22721 );
22722 }
22723
22724 #[tokio::test]
22725 async fn test_stream_llm_error_action_error_emits_terminal_error() {
22726 let mut primary = MockLLMProvider::new("primary");
22727 primary.set_error("Primary LLM is unavailable");
22728
22729 let agent = AgentBuilder::new()
22730 .system_prompt("You are helpful.")
22731 .llm(Arc::new(primary))
22732 .build()
22733 .unwrap();
22734
22735 let (_, chunks, final_response) = collect_stream_events(&agent, "Hello").await;
22736 assert!(
22737 final_response.is_none(),
22738 "default Error action must not produce Final"
22739 );
22740 assert!(
22741 chunks.iter().any(StreamChunk::is_error),
22742 "default Error action must surface a stream error"
22743 );
22744 }
22745
22746 struct MidStreamFailureProvider;
22748
22749 #[async_trait]
22750 impl LLMProvider for MidStreamFailureProvider {
22751 async fn complete(
22752 &self,
22753 _messages: &[ChatMessage],
22754 _config: Option<&LLMConfig>,
22755 ) -> std::result::Result<LLMResponse, LLMError> {
22756 Ok(LLMResponse::new("blocking path", FinishReason::Stop))
22757 }
22758
22759 async fn complete_stream(
22760 &self,
22761 _messages: &[ChatMessage],
22762 _config: Option<&LLMConfig>,
22763 ) -> std::result::Result<
22764 Box<dyn Stream<Item = std::result::Result<LLMChunk, LLMError>> + Unpin + Send>,
22765 LLMError,
22766 > {
22767 Ok(Box::new(futures::stream::iter(vec![
22768 Ok(LLMChunk::new("Partial ", false)),
22769 Err(LLMError::Network("connection dropped".to_string())),
22770 ])))
22771 }
22772
22773 fn provider_name(&self) -> &str {
22774 "mid-stream-failure"
22775 }
22776
22777 fn supports(&self, feature: LLMFeature) -> bool {
22778 matches!(feature, LLMFeature::Streaming)
22779 }
22780 }
22781
22782 #[tokio::test]
22783 async fn test_stream_mid_stream_failure_is_terminal() {
22784 let mut fallback = MockLLMProvider::new("fallback");
22785 fallback.set_response("Fallback must not run");
22786 let fallback_calls = fallback.clone();
22787
22788 let agent = AgentBuilder::new()
22789 .system_prompt("You are helpful.")
22790 .llm_alias("default", Arc::new(MidStreamFailureProvider))
22791 .llm_alias("backup", Arc::new(fallback))
22792 .recovery_manager(fallback_llm_recovery("backup"))
22793 .build()
22794 .unwrap();
22795
22796 let (content, chunks, final_response) = collect_stream_events(&agent, "Hello").await;
22797 assert_eq!(content, "Partial ");
22798 assert!(chunks.iter().any(StreamChunk::is_error));
22799 assert!(final_response.is_none());
22800 assert_eq!(
22801 fallback_calls.call_count(),
22802 0,
22803 "fallback must not run after a visible delta"
22804 );
22805 }
22806
22807 #[tokio::test]
22808 async fn test_buffered_streaming_draft_uses_fallback_llm() {
22809 use futures::StreamExt;
22810
22811 let mut primary = MockLLMProvider::new("primary");
22812 primary.set_error("Primary LLM is unavailable");
22813 let fallback = mock_with_response("fallback one two");
22814 let yaml = r#"
22815name: BufferedFallbackAgent
22816system_prompt: "You stream safely."
22817llm:
22818 default: default
22819streaming:
22820 enabled: true
22821 buffer_size: 8
22822runtime:
22823 optimization:
22824 enabled: true
22825 max_speculative_llm_calls_per_turn: 2
22826 speculative_state_transitions: true
22827 streaming_policy: buffer_until_routing_done
22828 max_parallel_runtime_tasks: 2
22829states:
22830 initial: triage
22831 states:
22832 triage:
22833 prompt: "Answer from triage."
22834 transitions:
22835 - to: billing
22836 guard:
22837 context:
22838 route:
22839 eq: billing
22840 timing: parallel
22841 billing:
22842 prompt: "Billing state."
22843"#;
22844 let agent = AgentBuilder::from_yaml(yaml)
22845 .unwrap()
22846 .llm_alias("default", Arc::new(primary))
22847 .llm_alias("backup", Arc::new(fallback))
22848 .recovery_manager(fallback_llm_recovery("backup"))
22849 .build()
22850 .unwrap();
22851
22852 let mut stream = agent.chat_stream("hello").await.unwrap();
22853 let mut content = String::new();
22854 let mut error = None;
22855 while let Some(chunk) = stream.next().await {
22856 match chunk {
22857 StreamChunk::Content { text } => content.push_str(&text),
22858 StreamChunk::Error { message } => error = Some(message),
22859 StreamChunk::Done {} => break,
22860 _ => {}
22861 }
22862 }
22863
22864 assert_eq!(error, None);
22865 assert_eq!(content, "fallback one two");
22866 }
22867
22868 #[tokio::test]
22869 async fn parity_llm_fallback_llm() {
22870 let build = || {
22871 let mut primary = MockLLMProvider::new("primary");
22872 primary.set_error("Primary LLM is unavailable");
22873 let mut fallback = MockLLMProvider::new("fallback");
22874 fallback.set_response("Fallback response works!");
22875 AgentBuilder::new()
22876 .system_prompt("You are helpful.")
22877 .llm_alias("default", Arc::new(primary))
22878 .llm_alias("backup", Arc::new(fallback))
22879 .recovery_manager(fallback_llm_recovery("backup"))
22880 .build()
22881 .unwrap()
22882 };
22883 assert_blocking_streaming_parity(build, "Hello").await;
22884 }
22885
22886 #[tokio::test]
22887 async fn parity_llm_fallback_response() {
22888 use ai_agents_recovery::{ErrorRecoveryConfig, LLMFailureAction, LLMRecoveryConfig};
22889 let build = || {
22890 let mut primary = MockLLMProvider::new("primary");
22891 primary.set_error("Primary LLM is unavailable");
22892 AgentBuilder::new()
22893 .system_prompt("You are helpful.")
22894 .llm(Arc::new(primary))
22895 .recovery_manager(RecoveryManager::new(ErrorRecoveryConfig {
22896 llm: LLMRecoveryConfig {
22897 on_failure: LLMFailureAction::FallbackResponse {
22898 message: "Service is temporarily unavailable.".to_string(),
22899 },
22900 ..Default::default()
22901 },
22902 ..Default::default()
22903 }))
22904 .build()
22905 .unwrap()
22906 };
22907 assert_blocking_streaming_parity(build, "Hello").await;
22908 }
22909
22910 #[tokio::test]
22911 async fn parity_basic_chat() {
22912 let build = || {
22913 AgentBuilder::new()
22914 .system_prompt("You are helpful.")
22915 .llm(Arc::new(mock_with_response("Plain answer")))
22916 .build()
22917 .unwrap()
22918 };
22919 assert_blocking_streaming_parity(build, "Hello").await;
22920 }
22921
22922 fn skills_with_parallel_transition_yaml(extra_optimization: &str, streaming: &str) -> String {
22928 format!(
22929 r#"
22930name: SkillsBesideTransitionAgent
22931system_prompt: "Use skills when they match."
22932llm:
22933 default: default
22934 router: router
22935observability:
22936 enabled: true
22937 export:
22938 write_raw_events: true
22939{streaming}
22940runtime:
22941 optimization:
22942 enabled: true
22943 speculative_state_transitions: true
22944{extra_optimization}
22945states:
22946 initial: triage
22947 states:
22948 triage:
22949 prompt: "Triage state."
22950 transitions:
22951 - to: billing
22952 guard:
22953 context:
22954 route:
22955 eq: billing
22956 timing: parallel
22957 billing:
22958 prompt: "Billing state."
22959skills:
22960 - id: helper
22961 description: "Answer helper requests"
22962 trigger: "User asks for helper"
22963 steps:
22964 - prompt: "Answer the helper request: {{{{ user_input }}}}"
22965 llm: skill
22966"#
22967 )
22968 }
22969
22970 struct RoleMocks {
22974 main: MockLLMProvider,
22975 router: MockLLMProvider,
22976 skill: MockLLMProvider,
22977 }
22978
22979 fn role_mocks(main: MockLLMProvider, router: MockLLMProvider) -> RoleMocks {
22980 RoleMocks {
22981 main,
22982 router,
22983 skill: mock_with_response("Skill step response"),
22984 }
22985 }
22986
22987 fn build_skills_beside_transition_agent(yaml: &str, mocks: RoleMocks) -> RuntimeAgent {
22988 AgentBuilder::from_yaml(yaml)
22989 .unwrap()
22990 .llm_alias("default", Arc::new(mocks.main))
22991 .llm_alias("router", Arc::new(mocks.router))
22992 .llm_alias("skill", Arc::new(mocks.skill))
22993 .build()
22994 .unwrap()
22995 }
22996
22997 fn branch_events_with_commit_behavior(agent: &RuntimeAgent, behavior: &str) -> usize {
22998 agent
22999 .observability()
23000 .unwrap()
23001 .raw_events()
23002 .iter()
23003 .filter(|event| event.dimensions.get("commit_behavior") == Some(&behavior.to_string()))
23004 .count()
23005 }
23006
23007 #[tokio::test]
23008 async fn test_speculative_transition_with_skills_and_no_skill_branch_routes_skill_serially() {
23009 let default_mock = mock_with_response("Draft response");
23010 let router_mock = mock_with_response("helper");
23011 let router_counter = router_mock.clone();
23012 let yaml = skills_with_parallel_transition_yaml(
23013 " max_speculative_llm_calls_per_turn: 2\n max_parallel_runtime_tasks: 2",
23014 "",
23015 );
23016 let agent =
23017 build_skills_beside_transition_agent(&yaml, role_mocks(default_mock, router_mock));
23018
23019 let response = agent.chat("please use helper").await.unwrap();
23020
23021 assert_eq!(
23022 response.metadata.as_ref().and_then(|m| m.get("skill_id")),
23023 Some(&serde_json::json!("helper")),
23024 "skill must route even without a skill branch: {response:?}"
23025 );
23026 assert_eq!(router_counter.call_count(), 1);
23027 assert!(branch_events_with_commit_behavior(&agent, "transition_decision") > 0);
23029 assert_eq!(
23030 branch_events_with_commit_behavior(&agent, "skill_selection"),
23031 0
23032 );
23033 }
23034
23035 #[tokio::test]
23036 async fn test_speculative_transition_with_skills_no_match_commits_draft() {
23037 let default_mock = mock_with_response("Draft response");
23038 let router_mock = mock_with_response("none");
23039 let router_counter = router_mock.clone();
23040 let yaml = skills_with_parallel_transition_yaml(
23041 " max_speculative_llm_calls_per_turn: 2\n max_parallel_runtime_tasks: 2",
23042 "",
23043 );
23044 let agent =
23045 build_skills_beside_transition_agent(&yaml, role_mocks(default_mock, router_mock));
23046
23047 let response = agent.chat("just chat").await.unwrap();
23048
23049 assert_eq!(response.content, "Draft response");
23050 assert!(
23051 response
23052 .metadata
23053 .as_ref()
23054 .is_none_or(|m| !m.contains_key("skill_id"))
23055 );
23056 assert_eq!(router_counter.call_count(), 1);
23057 assert!(branch_events_with_commit_behavior(&agent, "final_response") > 0);
23058 }
23059
23060 #[tokio::test]
23061 async fn test_speculative_transition_win_skips_serial_skill_selection() {
23062 let default_mock = mock_with_response("Billing answer");
23063 let router_mock = mock_with_response("none");
23064 let router_counter = router_mock.clone();
23065 let yaml = skills_with_parallel_transition_yaml(
23066 " max_speculative_llm_calls_per_turn: 2\n max_parallel_runtime_tasks: 2",
23067 "",
23068 );
23069 let agent =
23070 build_skills_beside_transition_agent(&yaml, role_mocks(default_mock, router_mock));
23071 agent
23072 .set_context("route", serde_json::json!("billing"))
23073 .unwrap();
23074
23075 let response = agent.chat("billing please").await.unwrap();
23076
23077 assert_eq!(agent.current_state().as_deref(), Some("billing"));
23078 assert_eq!(response.content, "Billing answer");
23079 assert_eq!(router_counter.call_count(), 1);
23081 }
23082
23083 #[tokio::test]
23084 async fn test_speculative_skill_capacity_exhausted_still_routes_skill_serially() {
23085 let default_mock = mock_with_response("Draft response");
23086 let router_mock = mock_with_response("helper");
23087 let router_counter = router_mock.clone();
23088 let yaml = skills_with_parallel_transition_yaml(
23090 " speculative_skill_routing: true\n max_speculative_llm_calls_per_turn: 2\n max_parallel_runtime_tasks: 2",
23091 "",
23092 );
23093 let agent =
23094 build_skills_beside_transition_agent(&yaml, role_mocks(default_mock, router_mock));
23095
23096 let response = agent.chat("please use helper").await.unwrap();
23097
23098 assert_eq!(
23099 response.metadata.as_ref().and_then(|m| m.get("skill_id")),
23100 Some(&serde_json::json!("helper"))
23101 );
23102 assert_eq!(router_counter.call_count(), 1);
23103 assert_eq!(
23104 branch_events_with_commit_behavior(&agent, "skill_selection"),
23105 0
23106 );
23107 }
23108
23109 #[tokio::test]
23110 async fn test_speculative_transition_and_skill_both_enabled_unchanged() {
23111 let default_mock = mock_with_response("Draft response");
23112 let router_mock = mock_with_response("helper");
23113 let router_counter = router_mock.clone();
23114 let yaml = skills_with_parallel_transition_yaml(
23115 " speculative_skill_routing: true\n max_speculative_llm_calls_per_turn: 3\n max_parallel_runtime_tasks: 3",
23116 "",
23117 );
23118 let agent =
23119 build_skills_beside_transition_agent(&yaml, role_mocks(default_mock, router_mock));
23120
23121 let response = agent.chat("please use helper").await.unwrap();
23122
23123 assert_eq!(
23124 response.metadata.as_ref().and_then(|m| m.get("skill_id")),
23125 Some(&serde_json::json!("helper"))
23126 );
23127 assert_eq!(router_counter.call_count(), 1);
23128 assert!(branch_events_with_commit_behavior(&agent, "skill_selection") > 0);
23130 }
23131
23132 const BUFFERED_STREAMING_YAML_FRAGMENT: &str = "streaming:\n enabled: true\n buffer_size: 16";
23133 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";
23134
23135 #[tokio::test]
23136 async fn test_buffered_streaming_skill_wins_after_transition_miss() {
23137 let mut default_mock = mock_with_response("draft one two");
23138 default_mock.set_latency(10);
23139 let router_mock = mock_with_response("helper");
23140 let router_counter = router_mock.clone();
23141 let yaml = skills_with_parallel_transition_yaml(
23142 BUFFERED_OPTIMIZATION_FRAGMENT,
23143 BUFFERED_STREAMING_YAML_FRAGMENT,
23144 );
23145 let agent =
23146 build_skills_beside_transition_agent(&yaml, role_mocks(default_mock, router_mock));
23147
23148 let (content, chunks, final_response) =
23149 collect_stream_events(&agent, "please use helper").await;
23150
23151 assert!(!chunks.iter().any(StreamChunk::is_error), "{chunks:?}");
23152 assert!(
23153 !content.contains("draft"),
23154 "buffered draft must be discarded when a skill wins: {content:?}"
23155 );
23156 let final_response = final_response.expect("Final");
23157 assert_eq!(
23158 final_response
23159 .metadata
23160 .as_ref()
23161 .and_then(|m| m.get("skill_id")),
23162 Some(&serde_json::json!("helper"))
23163 );
23164 assert_eq!(content, final_response.content);
23165 assert_eq!(router_counter.call_count(), 1);
23166 }
23167
23168 #[tokio::test]
23169 async fn test_buffered_streaming_skill_miss_releases_buffer_and_commits_draft() {
23170 let mut default_mock = mock_with_response("draft one two");
23171 default_mock.set_latency(10);
23172 let router_mock = mock_with_response("none");
23173 let router_counter = router_mock.clone();
23174 let yaml = skills_with_parallel_transition_yaml(
23175 BUFFERED_OPTIMIZATION_FRAGMENT,
23176 BUFFERED_STREAMING_YAML_FRAGMENT,
23177 );
23178 let agent =
23179 build_skills_beside_transition_agent(&yaml, role_mocks(default_mock, router_mock));
23180
23181 let (content, chunks, final_response) = collect_stream_events(&agent, "just chat").await;
23182
23183 assert!(!chunks.iter().any(StreamChunk::is_error), "{chunks:?}");
23184 assert_eq!(content, "draft one two");
23185 assert_eq!(final_response.expect("Final").content, "draft one two");
23186 assert_eq!(router_counter.call_count(), 1);
23187 }
23188
23189 #[tokio::test]
23190 async fn test_buffered_streaming_transition_win_skips_skill_selection() {
23191 let default_mock = mock_with_response("Billing answer");
23192 let router_mock = mock_with_response("none");
23193 let router_counter = router_mock.clone();
23194 let yaml = skills_with_parallel_transition_yaml(
23195 BUFFERED_OPTIMIZATION_FRAGMENT,
23196 BUFFERED_STREAMING_YAML_FRAGMENT,
23197 );
23198 let agent =
23199 build_skills_beside_transition_agent(&yaml, role_mocks(default_mock, router_mock));
23200 agent
23201 .set_context("route", serde_json::json!("billing"))
23202 .unwrap();
23203
23204 let (content, chunks, final_response) =
23205 collect_stream_events(&agent, "billing please").await;
23206
23207 assert!(!chunks.iter().any(StreamChunk::is_error), "{chunks:?}");
23208 assert_eq!(agent.current_state().as_deref(), Some("billing"));
23209 assert_eq!(content, "Billing answer");
23210 assert_eq!(final_response.expect("Final").content, "Billing answer");
23211 assert_eq!(router_counter.call_count(), 1);
23213 }
23214
23215 #[tokio::test]
23216 async fn parity_buffered_policy_with_skills() {
23217 let yaml = skills_with_parallel_transition_yaml(
23218 BUFFERED_OPTIMIZATION_FRAGMENT,
23219 BUFFERED_STREAMING_YAML_FRAGMENT,
23220 );
23221 let build = || {
23222 build_skills_beside_transition_agent(
23223 &yaml,
23224 role_mocks(
23225 mock_with_response("draft one two"),
23226 mock_with_response("helper"),
23227 ),
23228 )
23229 };
23230 let (blocking, _, _) = assert_blocking_streaming_parity(build, "please use helper").await;
23231 assert_eq!(
23232 blocking.metadata.as_ref().and_then(|m| m.get("skill_id")),
23233 Some(&serde_json::json!("helper"))
23234 );
23235 }
23236
23237 #[tokio::test]
23238 async fn parity_buffered_policy_with_cot() {
23239 let yaml = format!(
23240 r#"
23241name: BufferedCotAgent
23242system_prompt: "Think first."
23243llm:
23244 default: default
23245streaming:
23246 enabled: true
23247 buffer_size: 16
23248reasoning:
23249 mode: cot
23250runtime:
23251 optimization:
23252 enabled: true
23253 speculative_state_transitions: true
23254{BUFFERED_OPTIMIZATION_FRAGMENT}
23255states:
23256 initial: triage
23257 states:
23258 triage:
23259 prompt: "Triage state."
23260 transitions:
23261 - to: billing
23262 guard:
23263 context:
23264 route:
23265 eq: billing
23266 timing: parallel
23267 billing:
23268 prompt: "Billing state."
23269"#
23270 );
23271 let build = || {
23272 AgentBuilder::from_yaml(&yaml)
23273 .unwrap()
23274 .llm_alias(
23275 "default",
23276 Arc::new(mock_with_response(
23277 "<thinking>step by step</thinking>Reasoned answer",
23278 )),
23279 )
23280 .build()
23281 .unwrap()
23282 };
23283 let (blocking, streamed, _) = assert_blocking_streaming_parity(build, "hello").await;
23284 assert_eq!(blocking.content, "Reasoned answer");
23285 let mode = streamed
23286 .metadata
23287 .as_ref()
23288 .and_then(|m| m.get("reasoning"))
23289 .and_then(|r| r.get("mode_used"))
23290 .cloned();
23291 assert_eq!(
23293 mode,
23294 Some(serde_json::to_value(ReasoningMode::CoT).unwrap())
23295 );
23296 }
23297
23298 fn post_response_transition_yaml(states_extra: &str, billing_extra: &str) -> String {
23305 format!(
23306 r#"
23307name: PostResponseTransitionAgent
23308system_prompt: "You are helpful."
23309streaming:
23310 enabled: true
23311states:
23312 initial: intake
23313{states_extra}
23314 states:
23315 intake:
23316 prompt: "Intake"
23317 transitions:
23318 - to: billing
23319 guard:
23320 context:
23321 route:
23322 eq: billing
23323 billing:
23324 prompt: "Billing"
23325{billing_extra}
23326"#
23327 )
23328 }
23329
23330 fn build_post_response_transition_agent(yaml: &str, mock: MockLLMProvider) -> RuntimeAgent {
23331 let agent = AgentBuilder::from_yaml(yaml)
23332 .unwrap()
23333 .llm(Arc::new(mock))
23334 .build()
23335 .unwrap();
23336 agent
23337 .set_context("route", serde_json::json!("billing"))
23338 .unwrap();
23339 agent
23340 }
23341
23342 fn count_occurrences(haystack: &str, needle: &str) -> usize {
23343 haystack.matches(needle).count()
23344 }
23345
23346 #[tokio::test]
23347 async fn test_stream_transition_without_regeneration_emits_content_once() {
23348 let yaml = post_response_transition_yaml(" regenerate_on_transition: false", "");
23349 let agent =
23350 build_post_response_transition_agent(&yaml, mock_with_response("Intake answer"));
23351
23352 let (content, chunks, final_response) = collect_stream_events(&agent, "hello").await;
23353
23354 assert!(!chunks.iter().any(StreamChunk::is_error), "{chunks:?}");
23355 assert_eq!(agent.current_state().as_deref(), Some("billing"));
23356 assert_eq!(
23357 count_occurrences(&content, "Intake answer"),
23358 1,
23359 "committed content must not be emitted twice: {content:?}"
23360 );
23361 assert_eq!(final_response.expect("Final").content, content);
23362 assert!(
23363 chunks
23364 .iter()
23365 .any(|c| matches!(c, StreamChunk::StateTransition { .. }))
23366 );
23367 }
23368
23369 #[tokio::test]
23370 async fn test_stream_transition_without_regeneration_buffered_emits_content_once() {
23371 let yaml = post_response_transition_yaml(" regenerate_on_transition: false", "");
23372 let mut mock = mock_with_response("Intake answer");
23373 mock.set_tool_choice(Some(ToolChoice::Auto));
23375 let agent = build_post_response_transition_agent(&yaml, mock);
23376
23377 let (content, chunks, final_response) = collect_stream_events(&agent, "hello").await;
23378
23379 assert!(!chunks.iter().any(StreamChunk::is_error), "{chunks:?}");
23380 assert_eq!(agent.current_state().as_deref(), Some("billing"));
23381 assert_eq!(
23382 count_occurrences(&content, "Intake answer"),
23383 1,
23384 "{content:?}"
23385 );
23386 assert_eq!(final_response.expect("Final").content, content);
23387 }
23388
23389 #[tokio::test]
23390 async fn test_stream_state_regenerate_on_enter_false_emits_content_once() {
23391 let yaml = post_response_transition_yaml("", " regenerate_on_enter: false");
23392 let agent =
23393 build_post_response_transition_agent(&yaml, mock_with_response("Intake answer"));
23394
23395 let (content, chunks, final_response) = collect_stream_events(&agent, "hello").await;
23396
23397 assert!(!chunks.iter().any(StreamChunk::is_error), "{chunks:?}");
23398 assert_eq!(agent.current_state().as_deref(), Some("billing"));
23399 assert_eq!(
23400 count_occurrences(&content, "Intake answer"),
23401 1,
23402 "{content:?}"
23403 );
23404 assert_eq!(final_response.expect("Final").content, content);
23405 }
23406
23407 #[tokio::test]
23408 async fn test_stream_transition_with_regeneration_emits_replacement() {
23409 let yaml = post_response_transition_yaml("", "");
23410 let agent = build_post_response_transition_agent(
23411 &yaml,
23412 mock_with_responses(vec!["Intake answer", "Billing answer"]),
23413 );
23414
23415 let (content, chunks, final_response) = collect_stream_events(&agent, "hello").await;
23416
23417 assert!(!chunks.iter().any(StreamChunk::is_error), "{chunks:?}");
23418 assert_eq!(agent.current_state().as_deref(), Some("billing"));
23419 assert_eq!(count_occurrences(&content, "Intake answer"), 1);
23421 assert_eq!(count_occurrences(&content, "Billing answer"), 1);
23422 assert_eq!(final_response.expect("Final").content, "Billing answer");
23423 }
23424
23425 #[tokio::test]
23426 async fn test_blocking_transition_without_regeneration_unchanged() {
23427 let yaml = post_response_transition_yaml(" regenerate_on_transition: false", "");
23428 let agent =
23429 build_post_response_transition_agent(&yaml, mock_with_response("Intake answer"));
23430
23431 let response = agent.chat("hello").await.unwrap();
23432
23433 assert_eq!(response.content, "Intake answer");
23434 assert_eq!(agent.current_state().as_deref(), Some("billing"));
23435 }
23436
23437 #[tokio::test]
23438 async fn parity_transition_regenerate_off() {
23439 let yaml = post_response_transition_yaml(" regenerate_on_transition: false", "");
23440 let build =
23441 || build_post_response_transition_agent(&yaml, mock_with_response("Intake answer"));
23442 assert_blocking_streaming_parity(build, "hello").await;
23443 }
23444
23445 #[tokio::test]
23446 async fn parity_transition_regenerate_on() {
23447 let yaml = post_response_transition_yaml("", "");
23448 let build = || {
23449 build_post_response_transition_agent(
23450 &yaml,
23451 mock_with_responses(vec!["Intake answer", "Billing answer"]),
23452 )
23453 };
23454 let (blocking, _, _) = assert_blocking_streaming_parity(build, "hello").await;
23455 assert_eq!(blocking.content, "Billing answer");
23456 }
23457
23458 fn rejecting_process_processor() -> ProcessProcessor {
23463 use ai_agents_process::ProcessConfig;
23464 let validate_config = ai_agents_process::ValidateStage {
23465 id: Some("length_check".to_string()),
23466 condition: None,
23467 config: ai_agents_process::ValidateConfig {
23468 rules: vec![ai_agents_process::ValidationRule::MinLength {
23469 min_length: 10,
23470 on_fail: ai_agents_process::ValidationAction {
23471 action: ai_agents_process::ValidationActionType::Reject,
23472 message: None,
23473 },
23474 }],
23475 ..Default::default()
23476 },
23477 };
23478 ProcessProcessor::new(ProcessConfig {
23479 input: vec![ai_agents_process::ProcessStage::Validate(validate_config)],
23480 ..Default::default()
23481 })
23482 }
23483
23484 fn looks_like_rejection(content: &str) -> bool {
23486 content.contains("rejected")
23487 || content.contains("Input rejected")
23488 || content.contains("too short")
23489 || content.contains("Too short")
23490 || content.len() < 50
23491 }
23492
23493 #[tokio::test]
23494 async fn test_stream_input_rejection_is_final_response() {
23495 let responses = Arc::new(std::sync::atomic::AtomicUsize::new(0));
23496 let hooks = Arc::new(ResponseCountingHooks {
23497 responses: Arc::clone(&responses),
23498 });
23499 let mock = mock_with_response("Should not reach here.");
23500 let llm_calls = mock.clone();
23501 let agent = AgentBuilder::new()
23502 .system_prompt("You are helpful.")
23503 .llm(Arc::new(mock))
23504 .process_processor(rejecting_process_processor())
23505 .hooks(hooks.clone())
23506 .build()
23507 .unwrap();
23508
23509 let (content, chunks, final_response) = collect_stream_events(&agent, "Hi").await;
23510
23511 assert!(
23512 !chunks.iter().any(StreamChunk::is_error),
23513 "rejection is a response, not a stream error: {chunks:?}"
23514 );
23515 let final_response = final_response.expect("rejection must finalize as Final");
23516 assert!(
23517 looks_like_rejection(&final_response.content),
23518 "Expected rejection response, got: {}",
23519 final_response.content
23520 );
23521 assert_eq!(content, final_response.content);
23522 assert_eq!(
23523 llm_calls.call_count(),
23524 0,
23525 "rejected input must not reach the LLM"
23526 );
23527 assert_eq!(responses.load(Ordering::SeqCst), 1, "on_response must fire");
23528 }
23529
23530 #[tokio::test]
23531 async fn parity_input_rejection() {
23532 let build = || {
23533 AgentBuilder::new()
23534 .system_prompt("You are helpful.")
23535 .llm(Arc::new(mock_with_response("Should not reach here.")))
23536 .process_processor(rejecting_process_processor())
23537 .build()
23538 .unwrap()
23539 };
23540 let (blocking, _, _) = assert_blocking_streaming_parity(build, "Hi").await;
23541 assert!(
23542 looks_like_rejection(&blocking.content),
23543 "{}",
23544 blocking.content
23545 );
23546 }
23547
23548 fn pre_response_transition_yaml(streaming_policy: &str) -> String {
23549 format!(
23550 r#"
23551name: StreamingPreflightAgent
23552system_prompt: "You route before streaming."
23553runtime:
23554 optimization:
23555 enabled: true
23556 pre_response_deterministic_transitions: true
23557 streaming_policy: {streaming_policy}
23558streaming:
23559 enabled: true
23560 buffer_size: 16
23561states:
23562 initial: greeting
23563 states:
23564 greeting:
23565 prompt: "OLD_STATE_SENTINEL"
23566 transitions:
23567 - to: billing
23568 guard:
23569 context:
23570 topic:
23571 eq: billing
23572 timing: pre_response
23573 billing:
23574 prompt: "Billing state."
23575"#
23576 )
23577 }
23578
23579 fn build_pre_response_transition_agent(yaml: &str) -> RuntimeAgent {
23580 let agent = AgentBuilder::from_yaml(yaml)
23581 .unwrap()
23582 .llm(Arc::new(mock_with_response("Billing streamed response")))
23583 .build()
23584 .unwrap();
23585 agent
23586 .set_context("topic", serde_json::json!("billing"))
23587 .unwrap();
23588 agent
23589 }
23590
23591 #[tokio::test]
23592 async fn test_stream_buffered_policy_runs_pre_response_deterministic_transition() {
23593 let yaml = pre_response_transition_yaml("buffer_until_routing_done");
23594 let agent = build_pre_response_transition_agent(&yaml);
23595
23596 let (content, chunks, final_response) =
23597 collect_stream_events(&agent, "billing please").await;
23598
23599 assert!(!chunks.iter().any(StreamChunk::is_error), "{chunks:?}");
23600 assert_eq!(agent.current_state().as_deref(), Some("billing"));
23601 assert!(content.contains("Billing streamed response"));
23602 assert!(!content.contains("OLD_STATE_SENTINEL"));
23603 assert_eq!(final_response.expect("Final").content, content);
23604 }
23605
23606 #[tokio::test]
23607 async fn test_stream_disabled_policy_skips_preflight() {
23608 let yaml = pre_response_transition_yaml("disabled");
23609 let agent = build_pre_response_transition_agent(&yaml);
23610
23611 let (_, chunks, final_response) = collect_stream_events(&agent, "billing please").await;
23612
23613 assert!(!chunks.iter().any(StreamChunk::is_error), "{chunks:?}");
23614 assert!(final_response.is_some());
23615 assert_eq!(agent.current_state().as_deref(), Some("greeting"));
23619 }
23620
23621 #[tokio::test]
23622 async fn parity_pre_response_transition_buffered_policy() {
23623 let yaml = pre_response_transition_yaml("buffer_until_routing_done");
23624 let build = || build_pre_response_transition_agent(&yaml);
23625 assert_blocking_streaming_parity(build, "billing please").await;
23626 }
23627
23628 fn calculator_agent_with(mock: MockLLMProvider) -> RuntimeAgent {
23633 let mut tools = ai_agents_tools::ToolRegistry::new();
23634 tools
23635 .register(Arc::new(ai_agents_tools::CalculatorTool))
23636 .unwrap();
23637 AgentBuilder::new()
23638 .system_prompt("You are a calculator assistant.")
23639 .llm(Arc::new(mock))
23640 .tools(tools)
23641 .build()
23642 .unwrap()
23643 }
23644
23645 #[tokio::test]
23646 async fn test_stream_tool_start_events_precede_results_for_batch() {
23647 let mock = mock_with_responses(vec![
23648 r#"[{"tool": "calculator", "arguments": {"expression": "1+1"}}, {"tool": "calculator", "arguments": {"expression": "2+2"}}]"#,
23649 "Both answers are ready.",
23650 ]);
23651 let agent = calculator_agent_with(mock);
23652
23653 let (_, chunks, final_response) = collect_stream_events(&agent, "compute both").await;
23654
23655 assert!(!chunks.iter().any(StreamChunk::is_error), "{chunks:?}");
23656 let final_response = final_response.expect("Final");
23657 assert_eq!(final_response.tool_calls.as_ref().map(Vec::len), Some(2));
23658
23659 let tool_events: Vec<&StreamChunk> = chunks
23660 .iter()
23661 .filter(|c| {
23662 matches!(
23663 c,
23664 StreamChunk::ToolCallStart { .. }
23665 | StreamChunk::ToolResult { .. }
23666 | StreamChunk::ToolCallEnd { .. }
23667 )
23668 })
23669 .collect();
23670 assert_eq!(tool_events.len(), 6, "{tool_events:?}");
23671 assert!(matches!(tool_events[0], StreamChunk::ToolCallStart { .. }));
23673 assert!(matches!(tool_events[1], StreamChunk::ToolCallStart { .. }));
23674 assert!(matches!(
23675 tool_events[2],
23676 StreamChunk::ToolResult { success: true, .. }
23677 ));
23678 assert!(matches!(tool_events[3], StreamChunk::ToolCallEnd { .. }));
23679 assert!(matches!(
23680 tool_events[4],
23681 StreamChunk::ToolResult { success: true, .. }
23682 ));
23683 assert!(matches!(tool_events[5], StreamChunk::ToolCallEnd { .. }));
23684 }
23685
23686 #[tokio::test]
23687 async fn test_stream_clarification_final_carries_options_and_detection() {
23688 let responses = || {
23689 vec![
23690 r#"{"is_ambiguous":true,"confidence":0.2,"ambiguity_type":"missing_target","reasoning":"target missing","what_is_unclear":["target"],"detected_language":"en"}"#,
23691 r#"{"question":"What should I send?","options":["report","invoice"]}"#,
23692 ]
23693 };
23694 let (blocking_agent, _) = state_disambiguation_agent(responses(), true, None, true);
23695 let (streaming_agent, _) = state_disambiguation_agent(responses(), true, None, true);
23696
23697 let blocking = blocking_agent.chat("Send it").await.unwrap();
23698 let (_, chunks, streamed) = collect_stream_events(&streaming_agent, "Send it").await;
23699 assert!(!chunks.iter().any(StreamChunk::is_error), "{chunks:?}");
23700 let streamed = streamed.expect("clarification must finalize as Final");
23701
23702 assert_eq!(streamed.content, "What should I send?");
23703 let streamed_meta = streamed
23704 .metadata
23705 .as_ref()
23706 .and_then(|m| m.get("disambiguation"))
23707 .cloned()
23708 .expect("disambiguation metadata");
23709 for key in ["status", "options", "clarifying", "detection"] {
23710 assert!(
23711 streamed_meta.get(key).is_some(),
23712 "missing {key}: {streamed_meta}"
23713 );
23714 }
23715 assert_eq!(
23716 streamed_meta.get("detection").and_then(|d| d.get("type")),
23717 Some(&serde_json::json!("missing_target"))
23718 );
23719 assert_eq!(
23720 blocking
23721 .metadata
23722 .as_ref()
23723 .and_then(|m| m.get("disambiguation")),
23724 Some(&streamed_meta),
23725 "blocking and streaming clarification metadata must be identical"
23726 );
23727 }
23728
23729 struct FailingMemory {
23731 messages: parking_lot::RwLock<Vec<ChatMessage>>,
23732 fail_on_add: usize,
23733 adds: std::sync::atomic::AtomicUsize,
23734 }
23735
23736 #[async_trait]
23737 impl ai_agents_core::Memory for FailingMemory {
23738 async fn add_message(&self, message: ChatMessage) -> Result<()> {
23739 let n = self.adds.fetch_add(1, Ordering::SeqCst) + 1;
23740 if n == self.fail_on_add {
23741 return Err(AgentError::Other(format!(
23742 "simulated memory failure on add #{n}"
23743 )));
23744 }
23745 self.messages.write().push(message);
23746 Ok(())
23747 }
23748
23749 async fn get_messages(&self, limit: Option<usize>) -> Result<Vec<ChatMessage>> {
23750 let messages = self.messages.read();
23751 Ok(match limit {
23752 Some(n) if n < messages.len() => messages[messages.len() - n..].to_vec(),
23753 _ => messages.clone(),
23754 })
23755 }
23756
23757 async fn clear(&self) -> Result<()> {
23758 self.messages.write().clear();
23759 Ok(())
23760 }
23761
23762 fn len(&self) -> usize {
23763 self.messages.read().len()
23764 }
23765
23766 async fn restore(&self, snapshot: ai_agents_core::MemorySnapshot) -> Result<()> {
23767 *self.messages.write() = snapshot.messages;
23768 Ok(())
23769 }
23770 }
23771
23772 impl ai_agents_memory::Memory for FailingMemory {}
23773
23774 struct ReadFailingMemory {
23776 messages: parking_lot::RwLock<Vec<ChatMessage>>,
23777 fail_on_read: Option<usize>,
23778 reads: std::sync::atomic::AtomicUsize,
23779 }
23780
23781 impl ReadFailingMemory {
23782 fn new(fail_reads: bool) -> Self {
23784 Self::fail_on_read(fail_reads.then_some(1))
23785 }
23786
23787 fn fail_on_read(fail_on_read: Option<usize>) -> Self {
23789 Self {
23790 messages: parking_lot::RwLock::new(Vec::new()),
23791 fail_on_read,
23792 reads: std::sync::atomic::AtomicUsize::new(0),
23793 }
23794 }
23795
23796 fn read_count(&self) -> usize {
23798 self.reads.load(Ordering::SeqCst)
23799 }
23800 }
23801
23802 #[async_trait]
23803 impl ai_agents_core::Memory for ReadFailingMemory {
23804 async fn add_message(&self, message: ChatMessage) -> Result<()> {
23806 self.messages.write().push(message);
23807 Ok(())
23808 }
23809
23810 async fn get_messages(&self, limit: Option<usize>) -> Result<Vec<ChatMessage>> {
23812 let read = self.reads.fetch_add(1, Ordering::SeqCst) + 1;
23813 if self.fail_on_read == Some(read) {
23814 return Err(AgentError::Other(
23815 "simulated scope memory failure".to_string(),
23816 ));
23817 }
23818 let messages = self.messages.read();
23819 Ok(match limit {
23820 Some(n) if n < messages.len() => messages[messages.len() - n..].to_vec(),
23821 _ => messages.clone(),
23822 })
23823 }
23824
23825 async fn clear(&self) -> Result<()> {
23827 self.messages.write().clear();
23828 Ok(())
23829 }
23830
23831 fn len(&self) -> usize {
23833 self.messages.read().len()
23834 }
23835
23836 async fn restore(&self, snapshot: ai_agents_core::MemorySnapshot) -> Result<()> {
23838 *self.messages.write() = snapshot.messages;
23839 Ok(())
23840 }
23841 }
23842
23843 impl ai_agents_memory::Memory for ReadFailingMemory {}
23844
23845 fn scoped_planning_agent(
23847 memory: Arc<ReadFailingMemory>,
23848 planner: MockLLMProvider,
23849 ) -> RuntimeAgent {
23850 let yaml = r#"
23851name: ScopedPlanningAgent
23852system_prompt: "Plan safely."
23853reasoning:
23854 mode: plan_and_execute
23855tools: [calculator]
23856states:
23857 initial: current
23858 states:
23859 current:
23860 tools: [calculator]
23861"#;
23862 AgentBuilder::from_yaml(yaml)
23863 .unwrap()
23864 .llm(Arc::new(planner))
23865 .tool(Arc::new(CalculatorTool::new()))
23866 .tool(Arc::new(FileWriteTool::new()))
23867 .memory(memory)
23868 .build()
23869 .unwrap()
23870 }
23871
23872 fn scoped_disambiguation_agent(
23874 memory: Arc<ReadFailingMemory>,
23875 router: MockLLMProvider,
23876 include_available_tools: bool,
23877 ) -> RuntimeAgent {
23878 let yaml = format!(
23879 r#"
23880name: ScopedDisambiguationAgent
23881system_prompt: "Clarify safely."
23882llm:
23883 default: default
23884 router: router
23885disambiguation:
23886 enabled: true
23887 context:
23888 recent_messages: 0
23889 include_available_tools: {include_available_tools}
23890tools: [calculator]
23891states:
23892 initial: current
23893 states:
23894 current:
23895 tools: [calculator]
23896"#
23897 );
23898 AgentBuilder::from_yaml(&yaml)
23899 .unwrap()
23900 .llm_alias("default", Arc::new(mock_with_response("unused")))
23901 .llm_alias("router", Arc::new(router))
23902 .tool(Arc::new(CalculatorTool::new()))
23903 .tool(Arc::new(FileWriteTool::new()))
23904 .memory(memory)
23905 .build()
23906 .unwrap()
23907 }
23908
23909 #[tokio::test]
23910 async fn planning_scope_failure_stops_before_the_planner_in_blocking_and_streaming_turns() {
23911 use futures::StreamExt;
23912
23913 let direct_memory = Arc::new(ReadFailingMemory::new(true));
23914 let direct_planner = mock_with_response(r#"{"steps":[]}"#);
23915 let direct_calls = direct_planner.clone();
23916 let direct = scoped_planning_agent(direct_memory.clone(), direct_planner);
23917 let error = direct.generate_plan("plan this").await.unwrap_err();
23918 assert!(error.to_string().contains("simulated scope memory failure"));
23919 assert_eq!(direct_memory.read_count(), 1);
23920 assert_eq!(direct_calls.call_count(), 0);
23921
23922 let blocking_memory = Arc::new(ReadFailingMemory::new(true));
23923 let blocking_planner = mock_with_response(r#"{"steps":[]}"#);
23924 let blocking_calls = blocking_planner.clone();
23925 let blocking = scoped_planning_agent(blocking_memory.clone(), blocking_planner);
23926 let error = blocking.chat("plan this").await.unwrap_err();
23927 assert!(error.to_string().contains("simulated scope memory failure"));
23928 assert_eq!(blocking_memory.read_count(), 1);
23929 assert_eq!(blocking_calls.call_count(), 0);
23930 assert!(blocking.tool_call_history.read().is_empty());
23931
23932 let streaming_memory = Arc::new(ReadFailingMemory::new(true));
23933 let streaming_planner = mock_with_response(r#"{"steps":[]}"#);
23934 let streaming_calls = streaming_planner.clone();
23935 let streaming = scoped_planning_agent(streaming_memory.clone(), streaming_planner);
23936 let (_, chunks, final_response) = collect_stream_events(&streaming, "plan this").await;
23937 assert!(final_response.is_none());
23938 assert!(chunks.iter().any(|chunk| matches!(
23939 chunk,
23940 StreamChunk::Error { message } if message.contains("simulated scope memory failure")
23941 )));
23942 assert!(!chunks.iter().any(StreamChunk::is_done));
23943 assert_eq!(streaming_memory.read_count(), 1);
23944 assert_eq!(streaming_calls.call_count(), 0);
23945 assert!(streaming.tool_call_history.read().is_empty());
23946
23947 let legacy_memory = Arc::new(ReadFailingMemory::new(true));
23948 let legacy_planner = mock_with_response(r#"{"steps":[]}"#);
23949 let legacy_calls = legacy_planner.clone();
23950 let legacy = scoped_planning_agent(legacy_memory.clone(), legacy_planner);
23951 let mut stream = legacy.chat_stream("plan this").await.unwrap();
23952 let mut chunks = Vec::new();
23953 while let Some(chunk) = stream.next().await {
23954 chunks.push(chunk);
23955 }
23956 assert!(chunks.iter().any(|chunk| matches!(
23957 chunk,
23958 StreamChunk::Error { message } if message.contains("simulated scope memory failure")
23959 )));
23960 assert!(!chunks.iter().any(StreamChunk::is_done));
23961 assert_eq!(legacy_memory.read_count(), 1);
23962 assert_eq!(legacy_calls.call_count(), 0);
23963 }
23964
23965 #[tokio::test]
23966 async fn replanning_scope_failure_does_not_issue_a_second_planner_request() {
23967 let memory = Arc::new(ReadFailingMemory::fail_on_read(Some(4)));
23968 let planner = mock_with_response(
23969 r#"{"steps":[{"id":"step1","description":"invalid calculation","action_type":"tool","action_target":"calculator","args":{"expression":"not valid"},"dependencies":[]}]}"#,
23970 );
23971 let planner_calls = planner.clone();
23972 let agent = scoped_planning_agent(memory.clone(), planner);
23973
23974 let result = agent.chat("plan this").await;
23975 assert!(
23976 result.is_err(),
23977 "expected replan scope failure, got {result:?}; reads={}, planner_calls={}",
23978 memory.read_count(),
23979 planner_calls.call_count()
23980 );
23981 let error = result.unwrap_err();
23982 assert!(error.to_string().contains("simulated scope memory failure"));
23983 assert_eq!(memory.read_count(), 4);
23984 assert_eq!(planner_calls.call_count(), 1);
23985 let records = agent.tool_call_history.read();
23986 assert_eq!(records.len(), 1);
23987 assert_eq!(records[0].tool_id, "calculator");
23988 }
23989
23990 #[tokio::test]
23991 async fn planning_prompt_uses_only_the_effective_tool_scope() {
23992 let memory = Arc::new(ReadFailingMemory::new(false));
23993 let planner = mock_with_response(r#"{"steps":[]}"#);
23994 let planner_calls = planner.clone();
23995 let agent = scoped_planning_agent(memory.clone(), planner);
23996
23997 let plan = agent.generate_plan("plan this").await.unwrap();
23998 assert!(!plan.steps.is_empty());
23999 assert_eq!(memory.read_count(), 1);
24000 assert_eq!(planner_calls.call_count(), 1);
24001 let call = planner_calls.last_call().unwrap();
24002 let prompt = &call.messages[0].content;
24003 assert!(prompt.contains("- calculator ("), "{prompt}");
24004 assert!(!prompt.contains("file_write"), "{prompt}");
24005 }
24006
24007 #[tokio::test]
24008 async fn planning_with_no_granted_tools_keeps_a_normal_empty_scope() {
24009 let yaml = r#"
24010name: EmptyPlanningAgent
24011system_prompt: "Plan safely."
24012reasoning:
24013 mode: plan_and_execute
24014tools: []
24015"#;
24016 let planner = mock_with_response(r#"{"steps":[]}"#);
24017 let planner_calls = planner.clone();
24018 let agent = AgentBuilder::from_yaml(yaml)
24019 .unwrap()
24020 .llm(Arc::new(planner))
24021 .tool(Arc::new(CalculatorTool::new()))
24022 .build()
24023 .unwrap();
24024
24025 agent.generate_plan("plan this").await.unwrap();
24026 let call = planner_calls.last_call().unwrap();
24027 let prompt = &call.messages[0].content;
24028 assert!(prompt.contains("Available tools: none"), "{prompt}");
24029 assert!(!prompt.contains("- calculator ("), "{prompt}");
24030 }
24031
24032 #[tokio::test]
24033 async fn planning_filter_can_narrow_the_effective_scope_to_empty() {
24034 let yaml = r#"
24035name: FilteredPlanningAgent
24036system_prompt: "Plan safely."
24037reasoning:
24038 mode: plan_and_execute
24039 planning:
24040 available:
24041 tools: []
24042tools: [calculator]
24043states:
24044 initial: current
24045 states:
24046 current:
24047 tools: [calculator]
24048"#;
24049 let memory = Arc::new(ReadFailingMemory::new(false));
24050 let planner = mock_with_response(r#"{"steps":[]}"#);
24051 let planner_calls = planner.clone();
24052 let agent = AgentBuilder::from_yaml(yaml)
24053 .unwrap()
24054 .llm(Arc::new(planner))
24055 .tool(Arc::new(CalculatorTool::new()))
24056 .memory(memory.clone())
24057 .build()
24058 .unwrap();
24059
24060 agent.generate_plan("plan this").await.unwrap();
24061 assert_eq!(memory.read_count(), 1);
24062 let call = planner_calls.last_call().unwrap();
24063 let prompt = &call.messages[0].content;
24064 assert!(prompt.contains("Available tools: none"), "{prompt}");
24065 assert!(!prompt.contains("- calculator ("), "{prompt}");
24066 }
24067
24068 #[tokio::test]
24069 async fn disambiguation_scope_failure_never_reaches_a_model_or_final_response() {
24070 let direct_memory = Arc::new(ReadFailingMemory::new(true));
24071 let direct_router = mock_with_response("unused");
24072 let direct_calls = direct_router.clone();
24073 let direct = scoped_disambiguation_agent(direct_memory.clone(), direct_router, true);
24074 let error = direct.build_disambiguation_context().await.unwrap_err();
24075 assert!(error.to_string().contains("simulated scope memory failure"));
24076 assert_eq!(direct_memory.read_count(), 1);
24077 assert_eq!(direct_calls.call_count(), 0);
24078
24079 let blocking_memory = Arc::new(ReadFailingMemory::new(true));
24080 let blocking_router = mock_with_response("unused");
24081 let blocking_calls = blocking_router.clone();
24082 let blocking = scoped_disambiguation_agent(blocking_memory.clone(), blocking_router, true);
24083 let error = blocking.chat("send it").await.unwrap_err();
24084 assert!(error.to_string().contains("simulated scope memory failure"));
24085 assert_eq!(blocking_memory.read_count(), 1);
24086 assert_eq!(blocking_calls.call_count(), 0);
24087
24088 let streaming_memory = Arc::new(ReadFailingMemory::new(true));
24089 let streaming_router = mock_with_response("unused");
24090 let streaming_calls = streaming_router.clone();
24091 let streaming =
24092 scoped_disambiguation_agent(streaming_memory.clone(), streaming_router, true);
24093 let (_, chunks, final_response) = collect_stream_events(&streaming, "send it").await;
24094 assert!(final_response.is_none());
24095 assert!(chunks.iter().any(|chunk| matches!(
24096 chunk,
24097 StreamChunk::Error { message } if message.contains("simulated scope memory failure")
24098 )));
24099 assert!(!chunks.iter().any(StreamChunk::is_done));
24100 assert_eq!(streaming_memory.read_count(), 1);
24101 assert_eq!(streaming_calls.call_count(), 0);
24102
24103 let skipped_memory = Arc::new(ReadFailingMemory::new(true));
24104 let skipped_router = mock_with_response("unused");
24105 let skipped = scoped_disambiguation_agent(skipped_memory.clone(), skipped_router, false);
24106 let context = skipped.build_disambiguation_context().await.unwrap();
24107 assert!(context.available_tools.is_empty());
24108 assert_eq!(skipped_memory.read_count(), 0);
24109 }
24110
24111 #[tokio::test]
24112 async fn test_stream_memory_write_failure_surfaces_as_error() {
24113 let yaml = r#"
24116name: TransitionOnToolCallAgent
24117system_prompt: "You are helpful."
24118streaming:
24119 enabled: true
24120states:
24121 initial: intake
24122 states:
24123 intake:
24124 prompt: "Intake"
24125 transitions:
24126 - to: billing
24127 guard:
24128 context:
24129 route:
24130 eq: billing
24131 billing:
24132 prompt: "Billing"
24133"#;
24134 let build = |fail_on_add: usize| {
24135 let mut tools = ai_agents_tools::ToolRegistry::new();
24136 tools
24137 .register(Arc::new(ai_agents_tools::CalculatorTool))
24138 .unwrap();
24139 let agent = AgentBuilder::from_yaml(yaml)
24140 .unwrap()
24141 .llm(Arc::new(mock_with_responses(vec![
24142 r#"{"tool": "calculator", "arguments": {"expression": "1+1"}}"#,
24143 "Billing answer",
24144 ])))
24145 .tools(tools)
24146 .memory(Arc::new(FailingMemory {
24147 messages: parking_lot::RwLock::new(Vec::new()),
24148 fail_on_add,
24149 adds: std::sync::atomic::AtomicUsize::new(0),
24150 }))
24151 .build()
24152 .unwrap();
24153 agent
24154 .set_context("route", serde_json::json!("billing"))
24155 .unwrap();
24156 agent
24157 };
24158
24159 let blocking = build(2).chat("compute").await;
24160 assert!(
24161 blocking.is_err(),
24162 "blocking must surface the memory failure"
24163 );
24164
24165 let (_, chunks, final_response) = collect_stream_events(&build(2), "compute").await;
24166 assert!(
24167 final_response.is_none(),
24168 "streaming must not finalize after a memory failure"
24169 );
24170 assert!(
24171 chunks.iter().any(|c| matches!(c, StreamChunk::Error { message } if message.contains("simulated memory failure"))),
24172 "streaming must surface the memory failure: {chunks:?}"
24173 );
24174
24175 assert!(build(usize::MAX).chat("compute").await.is_ok());
24177 }
24178
24179 #[tokio::test]
24180 async fn parity_tool_execution() {
24181 let build = || {
24182 calculator_agent_with(mock_with_responses(vec![
24183 r#"{"tool": "calculator", "arguments": {"expression": "2+2"}}"#,
24184 "The answer is 4.",
24185 ]))
24186 };
24187 let (blocking, _, chunks) = assert_blocking_streaming_parity(build, "What is 2+2?").await;
24188 assert_eq!(blocking.content, "The answer is 4.");
24189 assert!(
24190 chunks
24191 .iter()
24192 .any(|c| matches!(c, StreamChunk::ToolResult { .. }))
24193 );
24194 }
24195
24196 #[tokio::test]
24197 async fn parity_disambiguation_clarification() {
24198 let build = || {
24199 state_disambiguation_agent(
24200 vec![
24201 r#"{"is_ambiguous":true,"confidence":0.2,"ambiguity_type":"missing_target","reasoning":"target missing","what_is_unclear":["target"],"detected_language":"en"}"#,
24202 r#"{"question":"What should I send?","options":null}"#,
24203 ],
24204 true,
24205 None,
24206 true,
24207 )
24208 .0
24209 };
24210 let (blocking, _, _) = assert_blocking_streaming_parity(build, "Send it").await;
24211 assert_eq!(blocking.content, "What should I send?");
24212 }
24213
24214 #[tokio::test]
24215 async fn runtime_disambiguation_uses_configured_recent_history_projection() {
24216 let yaml = r#"
24217name: DisambiguationContextAgent
24218system_prompt: "Help."
24219llm:
24220 default: default
24221 router: router
24222disambiguation:
24223 enabled: true
24224 detection:
24225 llm: router
24226 context:
24227 recent_messages: 1
24228 include_state: false
24229 include_available_tools: false
24230states:
24231 initial: private_state
24232 states:
24233 private_state:
24234 prompt: "PRIVATE_STATE_PROMPT"
24235"#;
24236 let main = mock_with_responses(vec!["FIRST_MAIN_MARKER", "SECOND_MAIN_MARKER"]);
24237 let router = mock_with_responses(vec![
24238 r#"{"is_ambiguous":false,"confidence":0.9,"ambiguity_type":null,"reasoning":"clear","what_is_unclear":[],"detected_language":"en"}"#,
24239 r#"{"is_ambiguous":false,"confidence":0.9,"ambiguity_type":null,"reasoning":"clear","what_is_unclear":[],"detected_language":"en"}"#,
24240 ]);
24241 let router_calls = router.clone();
24242 let agent = AgentBuilder::from_yaml(yaml)
24243 .unwrap()
24244 .llm_alias("default", Arc::new(main))
24245 .llm_alias("router", Arc::new(router))
24246 .build()
24247 .unwrap();
24248
24249 agent.chat("FIRST_USER_MARKER").await.unwrap();
24250 agent.chat("SECOND_USER_MARKER").await.unwrap();
24251
24252 let calls = router_calls.call_history();
24253 assert_eq!(calls.len(), 2);
24254 let second_prompt = &calls[1].messages.last().unwrap().content;
24255 assert!(
24256 second_prompt.contains("FIRST_MAIN_MARKER"),
24257 "{second_prompt}"
24258 );
24259 assert!(
24260 !second_prompt.contains("FIRST_USER_MARKER"),
24261 "{second_prompt}"
24262 );
24263 assert!(
24264 !second_prompt.contains("PRIVATE_STATE_PROMPT"),
24265 "{second_prompt}"
24266 );
24267 }
24268
24269 #[tokio::test]
24270 async fn runtime_zero_history_preserves_pending_clarification_across_turns() {
24271 let yaml = r#"
24272name: ZeroHistoryPendingAgent
24273system_prompt: "Help."
24274llm:
24275 default: default
24276 router: router
24277disambiguation:
24278 enabled: true
24279 context:
24280 recent_messages: 0
24281 include_available_tools: false
24282"#;
24283 let main = mock_with_response("Final answer");
24284 let main_calls = main.clone();
24285 let router = mock_with_responses(vec![
24286 r#"{"is_ambiguous":true,"confidence":0.1,"ambiguity_type":"missing_target","reasoning":"target missing","what_is_unclear":["target"],"detected_language":"en"}"#,
24287 r#"{"question":"Who should receive it?","options":null}"#,
24288 r#"{"status":"answered","selected_option":null,"enriched_input":"Send it to Ada","resolved":{"recipient":"Ada"}}"#,
24289 ]);
24290 let router_calls = router.clone();
24291 let agent = AgentBuilder::from_yaml(yaml)
24292 .unwrap()
24293 .llm_alias("default", Arc::new(main))
24294 .llm_alias("router", Arc::new(router))
24295 .build()
24296 .unwrap();
24297
24298 let clarification = agent.chat("Send it").await.unwrap();
24299 assert_eq!(clarification.content, "Who should receive it?");
24300 let response = agent.chat("Ada").await.unwrap();
24301 assert_eq!(response.content, "Final answer");
24302 assert_eq!(router_calls.call_count(), 3);
24303 assert_eq!(main_calls.call_count(), 1);
24304 let calls = router_calls.call_history();
24305 let parse_prompt = &calls[2].messages.last().unwrap().content;
24306 assert!(
24307 parse_prompt.contains("Who should receive it?"),
24308 "{parse_prompt}"
24309 );
24310 }
24311
24312 #[tokio::test]
24313 async fn parity_reflection_enabled() {
24314 let yaml = r#"
24315name: ReflectionAgent
24316system_prompt: "You are careful."
24317reflection:
24318 enabled: true
24319 criteria:
24320 - "Is the answer helpful?"
24321"#;
24322 let build = || {
24323 AgentBuilder::from_yaml(yaml)
24324 .unwrap()
24325 .llm(Arc::new(mock_with_responses(vec![
24326 "Main answer",
24327 "OVERALL: PASS\nCONFIDENCE: 0.9",
24328 ])))
24329 .build()
24330 .unwrap()
24331 };
24332 let (blocking, streamed, _) = assert_blocking_streaming_parity(build, "hello").await;
24333 assert_eq!(blocking.content, "Main answer");
24334 assert!(metadata_keys(&streamed).contains("reflection"));
24335 }
24336
24337 fn state_reflection_agent(
24338 global_retries: u32,
24339 state_retries: u32,
24340 main: MockLLMProvider,
24341 evaluator: MockLLMProvider,
24342 ) -> RuntimeAgent {
24343 let yaml = format!(
24344 r#"
24345name: StateReflectionAgent
24346system_prompt: "You are careful."
24347llm:
24348 default: default
24349 router: evaluator
24350reflection:
24351 enabled: true
24352 evaluator_llm: evaluator
24353 max_retries: {global_retries}
24354 criteria:
24355 - "Global criterion"
24356states:
24357 initial: active
24358 states:
24359 active:
24360 prompt: "Handle the active state."
24361 reflection:
24362 enabled: true
24363 evaluator_llm: evaluator
24364 max_retries: {state_retries}
24365 criteria:
24366 - "State criterion"
24367"#
24368 );
24369 AgentBuilder::from_yaml(&yaml)
24370 .unwrap()
24371 .llm_alias("default", Arc::new(main))
24372 .llm_alias("evaluator", Arc::new(evaluator))
24373 .build()
24374 .unwrap()
24375 }
24376
24377 #[test]
24378 fn routing_log_reports_the_effective_state_reflection_mode() {
24379 let yaml = r#"
24380name: ReflectionLogAgent
24381system_prompt: "Be concise."
24382reflection:
24383 enabled: true
24384states:
24385 initial: active
24386 states:
24387 active:
24388 prompt: "Answer directly."
24389 reflection:
24390 enabled: false
24391"#;
24392 let agent = AgentBuilder::from_yaml(yaml)
24393 .unwrap()
24394 .llm(Arc::new(mock_with_response("answer")))
24395 .build()
24396 .unwrap();
24397 assert!(matches!(
24398 agent.routing_reflection_mode(),
24399 ReflectionMode::Disabled
24400 ));
24401 }
24402
24403 #[tokio::test]
24404 async fn state_reflection_zero_retries_overrides_global_limit() {
24405 let main = mock_with_response("First answer");
24406 let main_calls = main.clone();
24407 let evaluator = mock_with_response("OVERALL: FAIL\nCONFIDENCE: 0.1");
24408 let evaluator_calls = evaluator.clone();
24409 let agent = state_reflection_agent(2, 0, main, evaluator);
24410
24411 let response = agent.chat("hello").await.unwrap();
24412 let reflection = response
24413 .metadata
24414 .as_ref()
24415 .and_then(|metadata| metadata.get("reflection"))
24416 .expect("reflection metadata");
24417
24418 assert_eq!(response.content, "First answer");
24419 assert_eq!(reflection["attempts"], 1);
24420 assert_eq!(main_calls.call_count(), 1);
24421 assert_eq!(evaluator_calls.call_count(), 1);
24422 }
24423
24424 #[tokio::test]
24425 async fn state_reflection_retry_limit_overrides_zero_global_limit() {
24426 let main = mock_with_responses(vec!["First answer", "Second answer", "Third answer"]);
24427 let main_calls = main.clone();
24428 let evaluator = mock_with_responses(vec![
24429 "OVERALL: FAIL\nCONFIDENCE: 0.1",
24430 "OVERALL: FAIL\nCONFIDENCE: 0.1",
24431 "OVERALL: FAIL\nCONFIDENCE: 0.1",
24432 ]);
24433 let evaluator_calls = evaluator.clone();
24434 let agent = state_reflection_agent(0, 2, main, evaluator);
24435
24436 let response = agent.chat("hello").await.unwrap();
24437 let reflection = response
24438 .metadata
24439 .as_ref()
24440 .and_then(|metadata| metadata.get("reflection"))
24441 .expect("reflection metadata");
24442
24443 assert_eq!(response.content, "Third answer");
24444 assert_eq!(reflection["attempts"], 3);
24445 assert_eq!(main_calls.call_count(), 3);
24446 assert_eq!(evaluator_calls.call_count(), 3);
24447 }
24448
24449 #[tokio::test]
24450 async fn state_reflection_override_is_preserved_in_event_stream_metadata() {
24451 let main = mock_with_responses(vec!["First answer", "Second answer", "Third answer"]);
24452 let main_calls = main.clone();
24453 let evaluator = mock_with_responses(vec![
24454 "OVERALL: FAIL\nCONFIDENCE: 0.1",
24455 "OVERALL: FAIL\nCONFIDENCE: 0.1",
24456 "OVERALL: FAIL\nCONFIDENCE: 0.1",
24457 ]);
24458 let evaluator_calls = evaluator.clone();
24459 let agent = state_reflection_agent(0, 2, main, evaluator);
24460
24461 let (content, chunks, final_response) = collect_stream_events(&agent, "hello").await;
24462 let final_response = final_response.expect("successful stream Final");
24463 let reflection = final_response
24464 .metadata
24465 .as_ref()
24466 .and_then(|metadata| metadata.get("reflection"))
24467 .expect("reflection metadata");
24468
24469 assert_eq!(content, "Third answer");
24470 assert!(!chunks.iter().any(StreamChunk::is_error));
24471 assert_eq!(reflection["attempts"], 3);
24472 assert_eq!(main_calls.call_count(), 3);
24473 assert_eq!(evaluator_calls.call_count(), 3);
24474 }
24475
24476 #[tokio::test]
24477 async fn state_reflection_override_preserves_legacy_stream_completion() {
24478 use futures::StreamExt;
24479
24480 let main = mock_with_response("First answer");
24481 let main_calls = main.clone();
24482 let evaluator = mock_with_response("OVERALL: FAIL\nCONFIDENCE: 0.1");
24483 let evaluator_calls = evaluator.clone();
24484 let agent = state_reflection_agent(2, 0, main, evaluator);
24485 let mut stream = agent.chat_stream("hello").await.unwrap();
24486 let mut content = String::new();
24487 let mut done = 0;
24488 while let Some(chunk) = stream.next().await {
24489 match chunk {
24490 StreamChunk::Content { text } => content.push_str(&text),
24491 StreamChunk::Done {} => done += 1,
24492 StreamChunk::Error { message } => panic!("unexpected error: {message}"),
24493 _ => {}
24494 }
24495 }
24496
24497 assert_eq!(content, "First answer");
24498 assert_eq!(done, 1);
24499 assert_eq!(main_calls.call_count(), 1);
24500 assert_eq!(evaluator_calls.call_count(), 1);
24501 }
24502
24503 #[tokio::test]
24504 async fn reflection_evaluator_error_emits_no_event_stream_final() {
24505 let main = mock_with_response("First answer");
24506 let mut evaluator = MockLLMProvider::new("evaluator");
24507 evaluator.set_error("judge failed");
24508 let agent = state_reflection_agent(0, 2, main, evaluator);
24509
24510 let (_, chunks, final_response) = collect_stream_events(&agent, "hello").await;
24511
24512 assert!(final_response.is_none());
24513 assert!(chunks.iter().any(StreamChunk::is_error));
24514 assert!(!chunks.iter().any(StreamChunk::is_done));
24515 }
24516
24517 #[tokio::test]
24518 async fn parity_cot_hidden_thinking() {
24519 let yaml = r#"
24520name: CotHiddenAgent
24521system_prompt: "Think first."
24522reasoning:
24523 mode: cot
24524 output: hidden
24525"#;
24526 let build = || {
24527 AgentBuilder::from_yaml(yaml)
24528 .unwrap()
24529 .llm(Arc::new(mock_with_response(
24530 "<thinking>step by step</thinking>Visible answer",
24531 )))
24532 .build()
24533 .unwrap()
24534 };
24535 let (blocking, streamed, chunks) = assert_blocking_streaming_parity(build, "hello").await;
24536 assert_eq!(blocking.content, "Visible answer");
24537 assert_eq!(content_chunks(&chunks).concat(), streamed.content);
24539 }
24540
24541 fn content_chunks(chunks: &[StreamChunk]) -> Vec<String> {
24546 chunks
24547 .iter()
24548 .filter_map(|c| match c {
24549 StreamChunk::Content { text } => Some(text.clone()),
24550 _ => None,
24551 })
24552 .collect()
24553 }
24554
24555 fn reflection_auto_agent(main: MockLLMProvider, judge: MockLLMProvider) -> RuntimeAgent {
24556 let yaml = r#"
24557name: ReflectionAutoAgent
24558system_prompt: "You are careful."
24559llm:
24560 default: default
24561 router: router
24562reflection:
24563 enabled: auto
24564 evaluator_llm: router
24565 criteria:
24566 - "Is the answer helpful?"
24567"#;
24568 AgentBuilder::from_yaml(yaml)
24569 .unwrap()
24570 .llm_alias("default", Arc::new(main))
24571 .llm_alias("router", Arc::new(judge))
24572 .build()
24573 .unwrap()
24574 }
24575
24576 #[tokio::test]
24577 async fn test_stream_reflection_auto_buffers_and_calls_judge_once_per_iteration() {
24578 let judge = mock_with_responses(vec!["YES", "OVERALL: PASS\nCONFIDENCE: 0.9"]);
24579 let judge_calls = judge.clone();
24580 let agent = reflection_auto_agent(mock_with_response("Main answer one two"), judge);
24581
24582 let (content, chunks, final_response) = collect_stream_events(&agent, "hello").await;
24583
24584 assert!(!chunks.iter().any(StreamChunk::is_error), "{chunks:?}");
24585 assert_eq!(
24586 content_chunks(&chunks).len(),
24587 1,
24588 "auto reflection must buffer the main response: {chunks:?}"
24589 );
24590 assert_eq!(content, "Main answer one two");
24591 assert_eq!(final_response.expect("Final").content, content);
24592 assert_eq!(judge_calls.call_count(), 2);
24594 }
24595
24596 #[tokio::test]
24597 async fn test_stream_reflection_auto_rewrite_is_streamed() {
24598 let judge = mock_with_responses(vec![
24599 "YES",
24600 "OVERALL: FAIL\nCONFIDENCE: 0.1",
24601 "OVERALL: PASS\nCONFIDENCE: 0.9",
24602 ]);
24603 let agent = reflection_auto_agent(
24604 mock_with_responses(vec!["First attempt", "Improved answer"]),
24605 judge,
24606 );
24607
24608 let (content, chunks, final_response) = collect_stream_events(&agent, "hello").await;
24609
24610 assert!(!chunks.iter().any(StreamChunk::is_error), "{chunks:?}");
24611 assert_eq!(content_chunks(&chunks).len(), 1, "{chunks:?}");
24612 assert_eq!(
24613 content, "Improved answer",
24614 "the rewritten answer is what streams"
24615 );
24616 assert_eq!(final_response.expect("Final").content, "Improved answer");
24617 }
24618
24619 fn reasoning_agent(mode: &str, output: &str) -> RuntimeAgent {
24620 let yaml = format!(
24621 r#"
24622name: ReasoningStreamAgent
24623system_prompt: "Think first."
24624reasoning:
24625 mode: {mode}
24626 output: {output}
24627"#
24628 );
24629 AgentBuilder::from_yaml(&yaml)
24630 .unwrap()
24631 .llm(Arc::new(mock_with_response(
24632 "<thinking>step by step</thinking>Visible answer",
24633 )))
24634 .build()
24635 .unwrap()
24636 }
24637
24638 #[tokio::test]
24639 async fn test_stream_cot_hidden_emits_no_thinking_tags() {
24640 let agent = reasoning_agent("cot", "hidden");
24641 let (content, chunks, final_response) = collect_stream_events(&agent, "hello").await;
24642 assert!(!chunks.iter().any(StreamChunk::is_error), "{chunks:?}");
24643 assert_eq!(content_chunks(&chunks).len(), 1, "{chunks:?}");
24644 assert!(!content.contains("<thinking>"), "{content:?}");
24645 assert_eq!(content, "Visible answer");
24646 assert_eq!(final_response.expect("Final").content, content);
24647 }
24648
24649 #[tokio::test]
24650 async fn test_stream_cot_visible_matches_final_format() {
24651 let agent = reasoning_agent("cot", "visible");
24652 let (content, chunks, final_response) = collect_stream_events(&agent, "hello").await;
24653 assert!(!chunks.iter().any(StreamChunk::is_error), "{chunks:?}");
24654 assert!(content.starts_with("Thinking:"), "{content:?}");
24655 assert!(content.contains("Answer:\nVisible answer"), "{content:?}");
24656 assert_eq!(final_response.expect("Final").content, content);
24657 }
24658
24659 #[tokio::test]
24660 async fn test_stream_react_buffers() {
24661 let agent = reasoning_agent("react", "hidden");
24662 let (content, chunks, final_response) = collect_stream_events(&agent, "hello").await;
24663 assert!(!chunks.iter().any(StreamChunk::is_error), "{chunks:?}");
24664 assert_eq!(content_chunks(&chunks).len(), 1, "{chunks:?}");
24665 assert_eq!(content, "Visible answer");
24666 assert_eq!(final_response.expect("Final").content, content);
24667 }
24668
24669 #[tokio::test]
24670 async fn test_stream_plain_mode_still_streams_deltas() {
24671 let agent = AgentBuilder::new()
24672 .system_prompt("You are helpful.")
24673 .llm(Arc::new(mock_with_response("one two three")))
24674 .build()
24675 .unwrap();
24676 let (content, chunks, final_response) = collect_stream_events(&agent, "hello").await;
24677 assert!(
24678 content_chunks(&chunks).len() >= 2,
24679 "plain turns must keep token-level streaming: {chunks:?}"
24680 );
24681 assert_eq!(content, "one two three");
24682 assert_eq!(final_response.expect("Final").content, content);
24683 }
24684
24685 struct ActorProbeHooks {
24687 seen: parking_lot::Mutex<Option<crate::TurnActorContext>>,
24688 }
24689
24690 #[async_trait]
24691 impl AgentHooks for ActorProbeHooks {
24692 async fn on_message_received(&self, _input: &str) {
24693 *self.seen.lock() = current_turn_actor_context();
24694 }
24695 }
24696
24697 fn actor_probe_agent(hooks: Arc<ActorProbeHooks>) -> RuntimeAgent {
24698 let yaml = r#"
24699name: ActorStreamAgent
24700system_prompt: "You are helpful."
24701observability:
24702 enabled: true
24703 export:
24704 write_raw_events: true
24705"#;
24706 AgentBuilder::from_yaml(yaml)
24707 .unwrap()
24708 .llm(Arc::new(mock_with_response("Hello actor")))
24709 .hooks(hooks)
24710 .build()
24711 .unwrap()
24712 }
24713
24714 async fn collect_actor_stream_final(
24715 agent: &RuntimeAgent,
24716 input: &str,
24717 actor_context: crate::TurnActorContext,
24718 ) -> AgentResponse {
24719 use futures::StreamExt;
24720 let mut events = agent
24721 .chat_stream_events_with_actor_context(input, actor_context)
24722 .await
24723 .expect("stream opens");
24724 let mut final_response = None;
24725 while let Some(event) = events.next().await {
24726 match event {
24727 AgentStreamEvent::Final(response) => final_response = Some(response),
24728 AgentStreamEvent::Chunk(StreamChunk::Error { message }) => {
24729 panic!("unexpected stream error: {message}")
24730 }
24731 AgentStreamEvent::Chunk(_) => {}
24732 }
24733 }
24734 final_response.expect("Final")
24735 }
24736
24737 #[tokio::test]
24738 async fn test_stream_events_with_actor_context_scopes_actor_for_turn() {
24739 let hooks = Arc::new(ActorProbeHooks {
24740 seen: parking_lot::Mutex::new(None),
24741 });
24742 let agent = actor_probe_agent(Arc::clone(&hooks));
24743 let actor_context = crate::TurnActorContext::new().with_origin_actor("customer_42");
24744
24745 let final_response = collect_actor_stream_final(&agent, "hi", actor_context).await;
24746
24747 assert_eq!(final_response.content, "Hello actor");
24748 assert_eq!(
24749 hooks
24750 .seen
24751 .lock()
24752 .as_ref()
24753 .and_then(|context| context.effective_actor_id().map(str::to_string)),
24754 Some("customer_42".to_string()),
24755 "the actor context must be visible inside the streaming turn"
24756 );
24757 assert!(
24758 agent.actor_id().is_none(),
24759 "a turn-scoped actor must not mutate the global actor ID"
24760 );
24761 let events = agent.observability().unwrap().raw_events();
24762 assert!(
24763 events
24764 .iter()
24765 .any(|event| event.dimensions.get("actor") == Some(&"customer_42".to_string())),
24766 "observation events must carry the actor dimension"
24767 );
24768 }
24769
24770 #[tokio::test]
24771 async fn test_stream_events_with_actor_context_matches_blocking_actor_context() {
24772 let actor_context = crate::TurnActorContext::new()
24773 .with_origin_actor("customer_42")
24774 .with_sender_agent("coordinator");
24775
24776 let blocking_hooks = Arc::new(ActorProbeHooks {
24777 seen: parking_lot::Mutex::new(None),
24778 });
24779 let blocking_agent = actor_probe_agent(Arc::clone(&blocking_hooks));
24780 let blocking = blocking_agent
24781 .chat_with_actor_context("hi", actor_context.clone())
24782 .await
24783 .unwrap();
24784
24785 let streaming_hooks = Arc::new(ActorProbeHooks {
24786 seen: parking_lot::Mutex::new(None),
24787 });
24788 let streaming_agent = actor_probe_agent(Arc::clone(&streaming_hooks));
24789 let streamed =
24790 collect_actor_stream_final(&streaming_agent, "hi", actor_context.clone()).await;
24791
24792 assert_eq!(blocking.content, streamed.content);
24793 assert_eq!(metadata_keys(&blocking), metadata_keys(&streamed));
24794 assert_eq!(
24795 *blocking_hooks.seen.lock(),
24796 *streaming_hooks.seen.lock(),
24797 "both entry points must expose the same turn actor context"
24798 );
24799 assert_eq!(*streaming_hooks.seen.lock(), Some(actor_context));
24800 }
24801
24802 #[tokio::test]
24803 async fn test_stream_events_with_actor_context_releases_root_turn_on_drop() {
24804 use futures::StreamExt;
24805 let agent = AgentBuilder::new()
24806 .system_prompt("You are helpful.")
24807 .llm(Arc::new(mock_with_response("one two three")))
24808 .build()
24809 .unwrap();
24810 {
24811 let mut events = agent
24812 .chat_stream_events_with_actor_context(
24813 "hi",
24814 crate::TurnActorContext::new().with_origin_actor("customer_42"),
24815 )
24816 .await
24817 .unwrap();
24818 let _first = events.next().await;
24820 }
24821 let next = tokio::time::timeout(Duration::from_secs(5), agent.chat("next")).await;
24822 assert!(
24823 matches!(next, Ok(Ok(_))),
24824 "the root turn must be released when the actor stream is dropped: {next:?}"
24825 );
24826 }
24827
24828 fn skill_scope_agent_with_router(states: &str) -> (RuntimeAgent, MockLLMProvider) {
24829 let yaml = format!(
24830 r#"
24831name: SkillScopeAgent
24832system_prompt: "Route skills."
24833skills:
24834 - id: alpha
24835 description: "Alpha"
24836 trigger: "alpha"
24837 steps:
24838 - prompt: "alpha {{{{ user_input }}}}"
24839 - id: beta
24840 description: "Beta"
24841 trigger: "beta"
24842 steps:
24843 - prompt: "beta {{{{ user_input }}}}"
24844{states}
24845"#
24846 );
24847 let router = mock_with_response("none");
24848 let calls = router.clone();
24849 let agent = AgentBuilder::from_yaml(&yaml)
24850 .unwrap()
24851 .llm(Arc::new(router))
24852 .build()
24853 .unwrap();
24854 (agent, calls)
24855 }
24856
24857 fn skill_scope_agent(states: &str) -> RuntimeAgent {
24858 skill_scope_agent_with_router(states).0
24859 }
24860
24861 fn available_skill_ids(agent: &RuntimeAgent) -> Vec<String> {
24862 let mut ids: Vec<String> = agent
24863 .get_available_skills()
24864 .into_iter()
24865 .map(|skill| skill.id.clone())
24866 .collect();
24867 ids.sort();
24868 ids
24869 }
24870
24871 #[test]
24872 fn state_skill_scope_characterizes_empty_inheritance_and_unknown_ids() {
24873 assert_eq!(
24874 available_skill_ids(&skill_scope_agent("")),
24875 vec!["alpha", "beta"]
24876 );
24877 assert_eq!(
24878 available_skill_ids(&skill_scope_agent(
24879 "states:\n initial: current\n states:\n current:\n skills: []\n"
24880 )),
24881 vec!["alpha", "beta"]
24882 );
24883 assert_eq!(
24884 available_skill_ids(&skill_scope_agent(
24885 "states:\n initial: current\n states:\n current:\n prompt: current\n"
24886 )),
24887 vec!["alpha", "beta"]
24888 );
24889 assert!(
24890 available_skill_ids(&skill_scope_agent(
24891 "states:\n initial: current\n states:\n current:\n skills: [unknown]\n"
24892 ))
24893 .is_empty()
24894 );
24895 assert_eq!(
24896 available_skill_ids(&skill_scope_agent(
24897 "states:\n initial: parent\n states:\n parent:\n skills: [alpha]\n initial: child\n states:\n child:\n skills: []\n"
24898 )),
24899 vec!["alpha"]
24900 );
24901 assert_eq!(
24902 available_skill_ids(&skill_scope_agent(
24903 "states:\n initial: parent\n states:\n parent:\n skills: [alpha]\n initial: child\n states:\n child:\n skills: [beta]\n"
24904 )),
24905 vec!["alpha", "beta"]
24906 );
24907 assert_eq!(
24908 available_skill_ids(&skill_scope_agent(
24909 "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"
24910 )),
24911 vec!["alpha", "beta"]
24912 );
24913 }
24914
24915 #[tokio::test]
24916 async fn state_skill_scope_reaches_the_router_candidate_prompt() {
24917 let (inherited, inherited_calls) = skill_scope_agent_with_router(
24918 "states:\n initial: parent\n states:\n parent:\n skills: [alpha]\n initial: child\n states:\n child:\n skills: []\n",
24919 );
24920 assert!(
24921 inherited
24922 .select_skill_candidate("route")
24923 .await
24924 .unwrap()
24925 .is_none()
24926 );
24927 let inherited_call = inherited_calls.last_call().unwrap();
24928 let inherited_prompt = &inherited_call.messages[0].content;
24929 assert!(inherited_prompt.contains("- alpha:"));
24930 assert!(!inherited_prompt.contains("- beta:"));
24931
24932 let (fallback_all, fallback_calls) = skill_scope_agent_with_router(
24933 "states:\n initial: current\n states:\n current:\n skills: []\n",
24934 );
24935 assert!(
24936 fallback_all
24937 .select_skill_candidate("route")
24938 .await
24939 .unwrap()
24940 .is_none()
24941 );
24942 let fallback_call = fallback_calls.last_call().unwrap();
24943 let fallback_prompt = &fallback_call.messages[0].content;
24944 assert!(fallback_prompt.contains("- alpha:"));
24945 assert!(fallback_prompt.contains("- beta:"));
24946
24947 let (omitted, omitted_calls) = skill_scope_agent_with_router(
24948 "states:\n initial: current\n states:\n current:\n prompt: current\n",
24949 );
24950 assert!(
24951 omitted
24952 .select_skill_candidate("route")
24953 .await
24954 .unwrap()
24955 .is_none()
24956 );
24957 let omitted_call = omitted_calls.last_call().unwrap();
24958 let omitted_prompt = &omitted_call.messages[0].content;
24959 assert!(omitted_prompt.contains("- alpha:"));
24960 assert!(omitted_prompt.contains("- beta:"));
24961
24962 let (unknown, unknown_calls) = skill_scope_agent_with_router(
24963 "states:\n initial: current\n states:\n current:\n skills: [unknown]\n",
24964 );
24965 assert!(
24966 unknown
24967 .select_skill_candidate("route")
24968 .await
24969 .unwrap()
24970 .is_none()
24971 );
24972 assert_eq!(unknown_calls.call_count(), 0);
24973 }
24974
24975 #[tokio::test]
24976 async fn parity_skill_route() {
24977 let yaml = skills_with_parallel_transition_yaml(
24978 " max_speculative_llm_calls_per_turn: 2\n max_parallel_runtime_tasks: 2",
24979 "",
24980 );
24981 let build = || {
24982 build_skills_beside_transition_agent(
24983 &yaml,
24984 role_mocks(
24985 mock_with_response("Draft response"),
24986 mock_with_response("helper"),
24987 ),
24988 )
24989 };
24990 assert_blocking_streaming_parity(build, "please use helper").await;
24991 }
24992
24993 fn required_context_agent(mock: MockLLMProvider, default: bool) -> RuntimeAgent {
24995 let default_yaml = if default {
24996 " default:\n brief: fallback\n"
24997 } else {
24998 ""
24999 };
25000 let yaml = format!(
25001 "name: RequiredContextAgent\nsystem_prompt: 'Voice: {{{{ context.voice.brief }}}}'\ncontext:\n voice:\n type: runtime\n required: true\n{default_yaml}"
25002 );
25003 AgentBuilder::from_yaml(&yaml)
25004 .unwrap()
25005 .llm(Arc::new(mock))
25006 .build()
25007 .unwrap()
25008 }
25009
25010 struct CountingContextProvider {
25011 marker: &'static str,
25012 calls: std::sync::atomic::AtomicUsize,
25013 }
25014
25015 #[async_trait]
25016 impl ContextProvider for CountingContextProvider {
25017 async fn get(&self, _key: &str, _current_context: &Value) -> Result<Value> {
25018 let call = self.calls.fetch_add(1, Ordering::SeqCst) + 1;
25019 Ok(serde_json::json!({"call": call, "marker": self.marker}))
25020 }
25021 }
25022
25023 struct FailOnceContextProvider {
25024 attempts: std::sync::atomic::AtomicUsize,
25025 }
25026
25027 #[async_trait]
25028 impl ContextProvider for FailOnceContextProvider {
25029 async fn get(&self, _key: &str, _current_context: &Value) -> Result<Value> {
25030 if self.attempts.fetch_add(1, Ordering::SeqCst) == 0 {
25031 return Err(AgentError::Other("context initialization failed".into()));
25032 }
25033 Ok(serde_json::json!({"brief": "ready"}))
25034 }
25035 }
25036
25037 fn session_context_agent(
25038 provider: Arc<CountingContextProvider>,
25039 refresh: &str,
25040 ) -> RuntimeAgent {
25041 let yaml = format!(
25042 "name: SessionContextAgent\nsystem_prompt: 'Call: {{{{ context.session_data.call }}}}'\ncontext:\n session_data:\n type: callback\n name: counter\n refresh: {refresh}\n"
25043 );
25044 let agent = AgentBuilder::from_yaml(&yaml)
25045 .unwrap()
25046 .llm(Arc::new(mock_with_response("ok")))
25047 .build()
25048 .unwrap();
25049 agent.register_context_provider("counter", provider);
25050 agent
25051 }
25052
25053 #[tokio::test]
25054 async fn session_context_characterizes_reset_and_restore_lifecycle() {
25055 let original_provider = Arc::new(CountingContextProvider {
25056 marker: "original",
25057 calls: std::sync::atomic::AtomicUsize::new(0),
25058 });
25059 let original = session_context_agent(original_provider.clone(), "per_session");
25060 original.chat("first").await.unwrap();
25061 original.chat("second").await.unwrap();
25062 assert_eq!(original_provider.calls.load(Ordering::SeqCst), 1);
25063 assert_eq!(original.get_context()["session_data"]["call"], 1);
25064 assert_eq!(original.get_context()["session_data"]["marker"], "original");
25065
25066 original.reset().await.unwrap();
25067 original.chat("after reset").await.unwrap();
25068 assert_eq!(original_provider.calls.load(Ordering::SeqCst), 1);
25069 assert_eq!(original.get_context()["session_data"]["call"], 1);
25070 assert_eq!(original.get_context()["session_data"]["marker"], "original");
25071 let snapshot = original.save_state().await.unwrap();
25072
25073 let fresh_provider = Arc::new(CountingContextProvider {
25074 marker: "fresh",
25075 calls: std::sync::atomic::AtomicUsize::new(0),
25076 });
25077 let fresh = session_context_agent(fresh_provider.clone(), "per_session");
25078 fresh.restore_state(snapshot.clone()).await.unwrap();
25079 fresh.chat("fresh restore").await.unwrap();
25080 assert_eq!(fresh_provider.calls.load(Ordering::SeqCst), 1);
25081 assert_eq!(fresh.get_context()["session_data"]["call"], 1);
25082 assert_eq!(fresh.get_context()["session_data"]["marker"], "fresh");
25083
25084 let warm_provider = Arc::new(CountingContextProvider {
25085 marker: "warm",
25086 calls: std::sync::atomic::AtomicUsize::new(0),
25087 });
25088 let warm = session_context_agent(warm_provider.clone(), "per_session");
25089 warm.chat("warmup").await.unwrap();
25090 assert_eq!(warm_provider.calls.load(Ordering::SeqCst), 1);
25091 warm.restore_state(snapshot).await.unwrap();
25092 warm.chat("warm restore").await.unwrap();
25093 assert_eq!(warm_provider.calls.load(Ordering::SeqCst), 1);
25094 assert_eq!(warm.get_context()["session_data"]["call"], 1);
25095 assert_eq!(warm.get_context()["session_data"]["marker"], "original");
25096 }
25097
25098 #[tokio::test]
25099 async fn once_context_is_not_refreshed_by_later_turns_or_reset() {
25100 let provider = Arc::new(CountingContextProvider {
25101 marker: "once",
25102 calls: std::sync::atomic::AtomicUsize::new(0),
25103 });
25104 let agent = session_context_agent(provider.clone(), "once");
25105
25106 agent.chat("first").await.unwrap();
25107 agent.chat("second").await.unwrap();
25108 agent.reset().await.unwrap();
25109 agent.chat("after reset").await.unwrap();
25110
25111 assert_eq!(provider.calls.load(Ordering::SeqCst), 1);
25112 assert_eq!(agent.get_context()["session_data"]["marker"], "once");
25113 }
25114
25115 #[tokio::test]
25116 async fn per_turn_context_refreshes_after_initialization_and_restore() {
25117 let original_provider = Arc::new(CountingContextProvider {
25118 marker: "original-per-turn",
25119 calls: std::sync::atomic::AtomicUsize::new(0),
25120 });
25121 let original = session_context_agent(original_provider.clone(), "per_turn");
25122 original.chat("first").await.unwrap();
25123 assert_eq!(original_provider.calls.load(Ordering::SeqCst), 2);
25124 assert_eq!(original.get_context()["session_data"]["call"], 2);
25125 original.chat("second").await.unwrap();
25126 assert_eq!(original_provider.calls.load(Ordering::SeqCst), 3);
25127 original.reset().await.unwrap();
25128 original.chat("after reset").await.unwrap();
25129 assert_eq!(original_provider.calls.load(Ordering::SeqCst), 4);
25130 let snapshot = original.save_state().await.unwrap();
25131
25132 let fresh_provider = Arc::new(CountingContextProvider {
25133 marker: "fresh-per-turn",
25134 calls: std::sync::atomic::AtomicUsize::new(0),
25135 });
25136 let fresh = session_context_agent(fresh_provider.clone(), "per_turn");
25137 fresh.restore_state(snapshot.clone()).await.unwrap();
25138 fresh.chat("fresh restore").await.unwrap();
25139 assert_eq!(fresh_provider.calls.load(Ordering::SeqCst), 2);
25140 assert_eq!(
25141 fresh.get_context()["session_data"]["marker"],
25142 "fresh-per-turn"
25143 );
25144 assert_eq!(fresh.get_context()["session_data"]["call"], 2);
25145
25146 let warm_provider = Arc::new(CountingContextProvider {
25147 marker: "warm-per-turn",
25148 calls: std::sync::atomic::AtomicUsize::new(0),
25149 });
25150 let warm = session_context_agent(warm_provider.clone(), "per_turn");
25151 warm.chat("warmup").await.unwrap();
25152 assert_eq!(warm_provider.calls.load(Ordering::SeqCst), 2);
25153 warm.restore_state(snapshot).await.unwrap();
25154 warm.chat("warm restore").await.unwrap();
25155 assert_eq!(warm_provider.calls.load(Ordering::SeqCst), 3);
25156 assert_eq!(
25157 warm.get_context()["session_data"]["marker"],
25158 "warm-per-turn"
25159 );
25160 assert_eq!(warm.get_context()["session_data"]["call"], 3);
25161 }
25162
25163 #[tokio::test]
25164 async fn test_context_initialization_retries_after_failure() {
25165 let mock = mock_with_response("Voice response");
25166 let calls = mock.clone();
25167 let yaml = "name: CallbackAgent\nsystem_prompt: 'Voice: {{ context.voice.brief }}'\ncontext:\n voice:\n type: callback\n name: flaky\n";
25168 let agent = AgentBuilder::from_yaml(yaml)
25169 .unwrap()
25170 .llm(Arc::new(mock))
25171 .build()
25172 .unwrap();
25173 let provider = Arc::new(FailOnceContextProvider {
25174 attempts: std::sync::atomic::AtomicUsize::new(0),
25175 });
25176 agent.register_context_provider("flaky", provider.clone());
25177
25178 assert!(agent.chat("first").await.is_err());
25179 assert_eq!(calls.call_count(), 0);
25180 assert_eq!(
25181 agent.chat("second").await.unwrap().content,
25182 "Voice response"
25183 );
25184 assert_eq!(provider.attempts.load(Ordering::SeqCst), 2);
25185 assert_eq!(calls.call_count(), 1);
25186 }
25187
25188 #[tokio::test]
25189 async fn test_required_context_blocks_chat_until_supplied_and_after_removal() {
25190 let mock = mock_with_response("Voice response");
25191 let calls = mock.clone();
25192 let agent = required_context_agent(mock, false);
25193
25194 let error = agent.chat("first").await.unwrap_err();
25195 assert!(
25196 error
25197 .to_string()
25198 .contains("Required context 'voice' not provided")
25199 );
25200 assert_eq!(calls.call_count(), 0);
25201 assert!(agent.memory.get_messages(None).await.unwrap().is_empty());
25202
25203 agent
25204 .set_context("voice.brief", serde_json::json!("ready"))
25205 .unwrap();
25206 assert_eq!(
25207 agent.chat("second").await.unwrap().content,
25208 "Voice response"
25209 );
25210 assert_eq!(calls.call_count(), 1);
25211
25212 agent.remove_context("voice");
25213 let error = agent.chat("third").await.unwrap_err();
25214 assert!(
25215 error
25216 .to_string()
25217 .contains("Required context 'voice' not provided")
25218 );
25219 assert_eq!(calls.call_count(), 1);
25220 }
25221
25222 #[tokio::test]
25223 async fn test_required_context_default_satisfies_presence_check() {
25224 let mock = mock_with_response("Fallback response");
25225 let calls = mock.clone();
25226 let agent = required_context_agent(mock, true);
25227
25228 assert_eq!(
25229 agent.chat("hello").await.unwrap().content,
25230 "Fallback response"
25231 );
25232 assert_eq!(
25233 agent.context_manager().get_path("voice.brief"),
25234 Some(serde_json::json!("fallback"))
25235 );
25236 assert_eq!(calls.call_count(), 1);
25237 }
25238
25239 #[tokio::test]
25240 async fn test_required_context_blocks_legacy_stream_before_model_call() {
25241 use futures::StreamExt;
25242
25243 let mock = mock_with_response("Voice response");
25244 let calls = mock.clone();
25245 let agent = required_context_agent(mock, false);
25246 let mut stream = agent.chat_stream("first").await.unwrap();
25247 assert!(
25248 matches!(stream.next().await, Some(StreamChunk::Error { message }) if message.contains("Required context 'voice' not provided"))
25249 );
25250 assert!(stream.next().await.is_none());
25251 drop(stream);
25252 assert_eq!(calls.call_count(), 0);
25253 assert!(agent.memory.get_messages(None).await.unwrap().is_empty());
25254
25255 agent
25256 .set_context("voice.brief", serde_json::json!("ready"))
25257 .unwrap();
25258 assert_eq!(
25259 agent.chat("second").await.unwrap().content,
25260 "Voice response"
25261 );
25262 }
25263
25264 #[tokio::test]
25265 async fn test_required_context_blocks_event_streams_without_final() {
25266 use futures::StreamExt;
25267
25268 for actor_scoped in [false, true] {
25269 let mock = mock_with_response("Voice response");
25270 let calls = mock.clone();
25271 let agent = required_context_agent(mock, false);
25272 let mut events = if actor_scoped {
25273 agent
25274 .chat_stream_events_with_actor_context(
25275 "first",
25276 crate::TurnActorContext::new().with_origin_actor("caller"),
25277 )
25278 .await
25279 .unwrap()
25280 } else {
25281 agent.chat_stream_events("first").await.unwrap()
25282 };
25283 assert!(
25284 matches!(events.next().await, Some(AgentStreamEvent::Chunk(StreamChunk::Error { message })) if message.contains("Required context 'voice' not provided"))
25285 );
25286 assert!(events.next().await.is_none());
25287 drop(events);
25288 assert_eq!(calls.call_count(), 0);
25289 assert!(agent.memory.get_messages(None).await.unwrap().is_empty());
25290
25291 agent
25292 .set_context("voice.brief", serde_json::json!("ready"))
25293 .unwrap();
25294 assert_eq!(
25295 agent.chat("second").await.unwrap().content,
25296 "Voice response"
25297 );
25298 }
25299 }
25300}