1use async_trait::async_trait;
2use futures::stream::{Stream, StreamExt};
3use parking_lot::RwLock;
4use serde_json::Value;
5use std::collections::{HashMap, HashSet};
6use std::future::Future;
7use std::pin::Pin;
8use std::sync::atomic::{AtomicBool, AtomicU64, Ordering};
9use std::sync::{Arc, Weak};
10use std::time::{Duration, Instant};
11use tracing::{debug, error, info, instrument, warn};
12
13const DISAMBIGUATION_STATE_GENERATION_KEY: &str = "_runtime.disambiguation_state_generation";
14const MAX_TOOL_FALLBACK_HOPS: usize = 16;
15
16pub(crate) type RootTurnGate = Arc<tokio::sync::Mutex<()>>;
18
19pub(crate) type RootTurnGateIdentityStack = Arc<[RootTurnGate]>;
21
22tokio::task_local! {
23 static RUNTIME_GATE_IDENTITY_STACK: RootTurnGateIdentityStack;
24}
25
26pub(crate) fn current_runtime_gate_identity_stack() -> RootTurnGateIdentityStack {
28 RUNTIME_GATE_IDENTITY_STACK
29 .try_with(Arc::clone)
30 .unwrap_or_default()
31}
32
33pub(crate) async fn scope_runtime_gate_identity_stack<F, T>(
35 identity_stack: &RootTurnGateIdentityStack,
36 future: F,
37) -> T
38where
39 F: Future<Output = T>,
40{
41 RUNTIME_GATE_IDENTITY_STACK
42 .scope(Arc::clone(identity_stack), future)
43 .await
44}
45
46pub(crate) type ToolResourceLocks = Arc<RwLock<HashMap<String, Weak<tokio::sync::Mutex<()>>>>>;
48
49struct ToolResourceGuards {
53 guards: Vec<tokio::sync::OwnedMutexGuard<()>>,
54 locks: ToolResourceLocks,
55}
56
57struct RootTurnAdmission {
61 guard: tokio::sync::OwnedMutexGuard<()>,
62 identity_stack: RootTurnGateIdentityStack,
63}
64
65#[derive(Clone)]
66struct StoredSessionRestore {
67 snapshot: AgentSnapshot,
68 metadata: Option<ai_agents_core::SessionMetadata>,
69}
70
71struct RuntimeSessionRestorePoint {
72 snapshot: AgentSnapshot,
73 metadata: ai_agents_core::SessionMetadata,
74 actor_id: Option<String>,
75 session_id: Option<String>,
76}
77
78impl Drop for ToolResourceGuards {
79 fn drop(&mut self) {
80 self.guards.clear();
81 self.locks.write().retain(|_, lock| lock.strong_count() > 0);
82 }
83}
84
85#[derive(Clone)]
89struct RuntimeSafetySnapshot {
90 version: u64,
91 emergency_deny: bool,
92 tool_security: ToolSecurityEngine,
93 tool_scope_override: Option<Vec<String>>,
94}
95
96#[derive(Clone, Copy)]
100struct ToolDecisionVersions {
101 policy: u64,
102 registry: u64,
103 runtime_control: u64,
104 state: Option<u64>,
105}
106
107#[derive(Clone, Debug, Default)]
111struct ToolFallbackState {
112 visited_canonical_ids: Vec<String>,
113}
114
115impl ToolFallbackState {
116 fn rejection_reason(&self, canonical_id: &str) -> Option<String> {
120 if self
121 .visited_canonical_ids
122 .iter()
123 .any(|visited| visited == canonical_id)
124 {
125 return Some(format!(
126 "Tool fallback cycle detected at '{canonical_id}' after [{}]",
127 self.visited_canonical_ids.join(" -> ")
128 ));
129 }
130 if self.visited_canonical_ids.len() > MAX_TOOL_FALLBACK_HOPS {
131 return Some(format!(
132 "Tool fallback chain exceeds the maximum of {MAX_TOOL_FALLBACK_HOPS} hops"
133 ));
134 }
135 None
136 }
137
138 fn with_current(mut self, canonical_id: String) -> Self {
142 self.visited_canonical_ids.push(canonical_id);
143 self
144 }
145
146 fn final_rejection_reason(
150 &self,
151 admitted_canonical_id: &str,
152 final_canonical_id: &str,
153 ) -> Option<String> {
154 if admitted_canonical_id == final_canonical_id {
155 return None;
156 }
157 if self
158 .visited_canonical_ids
159 .iter()
160 .any(|visited| visited == final_canonical_id)
161 {
162 return Some(format!(
163 "Tool fallback cycle detected after final resolution changed '{admitted_canonical_id}' to '{final_canonical_id}'"
164 ));
165 }
166 Some(format!(
167 "Tool canonical target changed after initial admission from '{admitted_canonical_id}' to '{final_canonical_id}'"
168 ))
169 }
170}
171
172#[derive(Clone, Copy, Debug)]
176struct ValidatedToolTimeout {
177 timer: Duration,
178 deadline_delta: chrono::Duration,
179}
180
181struct AvailableToolIdsSnapshot {
185 tool_ids: Vec<String>,
186 state_generation: Option<u64>,
187}
188
189#[derive(Clone)]
193struct ToolApprovalBinding {
194 canonical_id: String,
195 arguments: Value,
196 confirmation_required: bool,
197 policy_version: u64,
198 runtime_control_version: u64,
199 state_generation: Option<u64>,
200 reviewed_tool: Arc<dyn ai_agents_core::Tool>,
201}
202
203fn merge_approved_record(record: &mut Option<ToolApprovalRecord>) {
207 if record
208 .as_ref()
209 .is_some_and(|record| matches!(record.status, ToolApprovalStatus::Modified))
210 {
211 return;
212 }
213 *record = Some(ToolApprovalRecord {
214 status: ToolApprovalStatus::Approved,
215 reason: None,
216 modified_arguments: None,
217 });
218}
219
220impl ToolApprovalBinding {
221 fn is_stale(
223 &self,
224 canonical_id: &str,
225 arguments: &Value,
226 confirmation_required: bool,
227 versions: ToolDecisionVersions,
228 resolved_tool: &Arc<dyn ai_agents_core::Tool>,
229 ) -> bool {
230 self.canonical_id != canonical_id
231 || self.arguments != *arguments
232 || self.confirmation_required != confirmation_required
233 || self.policy_version != versions.policy
234 || self.runtime_control_version != versions.runtime_control
235 || self.state_generation != versions.state
236 || !Arc::ptr_eq(&self.reviewed_tool, resolved_tool)
237 }
238}
239
240use crate::turn_context::{current_turn_actor_context, scope_actor_context};
241
242use ai_agents_context::{ContextManager, ContextProvider, TemplateRenderer};
243use ai_agents_core::traits::storage::StorageCapability;
244use ai_agents_core::{
245 AgentError, AgentSnapshot, AgentStorage, ChatMessage, FinishReason, LLMChunk, LLMError,
246 LLMFeature, LLMProvider, LLMResponse, LLMToolDefinition, LLMToolRequest, PermissionOutcome,
247 Result, ToolActorContext, ToolApprovalRecord, ToolApprovalStatus, ToolCallClassification,
248 ToolCallSource, ToolCancellationToken, ToolChoice, ToolExecutionContext, ToolExecutionLimits,
249 ToolExecutionRecord, ToolExecutionRequest, ToolInvoker, ToolPolicyDecisionRecord, ToolResult,
250 ToolSafetyMetadata, decode_native_tool_call_markers, encode_native_tool_call_markers,
251 encode_native_tool_result_marker, inspect_native_history, native_readable_projection,
252};
253use ai_agents_disambiguation::{
254 AmbiguityDetectionResult, ClarificationObserver, ClarificationParseFuture,
255 ClarificationQuestion, ClarificationQuestionFuture, ConfirmationParseFuture,
256 DisambiguationConfig, DisambiguationContext, DisambiguationManager, DisambiguationResult,
257};
258use ai_agents_hitl::{
259 ApprovalHandler, ApprovalResolvedOutcome, ApprovalResult, ApprovalTrigger, HITLCheckResult,
260 HITLEngine, RejectAllHandler, TimeoutAction,
261};
262use ai_agents_hooks::{AgentHooks, NoopHooks};
263use ai_agents_llm::LLMRegistry;
264use ai_agents_memory::{
265 CompressResult, EvictionReason, Memory, MemoryBudgetEvent, MemoryCompressEvent,
266 MemoryEvictEvent, MemoryTokenBudget, OverflowStrategy,
267};
268use ai_agents_observability::{
269 EventStatus, EventType, ObservabilityManager, ObservationPurpose, SpanContext,
270 current_observation_context, new_session_id as new_observation_session_id,
271 resolve_language_from_context, with_observation_context, with_observation_purpose,
272};
273use ai_agents_process::{
274 ProcessData, ProcessProcessor, ProcessPurposeHint, ProcessStageFuture, ProcessStageObserver,
275};
276use ai_agents_reasoning::{
277 CriterionResult, EvaluationResult, Plan, PlanAction, PlanStatus, PlanStep, ReasoningConfig,
278 ReasoningMetadata, ReasoningMode, ReasoningOutput, ReflectionAttempt, ReflectionConfig,
279 ReflectionMetadata, StepFailureAction,
280};
281use ai_agents_recovery::{
282 ByRoleFilter, ContextOverflowAction, FilterConfig, KeepRecentFilter, LLMFailureAction,
283 MessageFilter, RecoveryManager, SkipPatternFilter, ToolFailureAction,
284};
285use ai_agents_relationships::RelationshipManager;
286use ai_agents_skills::{SkillDefinition, SkillExecutor, SkillRouter};
287use ai_agents_state::{
288 PromptMode, StateAction, StateMachine, StateMachineSnapshot, StateTransitionEvent, Transition,
289 TransitionContext, TransitionEvaluator, TransitionTiming, evaluate_guard,
290};
291use ai_agents_storage::{StorageConfig as StorageStorageConfig, create_storage};
292use ai_agents_tools::{
293 CommandRunner, ConditionEvaluator, DiagnosticsProvider, EvaluationContext, LLMGetter,
294 MAX_TOOL_TIMEOUT_MS, QuestionHandler, SecurityCheckResult, TodoItem, ToolCallRecord,
295 ToolRegistry, ToolSecurityConfig, ToolSecurityEngine,
296};
297
298use super::{
299 Agent, AgentInfo, AgentResponse, AgentStreamEvent, ParallelToolsConfig, StreamChunk,
300 StreamingConfig, ToolCall,
301};
302use crate::optimization::{
303 AwaitBeforeNextTurn, BackgroundMaintenanceQueue, BackgroundOverflowPolicy, MainResponseDraft,
304 MaintenanceMode, MaintenanceSequenceKey, RuntimeBranch, RuntimeBranchResult,
305 RuntimeBranchStatus, RuntimeCommitBehavior, RuntimeConfig, RuntimeOptimizationKind,
306 RuntimeTaskPriority, RuntimeTaskPurpose, ScheduledBranchSet, SkillCandidate,
307 StreamingDraftResult, TransitionCandidate, TurnBranchScheduler, TurnOptimizationContext,
308};
309use crate::spec::StorageConfig;
310
311enum ToolCallOutcome {
313 Continue,
315 TransitionFired,
317 Rejected(AgentResponse),
319}
320
321#[derive(Clone)]
322struct MainToolProtocol {
323 choice: Option<ToolChoice>,
324 tool_ids: Vec<String>,
325 definitions: Vec<LLMToolDefinition>,
326}
327
328struct MainProviderResponse {
329 response: LLMResponse,
330 used_native_tools: bool,
331}
332
333enum MainStreamSource {
338 Stream(Box<dyn Stream<Item = std::result::Result<LLMChunk, LLMError>> + Unpin + Send>),
339 StaticResponse(String),
340}
341
342#[derive(Clone)]
344struct ActiveNativeExchange {
345 exchange_id: String,
346 call_ids: Vec<String>,
347}
348
349struct CommittedTextResponse<'a> {
353 processed_input: &'a str,
354 input_context: &'a HashMap<String, Value>,
355 answer: String,
356 reasoning_mode: ReasoningMode,
357 auto_detected: bool,
358 iterations: u32,
359 thinking_content: Option<String>,
360 all_tool_calls: Vec<ToolCall>,
361}
362
363struct AgentResponseParts {
367 content: String,
368 all_tool_calls: Vec<ToolCall>,
369 reasoning_mode: ReasoningMode,
370 auto_detected: bool,
371 iterations: u32,
372 thinking: Option<String>,
373 reflection_metadata: Option<ReflectionMetadata>,
374}
375
376type RuntimeStreamTerminalSlot = Arc<RwLock<Option<AgentResponse>>>;
380
381fn new_runtime_stream_terminal_slot() -> RuntimeStreamTerminalSlot {
385 Arc::new(RwLock::new(None))
386}
387
388fn record_runtime_stream_final(slot: &RuntimeStreamTerminalSlot, response: AgentResponse) {
392 *slot.write() = Some(response);
393}
394
395#[derive(Clone, Copy)]
396struct DisambiguationOwnership {
397 epoch: u64,
398 state_generation: Option<u64>,
399}
400
401enum SkillRouteResult {
403 NoMatch,
405 Response { skill_id: String, content: String },
407 NeedsClarification {
409 response: AgentResponse,
410 ownership: Option<DisambiguationOwnership>,
411 },
412}
413
414enum ParallelTransitionSelection {
416 Candidate(TransitionCandidate),
418 NoMatch,
420 ReservationExhausted,
422}
423
424enum DisambiguationDispatch {
430 Proceed(String),
432 Terminal(AgentResponse),
434 RecheckSkill {
436 skill_id: String,
437 enriched_input: String,
438 disambiguation_epoch: u64,
439 state_generation: Option<u64>,
440 },
441}
442
443enum PostLoopResult {
444 NoTransition(String),
446 Transitioned { content: String, regenerated: bool },
449 NeedsRedispatch,
452}
453
454struct AppliedPostLoop {
456 content: String,
457 transitioned: bool,
458 regenerated: bool,
460}
461
462struct StateTransitionReservation<'a> {
463 reserved: &'a AtomicBool,
464}
465
466impl Drop for StateTransitionReservation<'_> {
467 fn drop(&mut self) {
468 self.reserved.store(false, Ordering::SeqCst);
469 }
470}
471
472struct RootTurnCleanup<'a> {
473 agent: &'a RuntimeAgent,
474}
475
476impl<'a> RootTurnCleanup<'a> {
477 fn new(agent: &'a RuntimeAgent) -> Self {
478 Self { agent }
479 }
480}
481
482impl Drop for RootTurnCleanup<'_> {
483 fn drop(&mut self) {
484 self.agent.end_root_turn();
485 }
486}
487
488#[derive(Debug)]
490struct RuntimeControlState {
491 snapshot_guard: RwLock<()>,
493 version: AtomicU64,
495 emergency_deny: Arc<AtomicBool>,
497 tool_security_override: RwLock<Option<ToolSecurityEngine>>,
499 tool_scope_override: RwLock<Option<Vec<String>>>,
501}
502
503impl Default for RuntimeControlState {
504 fn default() -> Self {
505 Self {
506 snapshot_guard: RwLock::new(()),
507 version: AtomicU64::new(1),
508 emergency_deny: Arc::new(AtomicBool::new(false)),
509 tool_security_override: RwLock::new(None),
510 tool_scope_override: RwLock::new(None),
511 }
512 }
513}
514
515#[derive(Clone)]
517pub struct RuntimeControlHandle {
518 state: Arc<RuntimeControlState>,
519}
520
521impl RuntimeControlHandle {
522 pub fn version(&self) -> u64 {
524 self.state.version.load(Ordering::SeqCst)
525 }
526
527 fn bump(&self) -> u64 {
528 self.state.version.fetch_add(1, Ordering::SeqCst) + 1
529 }
530
531 pub fn set_tool_security(&self, config: ToolSecurityConfig) -> u64 {
533 self.try_set_tool_security(config)
534 .expect("invalid tool security configuration")
535 }
536
537 pub fn try_set_tool_security(&self, config: ToolSecurityConfig) -> Result<u64> {
539 config.validate()?;
540 let _guard = self.state.snapshot_guard.write();
541 let generation = self.bump();
542 *self.state.tool_security_override.write() = Some(
543 ToolSecurityEngine::new_with_policy_version(config, generation),
544 );
545 Ok(generation)
546 }
547
548 pub fn clear_tool_security_override(&self) -> u64 {
550 let _guard = self.state.snapshot_guard.write();
551 *self.state.tool_security_override.write() = None;
552 self.bump()
553 }
554
555 pub fn set_tool_scope(&self, tool_ids: Vec<String>) -> u64 {
557 let _guard = self.state.snapshot_guard.write();
558 *self.state.tool_scope_override.write() = Some(tool_ids);
559 self.bump()
560 }
561
562 pub fn clear_tool_scope_override(&self) -> u64 {
564 let _guard = self.state.snapshot_guard.write();
565 *self.state.tool_scope_override.write() = None;
566 self.bump()
567 }
568
569 pub fn set_emergency_deny(&self, enabled: bool) -> u64 {
571 let _guard = self.state.snapshot_guard.write();
572 self.state.emergency_deny.store(enabled, Ordering::SeqCst);
573 self.bump()
574 }
575
576 pub fn cancel_all(&self) -> u64 {
578 self.set_emergency_deny(true)
579 }
580}
581
582pub struct RuntimeAgent {
583 info: AgentInfo,
584 llm_registry: Arc<LLMRegistry>,
585 memory: Arc<dyn Memory>,
586 tools: Arc<ToolRegistry>,
587 skills: Vec<SkillDefinition>,
588 skill_router: Option<SkillRouter>,
589 skill_executor: Option<SkillExecutor>,
590 base_system_prompt: String,
591 max_iterations: u32,
592 iteration_count: RwLock<u32>,
593 max_context_tokens: u32,
594 memory_token_budget: Option<MemoryTokenBudget>,
595 recovery_manager: RecoveryManager,
596 tool_security: ToolSecurityEngine,
597 process_processor: Option<ProcessProcessor>,
598 message_filters: RwLock<HashMap<String, Arc<dyn MessageFilter>>>,
599 state_machine: Option<Arc<StateMachine>>,
600 transition_evaluator: Option<Arc<dyn TransitionEvaluator>>,
601 context_manager: Arc<ContextManager>,
602 template_renderer: TemplateRenderer,
603 tool_call_history: RwLock<Vec<ToolCallRecord>>,
604 parallel_tools: ParallelToolsConfig,
605 streaming: StreamingConfig,
606 hooks: Arc<dyn AgentHooks>,
607 hitl_engine: Option<HITLEngine>,
608 approval_handler: Arc<dyn ApprovalHandler>,
609 storage_config: StorageConfig,
610 storage: RwLock<Option<Arc<dyn AgentStorage>>>,
611 storage_init: tokio::sync::Mutex<()>,
612 reasoning_config: ReasoningConfig,
613 reflection_config: ReflectionConfig,
614 disambiguation_manager: Option<DisambiguationManager>,
615 disambiguation_epoch: AtomicU64,
617 disambiguation_admission: tokio::sync::RwLock<()>,
619 state_transition_reserved: AtomicBool,
621 persona_manager: Option<Arc<ai_agents_persona::PersonaManager>>,
623 pending_skill_id: RwLock<Option<String>>,
627 current_plan: RwLock<Option<Plan>>,
628 declared_tool_ids: Option<Vec<String>>,
630 context_initialized: AtomicBool,
632 spawner: Option<Arc<crate::spawner::AgentSpawner>>,
634 spawner_registry: Option<Arc<crate::spawner::AgentRegistry>>,
636 redispatch_depth: RwLock<u32>,
639 active_turn_context: RwLock<Option<TurnOptimizationContext>>,
641 root_user_message_committed: AtomicBool,
643 active_native_exchanges: RwLock<Vec<ActiveNativeExchange>>,
645 actor_id: RwLock<Option<String>>,
647 fact_store: RwLock<Option<Arc<ai_agents_facts::FactStore>>>,
649 fact_extractor: RwLock<Option<Arc<dyn ai_agents_facts::FactExtractor>>>,
652 actor_facts_cache: Arc<RwLock<HashMap<String, Vec<ai_agents_core::KeyFact>>>>,
654 messages_since_extraction: Arc<RwLock<usize>>,
656 actor_memory_config: Option<ai_agents_facts::ActorMemoryConfig>,
658 facts_config: Option<ai_agents_facts::FactsConfig>,
660 session_metadata: RwLock<ai_agents_core::SessionMetadata>,
662 current_session_id: RwLock<Option<String>>,
664 relationship_manager: Option<Arc<RelationshipManager>>,
666 observability_manager: Option<Arc<ObservabilityManager>>,
668 runtime_config: RuntimeConfig,
670 background_maintenance: Arc<BackgroundMaintenanceQueue>,
672 resource_locks: ToolResourceLocks,
674 runtime_control: Arc<RuntimeControlState>,
676 root_turn_gate: RootTurnGate,
678}
679
680impl std::fmt::Debug for RuntimeAgent {
681 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
682 f.debug_struct("RuntimeAgent")
683 .field("info", &self.info)
684 .field("base_system_prompt", &self.base_system_prompt)
685 .field("max_iterations", &self.max_iterations)
686 .field("skills_count", &self.skills.len())
687 .field("max_context_tokens", &self.max_context_tokens)
688 .field("has_state_machine", &self.state_machine.is_some())
689 .field("parallel_tools", &self.parallel_tools)
690 .field("streaming", &self.streaming)
691 .field("has_hooks", &true)
692 .field("has_hitl", &self.hitl_engine.is_some())
693 .field("storage_type", &self.storage_config.storage_type())
694 .field("reasoning_mode", &self.reasoning_config.mode)
695 .field("reflection_enabled", &self.reflection_config.enabled)
696 .field("declared_tool_ids", &self.declared_tool_ids)
697 .field("has_persona", &self.persona_manager.is_some())
698 .field("has_observability", &self.observability_manager.is_some())
699 .finish_non_exhaustive()
700 }
701}
702
703struct ObservabilityClarificationObserver;
704
705impl ClarificationObserver for ObservabilityClarificationObserver {
706 fn observe_question<'a>(
708 &'a self,
709 future: ClarificationQuestionFuture<'a>,
710 ) -> ClarificationQuestionFuture<'a> {
711 Box::pin(async move {
712 with_observation_purpose(ObservationPurpose::DisambiguationClarification, future).await
713 })
714 }
715
716 fn observe_parse<'a>(
718 &'a self,
719 future: ClarificationParseFuture<'a>,
720 ) -> ClarificationParseFuture<'a> {
721 Box::pin(async move {
722 with_observation_purpose(ObservationPurpose::DisambiguationClarification, future).await
723 })
724 }
725
726 fn observe_confirmation_parse<'a>(
728 &'a self,
729 future: ConfirmationParseFuture<'a>,
730 ) -> ConfirmationParseFuture<'a> {
731 Box::pin(async move {
732 with_observation_purpose(ObservationPurpose::DisambiguationClarification, future).await
733 })
734 }
735}
736
737struct ObservabilityProcessStageObserver;
738
739impl ProcessStageObserver for ObservabilityProcessStageObserver {
740 fn observe<'a>(
742 &'a self,
743 hint: ProcessPurposeHint,
744 future: ProcessStageFuture<'a>,
745 ) -> ProcessStageFuture<'a> {
746 Box::pin(async move {
747 with_observation_purpose(observation_purpose_for_process(hint), future).await
748 })
749 }
750}
751
752struct RegistryLLMGetter {
753 registry: Arc<LLMRegistry>,
754}
755
756impl LLMGetter for RegistryLLMGetter {
757 fn get_llm(&self, alias: &str) -> Option<Arc<dyn LLMProvider>> {
758 self.registry.get(alias).ok()
759 }
760}
761
762impl RuntimeAgent {
763 #[allow(clippy::too_many_arguments)]
765 pub fn new(
766 info: AgentInfo,
767 llm_registry: Arc<LLMRegistry>,
768 memory: Arc<dyn Memory>,
769 tools: Arc<ToolRegistry>,
770 skills: Vec<SkillDefinition>,
771 system_prompt: String,
772 max_iterations: u32,
773 ) -> Self {
774 let (skill_router, skill_executor) = if !skills.is_empty() {
775 let router_llm = llm_registry.router().ok();
776 let router = router_llm.map(|llm| SkillRouter::new(llm, skills.clone()));
777 let executor = SkillExecutor::new(llm_registry.clone(), tools.clone());
778 (router, Some(executor))
779 } else {
780 (None, None)
781 };
782
783 let context_manager =
784 ContextManager::new(HashMap::new(), info.name.clone(), info.version.clone());
785
786 Self {
787 info,
788 llm_registry,
789 memory,
790 tools,
791 skills,
792 skill_router,
793 skill_executor,
794 base_system_prompt: system_prompt,
795 max_iterations,
796 iteration_count: RwLock::new(0),
797 max_context_tokens: 128000,
798 memory_token_budget: None,
799 recovery_manager: RecoveryManager::default(),
800 tool_security: ToolSecurityEngine::default(),
801 process_processor: None,
802 message_filters: RwLock::new(HashMap::new()),
803 state_machine: None,
804 transition_evaluator: None,
805 context_manager: Arc::new(context_manager),
806 template_renderer: TemplateRenderer::new(),
807 tool_call_history: RwLock::new(Vec::new()),
808 parallel_tools: ParallelToolsConfig::default(),
809 streaming: StreamingConfig::default(),
810 hooks: Arc::new(NoopHooks),
811 hitl_engine: None,
812 approval_handler: Arc::new(RejectAllHandler::new()),
813 storage_config: StorageConfig::default(),
814 storage: RwLock::new(None),
815 storage_init: tokio::sync::Mutex::new(()),
816 reasoning_config: ReasoningConfig::default(),
817 reflection_config: ReflectionConfig::default(),
818 disambiguation_manager: None,
819 disambiguation_epoch: AtomicU64::new(0),
820 disambiguation_admission: tokio::sync::RwLock::new(()),
821 state_transition_reserved: AtomicBool::new(false),
822 persona_manager: None,
823 pending_skill_id: RwLock::new(None),
824 current_plan: RwLock::new(None),
825 declared_tool_ids: None,
826 context_initialized: AtomicBool::new(false),
827 spawner: None,
828 spawner_registry: None,
829 redispatch_depth: RwLock::new(0),
830 active_turn_context: RwLock::new(None),
831 root_user_message_committed: AtomicBool::new(false),
832 active_native_exchanges: RwLock::new(Vec::new()),
833 actor_id: RwLock::new(None),
834 fact_store: RwLock::new(None),
835 fact_extractor: RwLock::new(None),
836 actor_facts_cache: Arc::new(RwLock::new(HashMap::new())),
837 messages_since_extraction: Arc::new(RwLock::new(0)),
838 actor_memory_config: None,
839 facts_config: None,
840 session_metadata: RwLock::new(ai_agents_core::SessionMetadata::default()),
841 current_session_id: RwLock::new(None),
842 relationship_manager: None,
843 observability_manager: None,
844 runtime_config: RuntimeConfig::default(),
845 background_maintenance: Arc::new(BackgroundMaintenanceQueue::default()),
846 resource_locks: new_tool_resource_locks(),
847 runtime_control: Arc::new(RuntimeControlState::default()),
848 root_turn_gate: Arc::new(tokio::sync::Mutex::new(())),
849 }
850 }
851
852 pub fn with_declared_tool_ids(mut self, ids: Option<Vec<String>>) -> Self {
853 self.declared_tool_ids = ids;
854 self
855 }
856
857 pub fn with_storage_config(mut self, config: StorageConfig) -> Self {
858 self.storage_config = config;
859 self
860 }
861
862 pub fn with_storage(self, storage: Arc<dyn AgentStorage>) -> Self {
863 *self.storage.write() = Some(storage);
864 self
865 }
866
867 pub(crate) fn with_shared_resource_locks(mut self, locks: ToolResourceLocks) -> Self {
868 self.resource_locks = locks;
869 self
870 }
871
872 pub fn with_reasoning(mut self, config: ReasoningConfig) -> Self {
873 self.reasoning_config = config;
874 self
875 }
876
877 pub fn with_reflection(mut self, config: ReflectionConfig) -> Self {
878 self.reflection_config = config;
879 self
880 }
881
882 pub fn with_relationships(mut self, manager: Arc<RelationshipManager>) -> Self {
884 self.relationship_manager = Some(manager);
885 self
886 }
887
888 pub fn with_observability(mut self, manager: Arc<ObservabilityManager>) -> Self {
890 self.observability_manager = Some(manager);
891 self
892 }
893
894 pub fn with_runtime_config(mut self, config: RuntimeConfig) -> Self {
896 let max_tasks = config.optimization.post_turn.max_background_tasks;
897 self.background_maintenance = Arc::new(BackgroundMaintenanceQueue::new(max_tasks));
898 self.runtime_config = config;
899 self
900 }
901
902 pub fn runtime_config(&self) -> &RuntimeConfig {
904 &self.runtime_config
905 }
906
907 pub async fn flush_background_tasks(&self) -> Result<()> {
909 self.background_maintenance.flush_all().await
910 }
911
912 pub async fn flush_background_tasks_for_actor(&self, actor_id: &str) -> Result<()> {
914 self.background_maintenance.flush_scope(actor_id).await
915 }
916
917 pub async fn flush_background_tasks_for_purpose(
919 &self,
920 purpose: RuntimeTaskPurpose,
921 ) -> Result<()> {
922 self.background_maintenance.flush_purpose(purpose).await
923 }
924
925 pub async fn flush_background_tasks_for_actor_purpose(
927 &self,
928 actor_id: &str,
929 purpose: RuntimeTaskPurpose,
930 ) -> Result<()> {
931 self.background_maintenance
932 .flush_scope_purpose(actor_id, purpose)
933 .await
934 }
935
936 pub async fn shutdown_background_tasks(&self) -> Result<()> {
938 self.flush_background_tasks().await
939 }
940
941 pub fn observability(&self) -> Option<Arc<ObservabilityManager>> {
943 self.observability_manager.clone()
944 }
945
946 async fn export_observability_if_configured(&self) {
948 let Some(manager) = self.observability_manager.as_ref() else {
949 return;
950 };
951 let export = &manager.config().export;
952 if !export.write_report && !export.write_raw_events {
953 return;
954 }
955 if let Err(error) = manager.export().await {
956 warn!(error = %error, "Observability export failed");
957 }
958 }
959
960 pub fn relationship_manager(&self) -> Option<Arc<RelationshipManager>> {
962 self.relationship_manager.clone()
963 }
964
965 fn current_turn_actor_context(&self) -> Option<crate::TurnActorContext> {
966 current_turn_actor_context()
967 }
968
969 fn effective_actor_id(&self) -> Option<String> {
970 self.current_turn_actor_context()
971 .and_then(|ctx| ctx.effective_actor_id().map(|id| id.to_string()))
972 .or_else(|| self.actor_id.read().clone())
973 }
974
975 fn effective_origin_actor_id(&self) -> Option<String> {
976 self.current_turn_actor_context()
977 .and_then(|ctx| ctx.origin_actor_id.clone())
978 .or_else(|| self.actor_id.read().clone())
979 }
980
981 fn record_session_actor_if_needed(&self) {
982 if let Some(actor_id) = self.effective_origin_actor_id() {
983 let mut meta = self.session_metadata.write();
984 meta.actor_id = Some(actor_id.clone());
985 if !meta.actors.iter().any(|a| a == &actor_id) {
986 meta.actors.push(actor_id);
987 }
988 }
989 }
990
991 fn outbound_actor_context(&self) -> crate::TurnActorContext {
992 let mut context = self.current_turn_actor_context().unwrap_or_default();
993 if context.origin_actor_id.is_none() {
994 context.origin_actor_id = self.effective_origin_actor_id();
995 }
996 context.sender_agent_id = Some(self.info.id.clone());
997 context
998 }
999
1000 fn observation_session_id(&self) -> Option<String> {
1002 let mut current = self.current_session_id.write();
1003 if current.is_none() {
1004 *current = Some(new_observation_session_id());
1005 }
1006 current.clone()
1007 }
1008
1009 fn build_observation_context(&self, actor_id: Option<String>) -> Option<SpanContext> {
1011 let manager = self.observability_manager.as_ref()?;
1012 let context = self.build_context_with_overlays();
1013 let language = resolve_language_from_context(manager.config(), &context);
1014 let context = current_observation_context()
1015 .map(|parent| parent.child_for_agent(self.info.id.clone()).with_new_turn())
1016 .unwrap_or_else(|| SpanContext::new_root(self.info.id.clone()));
1017 Some(
1018 context
1019 .with_actor(actor_id.or_else(|| self.effective_actor_id()))
1020 .with_session(self.observation_session_id())
1021 .with_state(self.current_state())
1022 .with_language(Some(language)),
1023 )
1024 }
1025
1026 fn current_runtime_observation_context(
1028 &self,
1029 purpose: ObservationPurpose,
1030 ) -> Option<SpanContext> {
1031 let manager = self.observability_manager.as_ref()?;
1032 let context = self.build_context_with_overlays();
1033 let language = resolve_language_from_context(manager.config(), &context);
1034 let mut observation = current_observation_context()
1035 .unwrap_or_else(|| SpanContext::new_root(self.info.id.clone()));
1036 observation.agent_id = self.info.id.clone();
1037 observation.actor_id = self.effective_actor_id();
1038 observation.session_id = self.observation_session_id();
1039 observation.state = self.current_state();
1040 observation.language = Some(language);
1041 observation.purpose = purpose;
1042 Some(observation)
1043 }
1044
1045 async fn observe_purpose<F, T>(&self, purpose: ObservationPurpose, future: F) -> T
1047 where
1048 F: Future<Output = T>,
1049 {
1050 if let Some(context) = self.current_runtime_observation_context(purpose) {
1051 with_observation_context(context, future).await
1052 } else {
1053 future.await
1054 }
1055 }
1056
1057 fn chat_with_actor_context_boxed<'a>(
1061 &'a self,
1062 input: &'a str,
1063 actor_context: crate::TurnActorContext,
1064 ) -> Pin<Box<dyn Future<Output = Result<AgentResponse>> + Send + 'a>> {
1065 Box::pin(async move {
1066 let RootTurnAdmission {
1067 guard,
1068 identity_stack,
1069 } = self.acquire_root_turn().await?;
1070 let result = scope_runtime_gate_identity_stack(&identity_stack, async move {
1071 let actor_id = actor_context.effective_actor_id().map(str::to_string);
1072 let run = async move {
1073 scope_actor_context(
1074 actor_context,
1075 Box::pin(async move { self.run_loop(input).await }),
1076 )
1077 .await
1078 };
1079 let result = if let Some(context) = self.build_observation_context(actor_id) {
1080 with_observation_context(context, run).await
1081 } else {
1082 run.await
1083 };
1084 self.export_observability_if_configured().await;
1085 result
1086 })
1087 .await;
1088 drop(guard);
1089 result
1090 })
1091 }
1092
1093 async fn acquire_root_turn(&self) -> Result<RootTurnAdmission> {
1095 let gate_identity = Arc::clone(&self.root_turn_gate);
1096 let current_identity_stack = current_runtime_gate_identity_stack();
1097 if current_identity_stack
1098 .iter()
1099 .any(|owned_gate| Arc::ptr_eq(owned_gate, &gate_identity))
1100 {
1101 return Err(AgentError::Other(format!(
1102 "RuntimeAgent '{}' rejected reentrant root turn ownership",
1103 self.info.id
1104 )));
1105 }
1106 let guard = Arc::clone(&gate_identity).lock_owned().await;
1107 let mut identity_stack = Vec::with_capacity(current_identity_stack.len() + 1);
1111 identity_stack.extend(current_identity_stack.iter().cloned());
1112 identity_stack.push(gate_identity);
1113 Ok(RootTurnAdmission {
1114 guard,
1115 identity_stack: identity_stack.into(),
1116 })
1117 }
1118
1119 pub async fn chat_with_actor_context(
1123 &self,
1124 input: &str,
1125 actor_context: crate::TurnActorContext,
1126 ) -> Result<AgentResponse> {
1127 self.chat_with_actor_context_boxed(input, actor_context)
1128 .await
1129 }
1130
1131 pub async fn chat_as_actor(&self, actor_id: &str, input: &str) -> Result<AgentResponse> {
1133 let actor_context = crate::TurnActorContext::new().with_origin_actor(actor_id);
1134 self.chat_with_actor_context(input, actor_context).await
1135 }
1136
1137 pub async fn load_actor_relationship(&self) -> Result<()> {
1139 self.maybe_load_actor_relationship().await;
1140 Ok(())
1141 }
1142
1143 pub async fn update_relationship_dimension(
1145 &self,
1146 dimension: &str,
1147 delta: f64,
1148 reason: Option<&str>,
1149 ) -> Result<ai_agents_relationships::DimensionChange> {
1150 self.update_relationship_dimension_for_perspective(
1151 ai_agents_relationships::RelationshipPerspective::AgentToActor,
1152 dimension,
1153 delta,
1154 reason,
1155 )
1156 .await
1157 }
1158
1159 pub async fn update_relationship_dimension_for_perspective(
1163 &self,
1164 perspective: ai_agents_relationships::RelationshipPerspective,
1165 dimension: &str,
1166 delta: f64,
1167 reason: Option<&str>,
1168 ) -> Result<ai_agents_relationships::DimensionChange> {
1169 let manager = self
1170 .relationship_manager
1171 .as_ref()
1172 .ok_or_else(|| AgentError::Config("Relationship memory is not configured".into()))?;
1173 let actor_id = self.effective_actor_id().ok_or_else(|| {
1174 AgentError::Config("No actor ID set. Use set_actor_id() first".into())
1175 })?;
1176 let change = manager.update_dimension_for_perspective(
1177 &actor_id,
1178 perspective,
1179 dimension,
1180 delta,
1181 1.0,
1182 reason.unwrap_or("manual relationship update"),
1183 )?;
1184 self.persist_actor_relationship(&actor_id).await?;
1185 info!(
1186 actor_id = %actor_id,
1187 perspective = %change.perspective,
1188 dimension = %change.dimension,
1189 delta = change.delta,
1190 current = change.current,
1191 "relationship updated manually"
1192 );
1193 self.hooks
1194 .on_relationship_change(&actor_id, std::slice::from_ref(&change))
1195 .await;
1196 Ok(change)
1197 }
1198
1199 pub fn reasoning_config(&self) -> &ReasoningConfig {
1200 &self.reasoning_config
1201 }
1202
1203 pub fn reflection_config(&self) -> &ReflectionConfig {
1204 &self.reflection_config
1205 }
1206
1207 pub fn with_facts_config(
1210 mut self,
1211 actor_memory_config: Option<ai_agents_facts::ActorMemoryConfig>,
1212 facts_config: Option<ai_agents_facts::FactsConfig>,
1213 ) -> Self {
1214 self.actor_memory_config = actor_memory_config;
1215 self.facts_config = facts_config;
1216 self
1217 }
1218
1219 pub fn with_facts(
1222 mut self,
1223 store: Arc<ai_agents_facts::FactStore>,
1224 extractor: Option<Arc<dyn ai_agents_facts::FactExtractor>>,
1225 actor_memory_config: Option<ai_agents_facts::ActorMemoryConfig>,
1226 facts_config: Option<ai_agents_facts::FactsConfig>,
1227 ) -> Self {
1228 *self.fact_store.write() = Some(store);
1229 *self.fact_extractor.write() = extractor;
1230 self.actor_memory_config = actor_memory_config;
1231 self.facts_config = facts_config;
1232 self
1233 }
1234
1235 pub fn fact_store(&self) -> Option<Arc<ai_agents_facts::FactStore>> {
1237 self.fact_store.read().clone()
1238 }
1239
1240 pub fn actor_id(&self) -> Option<String> {
1242 self.actor_id.read().clone()
1243 }
1244
1245 pub fn set_actor_id(&self, actor_id: &str) -> ai_agents_core::Result<()> {
1247 *self.actor_id.write() = Some(actor_id.to_string());
1248 {
1249 let mut meta = self.session_metadata.write();
1250 meta.actor_id = Some(actor_id.to_string());
1251 if !meta.actors.iter().any(|a| a == actor_id) {
1252 meta.actors.push(actor_id.to_string());
1253 }
1254 }
1255 Ok(())
1256 }
1257
1258 pub fn clear_actor_id(&self) {
1260 *self.actor_id.write() = None;
1261 self.session_metadata.write().actor_id = None;
1262 }
1263
1264 pub fn set_user_id(&self, user_id: &str) -> ai_agents_core::Result<()> {
1266 self.set_actor_id(user_id)
1267 }
1268
1269 pub async fn load_actor_memory(&self) -> ai_agents_core::Result<()> {
1271 let actor_id = match self.effective_actor_id() {
1272 Some(id) => id,
1273 None => return Ok(()),
1274 };
1275
1276 let store_opt = self.fact_store.read().clone();
1277 if let Some(store) = store_opt {
1278 let facts = store.get_facts(&actor_id).await?;
1279 let count = facts.len();
1280 self.actor_facts_cache
1281 .write()
1282 .insert(actor_id.clone(), facts);
1283 self.hooks.on_actor_memory_loaded(&actor_id, count).await;
1284 tracing::debug!("loaded {} facts for actor {}", count, actor_id);
1285 }
1286
1287 Ok(())
1288 }
1289
1290 async fn maybe_load_actor_memory(&self) {
1292 let Some(actor_id) = self.effective_actor_id() else {
1293 return;
1294 };
1295 if self.actor_facts_cache.read().contains_key(&actor_id) {
1296 return;
1297 }
1298 let _ = self.load_actor_memory().await;
1299 }
1300
1301 async fn pre_turn_session_lifecycle(&self) {
1303 if *self.redispatch_depth.read() > 0 {
1304 return;
1305 }
1306 self.resolve_actor_id_from_context();
1307 self.await_background_before_next_turn().await;
1308 self.record_session_actor_if_needed();
1309 self.maybe_load_actor_memory().await;
1310 self.maybe_load_actor_relationship().await;
1311 *self.messages_since_extraction.write() += 1;
1312 }
1313
1314 async fn post_turn_session_lifecycle(&self) -> Result<()> {
1316 if *self.redispatch_depth.read() > 0 {
1317 return Ok(());
1318 }
1319 *self.messages_since_extraction.write() += 1;
1320 self.run_post_turn_maintenance().await
1321 }
1322
1323 fn begin_root_turn(&self) {
1325 if *self.redispatch_depth.read() == 0 {
1326 let mut guard = self.active_turn_context.write();
1327 if guard.is_none() {
1328 self.root_user_message_committed
1329 .store(false, Ordering::SeqCst);
1330 self.active_native_exchanges.write().clear();
1331 let max_calls = self
1332 .runtime_config
1333 .optimization
1334 .max_speculative_llm_calls_per_turn;
1335 *guard = Some(TurnOptimizationContext::new(
1336 String::new(),
1337 HashMap::new(),
1338 max_calls,
1339 ));
1340 }
1341 }
1342 }
1343
1344 fn update_active_turn_context(
1345 &self,
1346 processed_input: &str,
1347 input_context: HashMap<String, Value>,
1348 ) {
1349 if *self.redispatch_depth.read() > 0 {
1350 return;
1351 }
1352 let max_calls = self
1353 .runtime_config
1354 .optimization
1355 .max_speculative_llm_calls_per_turn;
1356 let mut guard = self.active_turn_context.write();
1357 match guard.as_mut() {
1358 Some(context) => {
1359 context.processed_input = processed_input.to_string();
1360 context.input_context = input_context;
1361 context.max_speculative_llm_calls = max_calls;
1362 }
1363 None => {
1364 *guard = Some(TurnOptimizationContext::new(
1365 processed_input,
1366 input_context,
1367 max_calls,
1368 ));
1369 }
1370 }
1371 }
1372
1373 async fn commit_root_user_message(&self, processed_input: &str) -> Result<()> {
1375 if *self.redispatch_depth.read() > 0 {
1376 return Ok(());
1377 }
1378 if !self
1379 .root_user_message_committed
1380 .swap(true, Ordering::SeqCst)
1381 {
1382 self.memory
1383 .add_message(ChatMessage::user(processed_input))
1384 .await?;
1385 if let Some(context) = self.active_turn_context.write().as_mut() {
1386 context.mark_user_message_committed();
1387 }
1388 }
1389 Ok(())
1390 }
1391
1392 fn end_root_turn(&self) {
1394 if *self.redispatch_depth.read() == 0 {
1395 self.root_user_message_committed
1396 .store(false, Ordering::SeqCst);
1397 *self.active_turn_context.write() = None;
1398 self.active_native_exchanges.write().clear();
1399 }
1400 }
1401
1402 fn reserve_active_speculative_llm_call(&self, kind: RuntimeOptimizationKind) -> bool {
1403 self.begin_root_turn();
1404 let mut guard = self.active_turn_context.write();
1405 let Some(context) = guard.as_mut() else {
1406 return false;
1407 };
1408 context.reserve_speculative_llm_call_for(kind)
1409 }
1410
1411 fn branch_context_preview(&self) -> String {
1412 let context = self.build_context_with_overlays();
1413 let mut value = serde_json::to_string_pretty(&context).unwrap_or_else(|_| "{}".to_string());
1414 const MAX_CONTEXT_PREVIEW_CHARS: usize = 2048;
1415 if value.chars().count() > MAX_CONTEXT_PREVIEW_CHARS {
1416 value = value
1417 .chars()
1418 .take(MAX_CONTEXT_PREVIEW_CHARS)
1419 .collect::<String>();
1420 value.push_str("...");
1421 }
1422 value
1423 }
1424
1425 async fn await_background_before_next_turn(&self) {
1427 let optimization = &self.runtime_config.optimization;
1428 if !optimization.enabled {
1429 return;
1430 }
1431 let actor_id = self.effective_actor_id();
1432 let post = &optimization.post_turn;
1433 self.await_background_task(
1434 post.facts.await_before_next_turn,
1435 RuntimeTaskPurpose::PostTurnFacts,
1436 actor_id.as_deref(),
1437 "facts",
1438 )
1439 .await;
1440 self.await_background_task(
1441 post.relationships.await_before_next_turn,
1442 RuntimeTaskPurpose::PostTurnRelationship,
1443 actor_id.as_deref(),
1444 "relationships",
1445 )
1446 .await;
1447 }
1448
1449 async fn await_background_task(
1450 &self,
1451 policy: AwaitBeforeNextTurn,
1452 purpose: RuntimeTaskPurpose,
1453 actor_id: Option<&str>,
1454 label: &str,
1455 ) {
1456 match policy {
1457 AwaitBeforeNextTurn::Never => {}
1458 AwaitBeforeNextTurn::Always => {
1459 if let Err(error) = self.flush_background_tasks_for_purpose(purpose).await {
1460 warn!(label = label, error = %error, "background maintenance flush failed");
1461 }
1462 }
1463 AwaitBeforeNextTurn::SameActor => {
1464 if let Some(actor_id) = actor_id
1465 && let Err(error) = self
1466 .flush_background_tasks_for_actor_purpose(actor_id, purpose)
1467 .await
1468 {
1469 warn!(label = label, actor_id = %actor_id, error = %error, "actor background maintenance flush failed");
1470 }
1471 }
1472 }
1473 }
1474
1475 async fn run_post_turn_maintenance(&self) -> Result<()> {
1477 let optimization = &self.runtime_config.optimization;
1478 if !optimization.enabled {
1479 self.auto_extract_facts().await;
1480 self.auto_update_relationship().await;
1481 return Ok(());
1482 }
1483
1484 let facts_mode = effective_maintenance_mode(
1485 optimization.post_turn.facts.mode,
1486 optimization.parallel_post_turn_memory,
1487 );
1488 let relationships_mode = effective_maintenance_mode(
1489 optimization.post_turn.relationships.mode,
1490 optimization.parallel_post_turn_memory,
1491 );
1492
1493 match (facts_mode, relationships_mode) {
1494 (MaintenanceMode::InlineSerial, MaintenanceMode::InlineSerial) => {
1495 self.auto_extract_facts().await;
1496 self.auto_update_relationship().await;
1497 }
1498 (MaintenanceMode::InlineParallel, MaintenanceMode::InlineParallel) => {
1499 let facts = self.auto_extract_facts();
1500 let relationships = self.auto_update_relationship();
1501 tokio::join!(facts, relationships);
1502 }
1503 (MaintenanceMode::Background, MaintenanceMode::Background) => {
1504 self.schedule_facts_background().await?;
1505 self.schedule_relationship_background().await?;
1506 }
1507 (MaintenanceMode::Background, MaintenanceMode::InlineParallel)
1508 | (MaintenanceMode::Background, MaintenanceMode::InlineSerial) => {
1509 self.schedule_facts_background().await?;
1510 self.auto_update_relationship().await;
1511 }
1512 (MaintenanceMode::InlineParallel, MaintenanceMode::Background)
1513 | (MaintenanceMode::InlineSerial, MaintenanceMode::Background) => {
1514 self.auto_extract_facts().await;
1515 self.schedule_relationship_background().await?;
1516 }
1517 _ => {
1518 self.auto_extract_facts().await;
1519 self.auto_update_relationship().await;
1520 }
1521 }
1522 Ok(())
1523 }
1524
1525 async fn schedule_facts_background(&self) -> Result<()> {
1526 let policy = self.runtime_config.optimization.post_turn.facts.clone();
1527 let should_extract = self
1528 .facts_config
1529 .as_ref()
1530 .map(|c| c.enabled && c.auto_extract)
1531 .unwrap_or(false);
1532 if !should_extract {
1533 return Ok(());
1534 }
1535 let msgs_since = *self.messages_since_extraction.read();
1536 if msgs_since < 2 {
1537 return Ok(());
1538 }
1539 let Some(actor_id) = self.effective_actor_id() else {
1540 self.record_skipped_maintenance(
1541 "facts",
1542 ObservationPurpose::FactsExtraction,
1543 "missing_actor",
1544 Some(&policy),
1545 );
1546 return Ok(());
1547 };
1548 let Some(extractor) = self.fact_extractor.read().clone() else {
1549 return Ok(());
1550 };
1551 let messages = match self.memory.get_messages(None).await {
1552 Ok(messages) => messages,
1553 Err(error) => {
1554 warn!(error = %error, "failed to snapshot messages for fact extraction");
1555 return Ok(());
1556 }
1557 };
1558 let messages = Self::readable_native_messages(messages)?;
1559 let recent: Vec<_> = messages
1560 .iter()
1561 .rev()
1562 .take(msgs_since)
1563 .rev()
1564 .cloned()
1565 .collect();
1566 if recent.is_empty() {
1567 return Ok(());
1568 }
1569 let existing = self
1570 .actor_facts_cache
1571 .read()
1572 .get(&actor_id)
1573 .cloned()
1574 .unwrap_or_default();
1575 let categories = self
1576 .facts_config
1577 .as_ref()
1578 .map(|c| c.custom_categories.clone())
1579 .unwrap_or_default();
1580 let store = self.fact_store.read().clone();
1581 let cache = Arc::clone(&self.actor_facts_cache);
1582 let counter = Arc::clone(&self.messages_since_extraction);
1583 let hooks = Arc::clone(&self.hooks);
1584 let agent_id = self.info.id.clone();
1585 let observation = current_observation_context();
1586 let key = MaintenanceSequenceKey::actor(
1587 agent_id,
1588 actor_id.clone(),
1589 RuntimeTaskPurpose::PostTurnFacts,
1590 );
1591 let actor_for_task = actor_id.clone();
1592 let task = async move {
1593 let run = async move {
1594 let facts = extractor
1595 .extract(&recent, &existing, Some(&actor_for_task), &categories)
1596 .await?;
1597 if !facts.is_empty() {
1598 if let Some(store) = store {
1599 let authoritative = store.add_facts(&actor_for_task, facts.clone()).await?;
1600 cache.write().insert(actor_for_task.clone(), authoritative);
1601 } else {
1602 cache
1603 .write()
1604 .entry(actor_for_task.clone())
1605 .or_default()
1606 .extend(facts.clone());
1607 }
1608 {
1609 let mut count = counter.write();
1610 if *count <= msgs_since {
1611 *count = 0;
1612 } else {
1613 *count -= msgs_since;
1614 }
1615 }
1616 hooks.on_facts_extracted(&actor_for_task, &facts).await;
1617 }
1618 Ok(())
1619 };
1620 if let Some(context) = observation {
1621 with_observation_context(
1622 context.with_purpose(ObservationPurpose::FactsExtraction),
1623 run,
1624 )
1625 .await
1626 } else {
1627 run.await
1628 }
1629 };
1630 self.spawn_or_handle_background(Some(key), task, "facts", &policy)
1631 .await
1632 }
1633
1634 async fn schedule_relationship_background(&self) -> Result<()> {
1635 let policy = self
1636 .runtime_config
1637 .optimization
1638 .post_turn
1639 .relationships
1640 .clone();
1641 let Some(manager) = self.relationship_manager.as_ref().cloned() else {
1642 return Ok(());
1643 };
1644 let Some(actor_id) = self.effective_actor_id() else {
1645 self.record_skipped_maintenance(
1646 "relationships",
1647 ObservationPurpose::RelationshipUpdate,
1648 "missing_actor",
1649 Some(&policy),
1650 );
1651 return Ok(());
1652 };
1653 let recent_messages = manager.config().auto_update.recent_messages;
1654 let messages = match self.memory.get_messages(Some(recent_messages)).await {
1655 Ok(messages) => messages,
1656 Err(error) => {
1657 warn!(actor = %actor_id, error = %error, "failed to snapshot messages for relationship update");
1658 return Ok(());
1659 }
1660 };
1661 let messages = Self::readable_native_messages(messages)?;
1662 let storage = self.storage.read().clone();
1663 let hooks = Arc::clone(&self.hooks);
1664 let agent_id = self.info.id.clone();
1665 let observation = current_observation_context();
1666 let key = MaintenanceSequenceKey::actor(
1667 agent_id.clone(),
1668 actor_id.clone(),
1669 RuntimeTaskPurpose::PostTurnRelationship,
1670 );
1671 let actor_for_task = actor_id.clone();
1672 let task = async move {
1673 let run = async move {
1674 if manager.config().auto_update.enabled {
1675 let update = manager.auto_update(&actor_for_task, &messages).await?;
1676 if !update.changes.is_empty() {
1677 hooks
1678 .on_relationship_change(&actor_for_task, &update.changes)
1679 .await;
1680 }
1681 if let Some(ref event) = update.event {
1682 hooks.on_notable_event(&actor_for_task, event).await;
1683 }
1684 }
1685 if manager.config().persistence.enabled
1686 && let (Some(storage), Some(value)) =
1687 (storage, manager.relationship_as_value(&actor_for_task)?)
1688 {
1689 storage
1690 .save_relationship(&agent_id, &actor_for_task, &value)
1691 .await?;
1692 }
1693 Ok(())
1694 };
1695 if let Some(context) = observation {
1696 with_observation_context(
1697 context.with_purpose(ObservationPurpose::RelationshipUpdate),
1698 run,
1699 )
1700 .await
1701 } else {
1702 run.await
1703 }
1704 };
1705 self.spawn_or_handle_background(Some(key), task, "relationships", &policy)
1706 .await
1707 }
1708
1709 async fn spawn_or_handle_background<F>(
1711 &self,
1712 key: Option<MaintenanceSequenceKey>,
1713 task: F,
1714 label: &'static str,
1715 policy: &crate::optimization::config::MaintenanceTaskPolicy,
1716 ) -> Result<()>
1717 where
1718 F: Future<Output = Result<()>> + Send + 'static,
1719 {
1720 if self.background_maintenance.is_full() {
1721 match self
1722 .runtime_config
1723 .optimization
1724 .post_turn
1725 .on_background_overflow
1726 {
1727 BackgroundOverflowPolicy::RunInline => {
1728 record_background_maintenance_event(
1729 self.observability_manager.as_ref(),
1730 label,
1731 EventStatus::Success,
1732 0,
1733 "inline_overflow",
1734 None,
1735 Some(policy),
1736 );
1737 let start = Instant::now();
1738 match task.await {
1739 Ok(()) => record_background_maintenance_event(
1740 self.observability_manager.as_ref(),
1741 label,
1742 EventStatus::Success,
1743 start.elapsed().as_millis() as u64,
1744 "inline_completed",
1745 None,
1746 Some(policy),
1747 ),
1748 Err(error) => {
1749 warn!(label = label, error = %error, "inline maintenance fallback failed");
1750 record_background_maintenance_event(
1751 self.observability_manager.as_ref(),
1752 label,
1753 EventStatus::Error,
1754 start.elapsed().as_millis() as u64,
1755 "inline_failed",
1756 Some(error.to_string()),
1757 Some(policy),
1758 );
1759 return Err(error);
1760 }
1761 }
1762 }
1763 BackgroundOverflowPolicy::Drop => {
1764 self.record_skipped_maintenance(
1765 label,
1766 ObservationPurpose::Other(label.to_string()),
1767 "queue_full",
1768 Some(policy),
1769 );
1770 }
1771 BackgroundOverflowPolicy::Error => {
1772 record_background_maintenance_event(
1773 self.observability_manager.as_ref(),
1774 label,
1775 EventStatus::Error,
1776 0,
1777 "queue_full",
1778 None,
1779 Some(policy),
1780 );
1781 warn!(label = label, "background maintenance queue full");
1782 return Err(AgentError::Other(format!(
1783 "background maintenance queue is full for {}",
1784 label
1785 )));
1786 }
1787 }
1788 return Ok(());
1789 }
1790
1791 record_background_maintenance_event(
1792 self.observability_manager.as_ref(),
1793 label,
1794 EventStatus::Success,
1795 0,
1796 "scheduled",
1797 None,
1798 Some(policy),
1799 );
1800 let manager = self.observability_manager.clone();
1801 let policy_for_task = policy.clone();
1802 let observed_task = async move {
1803 let start = Instant::now();
1804 let result = task.await;
1805 match &result {
1806 Ok(()) => record_background_maintenance_event(
1807 manager.as_ref(),
1808 label,
1809 EventStatus::Success,
1810 start.elapsed().as_millis() as u64,
1811 "completed",
1812 None,
1813 Some(&policy_for_task),
1814 ),
1815 Err(error) => record_background_maintenance_event(
1816 manager.as_ref(),
1817 label,
1818 EventStatus::Error,
1819 start.elapsed().as_millis() as u64,
1820 "failed",
1821 Some(error.to_string()),
1822 Some(&policy_for_task),
1823 ),
1824 }
1825 result
1826 };
1827
1828 if let Err(error) = self.background_maintenance.spawn(key, observed_task) {
1829 record_background_maintenance_event(
1830 self.observability_manager.as_ref(),
1831 label,
1832 EventStatus::Error,
1833 0,
1834 "spawn_failed",
1835 Some(error.to_string()),
1836 Some(policy),
1837 );
1838 warn!(label = label, error = %error, "background maintenance spawn failed");
1839 return Err(error);
1840 }
1841 Ok(())
1842 }
1843
1844 fn record_skipped_maintenance(
1846 &self,
1847 label: &str,
1848 purpose: ObservationPurpose,
1849 reason: &str,
1850 policy: Option<&crate::optimization::config::MaintenanceTaskPolicy>,
1851 ) {
1852 if let Some(manager) = self.observability_manager.as_ref() {
1853 let mut tags = background_maintenance_tags(label, "skipped", Some(reason), policy);
1854 tags.insert("runtime.skip_reason".to_string(), reason.to_string());
1855 manager.record_lifecycle_event(
1856 EventType::MemoryOperation {
1857 operation: format!("{}_maintenance", label),
1858 },
1859 purpose,
1860 EventStatus::Skipped,
1861 0,
1862 tags,
1863 None,
1864 );
1865 }
1866 }
1867
1868 pub fn actor_facts(&self) -> Vec<ai_agents_core::KeyFact> {
1870 let Some(actor_id) = self.effective_actor_id() else {
1871 return Vec::new();
1872 };
1873 self.actor_facts_cache
1874 .read()
1875 .get(&actor_id)
1876 .cloned()
1877 .unwrap_or_default()
1878 }
1879
1880 pub fn relationship_memory_text(&self) -> Option<String> {
1882 self.format_relationship_for_context().map(|(_, text)| text)
1883 }
1884
1885 pub async fn extract_facts(
1887 &self,
1888 last_n: usize,
1889 ) -> ai_agents_core::Result<Vec<ai_agents_core::KeyFact>> {
1890 self.extract_facts_with_source(last_n, "manual").await
1891 }
1892
1893 async fn extract_facts_with_source(
1894 &self,
1895 last_n: usize,
1896 source: &'static str,
1897 ) -> ai_agents_core::Result<Vec<ai_agents_core::KeyFact>> {
1898 let extractor = match self.fact_extractor.read().clone() {
1899 Some(e) => e,
1900 None => return Ok(vec![]),
1901 };
1902
1903 let messages = Self::readable_native_messages(self.memory.get_messages(None).await?)?;
1904 let recent: Vec<_> = messages.iter().rev().take(last_n).rev().cloned().collect();
1905
1906 if recent.is_empty() {
1907 return Ok(vec![]);
1908 }
1909
1910 let actor_id = self.effective_actor_id();
1911 let existing = actor_id
1912 .as_ref()
1913 .and_then(|aid| self.actor_facts_cache.read().get(aid).cloned())
1914 .unwrap_or_default();
1915
1916 let categories = self
1917 .facts_config
1918 .as_ref()
1919 .map(|c| c.custom_categories.clone())
1920 .unwrap_or_default();
1921
1922 let facts = self
1923 .observe_purpose(
1924 ObservationPurpose::FactsExtraction,
1925 extractor.extract(&recent, &existing, actor_id.as_deref(), &categories),
1926 )
1927 .await?;
1928
1929 if !facts.is_empty() {
1931 let fact_store_opt = self.fact_store.read().clone();
1932 let mut stored_total = 0usize;
1933 let mut cache_updated = false;
1934 if let (Some(store), Some(aid)) = (fact_store_opt, &actor_id) {
1935 let authoritative = store.add_facts(aid, facts.clone()).await?;
1937 stored_total = authoritative.len();
1938 self.actor_facts_cache
1939 .write()
1940 .insert(aid.clone(), authoritative);
1941 cache_updated = true;
1942 } else if let Some(aid) = &actor_id {
1943 let mut cache = self.actor_facts_cache.write();
1944 let entry = cache.entry(aid.clone()).or_default();
1945 entry.extend(facts.clone());
1946 stored_total = entry.len();
1947 cache_updated = true;
1948 }
1949
1950 info!(
1951 actor_id = %actor_id.as_deref().unwrap_or("<none>"),
1952 source = source,
1953 requested_messages = last_n,
1954 message_count = recent.len(),
1955 extracted_count = facts.len(),
1956 cache_updated = cache_updated,
1957 stored_total = stored_total,
1958 "facts extracted"
1959 );
1960
1961 if let Some(ref aid) = actor_id {
1962 self.hooks.on_facts_extracted(aid, &facts).await;
1963 }
1964 }
1965
1966 Ok(facts)
1967 }
1968
1969 fn resolve_actor_id_from_context(&self) {
1972 if self
1973 .current_turn_actor_context()
1974 .and_then(|ctx| ctx.effective_actor_id().map(str::to_string))
1975 .is_some()
1976 {
1977 return;
1978 }
1979
1980 if let Some(ref am_config) = self.actor_memory_config
1981 && am_config.identification.method == ai_agents_facts::IdentificationMethod::FromContext
1982 && let Some(ref path) = am_config.identification.context_path
1983 {
1984 let val = self
1986 .context_manager
1987 .get_path(path)
1988 .or_else(|| self.context_manager.get(path));
1989 if let Some(val) = val
1990 && let Some(id_str) = val.as_str()
1991 {
1992 let current = self.actor_id.read().clone();
1993 if current.as_deref() != Some(id_str) {
1994 *self.actor_id.write() = Some(id_str.to_string());
1995 let mut meta = self.session_metadata.write();
1996 meta.actor_id = Some(id_str.to_string());
1997 if !meta.actors.iter().any(|a| a == id_str) {
1998 meta.actors.push(id_str.to_string());
1999 }
2000 }
2001 }
2002 }
2003 }
2004
2005 fn format_actor_facts_for_context(&self) -> String {
2007 let should_inject = self
2009 .facts_config
2010 .as_ref()
2011 .map(|c| c.inject_in_context)
2012 .unwrap_or(true);
2013 if !should_inject {
2014 return String::new();
2015 }
2016
2017 let Some(actor_id) = self.effective_actor_id() else {
2018 return String::new();
2019 };
2020
2021 let facts = self
2022 .actor_facts_cache
2023 .read()
2024 .get(&actor_id)
2025 .cloned()
2026 .unwrap_or_default();
2027 if facts.is_empty() {
2028 return String::new();
2029 }
2030
2031 let am_config = self.actor_memory_config.as_ref();
2032 let facts_budget = self
2035 .memory_token_budget
2036 .as_ref()
2037 .map(|b| b.allocation.facts as usize)
2038 .filter(|n| *n > 0);
2039 let default_max = am_config.map(|c| c.injection.max_tokens).unwrap_or(800);
2040 let max_tokens = facts_budget.unwrap_or(default_max);
2041
2042 let filtered: Vec<ai_agents_core::KeyFact> = if let Some(cfg) = am_config {
2044 if cfg.injection.mode == ai_agents_facts::InjectionMode::OnDemand {
2045 return String::new();
2046 }
2047 if cfg.injection.mode == ai_agents_facts::InjectionMode::Category
2048 && !cfg.injection.categories.is_empty()
2049 {
2050 facts
2051 .iter()
2052 .filter(|f| {
2053 cfg.injection
2054 .categories
2055 .iter()
2056 .any(|c| f.category.to_string() == *c)
2057 })
2058 .cloned()
2059 .collect()
2060 } else {
2061 facts.clone()
2062 }
2063 } else {
2064 facts.clone()
2065 };
2066
2067 if filtered.is_empty() {
2068 return String::new();
2069 }
2070
2071 if let Some(store) = self.fact_store.read().clone() {
2072 store.format_for_context(&filtered, max_tokens)
2073 } else {
2074 String::new()
2075 }
2076 }
2077
2078 fn build_context_with_staged(&self, staged: &HashMap<String, Value>) -> HashMap<String, Value> {
2079 let context = self.build_context_with_overlays();
2080 let mut root = Value::Object(context.into_iter().collect());
2081 for (path, value) in staged {
2082 if let Ok(updated) = ai_agents_core::set_dot_path(root.clone(), path, value.clone()) {
2083 root = updated;
2084 }
2085 }
2086 match root {
2087 Value::Object(obj) => obj.into_iter().collect(),
2088 _ => HashMap::new(),
2089 }
2090 }
2091
2092 fn build_context_with_overlays(&self) -> HashMap<String, Value> {
2093 let mut context = self.context_manager.get_all();
2094 let mut root = Value::Object(context.clone().into_iter().collect());
2095
2096 if let Some(turn_ctx) = self.current_turn_actor_context() {
2097 if let Some(ref origin_actor_id) = turn_ctx.origin_actor_id
2098 && let Ok(updated) = ai_agents_core::set_dot_path(
2099 root.clone(),
2100 "interaction.origin_actor_id",
2101 serde_json::json!(origin_actor_id),
2102 )
2103 {
2104 root = updated;
2105 }
2106 if let Some(ref sender_agent_id) = turn_ctx.sender_agent_id
2107 && let Ok(updated) = ai_agents_core::set_dot_path(
2108 root.clone(),
2109 "interaction.sender_agent_id",
2110 serde_json::json!(sender_agent_id),
2111 )
2112 {
2113 root = updated;
2114 }
2115 }
2116
2117 if let Some(ref actor_id) = self.effective_actor_id()
2118 && let Ok(updated) = ai_agents_core::set_dot_path(
2119 root.clone(),
2120 "interaction.actor_id",
2121 serde_json::json!(actor_id),
2122 )
2123 {
2124 root = updated;
2125 }
2126
2127 if let Some(manager) = self.relationship_manager.as_ref()
2128 && let Some(actor_id) = self.effective_actor_id()
2129 && let Some(value) = manager.to_context_value(&actor_id)
2130 && let Ok(updated) = ai_agents_core::set_dot_path(
2131 root.clone(),
2132 &manager.config().injection.context_path,
2133 value,
2134 )
2135 {
2136 root = updated;
2137 }
2138
2139 if let Value::Object(obj) = root {
2140 context = obj.into_iter().collect();
2141 }
2142
2143 context
2144 }
2145
2146 fn resolve_actor_name_from_context(&self) -> Option<String> {
2147 for path in ["actor.name", "user.name", "player.name", "customer.name"] {
2148 if let Some(value) = self.context_manager.get_path(path)
2149 && let Some(name) = value.as_str()
2150 {
2151 return Some(name.to_string());
2152 }
2153 }
2154 None
2155 }
2156
2157 async fn maybe_load_actor_relationship(&self) {
2158 let Some(manager) = self.relationship_manager.as_ref() else {
2159 return;
2160 };
2161 let Some(actor_id) = self.effective_actor_id() else {
2162 return;
2163 };
2164
2165 let mut should_fire_loaded = false;
2166 if manager.get(&actor_id).is_none() {
2167 let mut loaded = false;
2168 if manager.config().persistence.enabled {
2169 let storage = self.storage.read().clone();
2170 if let Some(storage) = storage {
2171 match storage.load_relationship(&self.info.id, &actor_id).await {
2172 Ok(Some(value)) => match manager.insert_from_value(value) {
2173 Ok(_) => loaded = true,
2174 Err(e) => {
2175 warn!(actor = %actor_id, error = %e, "failed to restore relationship")
2176 }
2177 },
2178 Ok(None) => {}
2179 Err(e) => {
2180 warn!(actor = %actor_id, error = %e, "failed to load relationship")
2181 }
2182 }
2183 }
2184 }
2185
2186 if !loaded {
2187 manager.get_or_create(&actor_id, self.resolve_actor_name_from_context().as_deref());
2188 }
2189 should_fire_loaded = true;
2190 }
2191
2192 let actor_name = self.resolve_actor_name_from_context();
2193 let relationship = manager.touch_interaction(&actor_id, actor_name.as_deref());
2194 if should_fire_loaded {
2195 self.hooks
2196 .on_relationship_loaded(&actor_id, &relationship)
2197 .await;
2198 }
2199 }
2200
2201 fn format_relationship_for_context(&self) -> Option<(String, String)> {
2202 let manager = self.relationship_manager.as_ref()?;
2203 if !manager.config().injection.enabled {
2204 return None;
2205 }
2206 let actor_id = self.effective_actor_id()?;
2207 let relationship = manager.get(&actor_id)?;
2208 let local_cap = manager.config().injection.max_tokens;
2209 let global_cap = self
2210 .memory_token_budget
2211 .as_ref()
2212 .map(|b| b.allocation.relationships as usize)
2213 .filter(|n| *n > 0);
2214 let max_tokens = global_cap.map(|g| g.min(local_cap)).unwrap_or(local_cap);
2215 let text = ai_agents_relationships::format_relationship(
2216 &relationship,
2217 &manager.config().injection.format,
2218 max_tokens,
2219 );
2220 if text.is_empty() {
2221 None
2222 } else {
2223 Some((manager.config().injection.prompt_variable.clone(), text))
2224 }
2225 }
2226
2227 async fn persist_actor_relationship(&self, actor_id: &str) -> Result<()> {
2228 let Some(manager) = self.relationship_manager.as_ref() else {
2229 return Ok(());
2230 };
2231 if !manager.config().persistence.enabled {
2232 return Ok(());
2233 }
2234 let storage = self.storage.read().clone();
2235 let Some(storage) = storage else {
2236 return Ok(());
2237 };
2238 if let Some(value) = manager.relationship_as_value(actor_id)? {
2239 storage
2240 .save_relationship(&self.info.id, actor_id, &value)
2241 .await?;
2242 }
2243 Ok(())
2244 }
2245
2246 async fn auto_update_relationship(&self) {
2247 let Some(manager) = self.relationship_manager.as_ref() else {
2248 return;
2249 };
2250 let Some(actor_id) = self.effective_actor_id() else {
2251 return;
2252 };
2253 if !manager.config().auto_update.enabled {
2254 let _ = self.persist_actor_relationship(&actor_id).await;
2255 return;
2256 }
2257
2258 let recent_messages = manager.config().auto_update.recent_messages;
2259 let messages = match self.memory.get_messages(Some(recent_messages)).await {
2260 Ok(messages) => messages,
2261 Err(e) => {
2262 warn!(actor = %actor_id, error = %e, "failed to read messages for relationship update");
2263 return;
2264 }
2265 };
2266 let messages = match Self::readable_native_messages(messages) {
2267 Ok(messages) => messages,
2268 Err(error) => {
2269 warn!(actor = %actor_id, error = %error, "failed to project native history for relationship update");
2270 return;
2271 }
2272 };
2273
2274 match self
2275 .observe_purpose(
2276 ObservationPurpose::RelationshipUpdate,
2277 manager.auto_update(&actor_id, &messages),
2278 )
2279 .await
2280 {
2281 Ok(update) => {
2282 if !update.changes.is_empty() {
2283 self.hooks
2284 .on_relationship_change(&actor_id, &update.changes)
2285 .await;
2286 }
2287 if let Some(ref event) = update.event {
2288 self.hooks.on_notable_event(&actor_id, event).await;
2289 }
2290 let persisted = match self.persist_actor_relationship(&actor_id).await {
2291 Ok(()) => true,
2292 Err(e) => {
2293 warn!(actor = %actor_id, error = %e, "failed to persist relationship");
2294 false
2295 }
2296 };
2297 if !update.changes.is_empty() || update.event.is_some() {
2298 let changed_dimensions: Vec<String> = update
2299 .changes
2300 .iter()
2301 .map(|change| format!("{}:{}", change.perspective, change.dimension))
2302 .collect();
2303 info!(
2304 actor_id = %actor_id,
2305 change_count = update.changes.len(),
2306 changed_dimensions = ?changed_dimensions,
2307 event_present = update.event.is_some(),
2308 persisted = persisted,
2309 "relationship updated"
2310 );
2311 } else {
2312 debug!(actor_id = %actor_id, persisted = persisted, "relationship evaluation ran but found no changes");
2313 }
2314 }
2315 Err(e) => warn!(actor = %actor_id, error = %e, "relationship update failed"),
2316 }
2317 }
2318
2319 async fn auto_extract_facts(&self) {
2321 let should_extract = self
2322 .facts_config
2323 .as_ref()
2324 .map(|c| c.enabled && c.auto_extract)
2325 .unwrap_or(false);
2326
2327 if !should_extract {
2328 debug!("fact extraction skipped because auto extraction is disabled");
2329 return;
2330 }
2331
2332 let msgs_since = *self.messages_since_extraction.read();
2333 if msgs_since < 2 {
2334 debug!(
2335 messages_since_extraction = msgs_since,
2336 "fact extraction skipped until threshold is reached"
2337 );
2338 return;
2339 }
2340
2341 match self.extract_facts_with_source(msgs_since, "auto").await {
2342 Ok(facts) => {
2343 if !facts.is_empty() {
2344 *self.messages_since_extraction.write() = 0;
2345 } else {
2346 debug!("fact extraction ran but found no new facts");
2347 }
2348 }
2349 Err(e) => {
2350 warn!("fact extraction failed: {}", e);
2351 }
2352 }
2353 }
2354
2355 pub fn with_persona(mut self, manager: Arc<ai_agents_persona::PersonaManager>) -> Self {
2356 self.persona_manager = Some(manager);
2357 self
2358 }
2359
2360 pub fn persona_manager(&self) -> Option<&Arc<ai_agents_persona::PersonaManager>> {
2361 self.persona_manager.as_ref()
2362 }
2363
2364 pub fn with_disambiguation(mut self, config: DisambiguationConfig) -> Self {
2365 if config.is_enabled() {
2366 let manager = DisambiguationManager::new(config, Arc::clone(&self.llm_registry))
2367 .with_clarification_observer(Arc::new(ObservabilityClarificationObserver));
2368 self.disambiguation_manager = Some(manager);
2369 }
2370 self
2371 }
2372
2373 pub fn disambiguation_manager(&self) -> Option<&DisambiguationManager> {
2374 self.disambiguation_manager.as_ref()
2375 }
2376
2377 pub fn has_disambiguation(&self) -> bool {
2378 self.disambiguation_manager
2379 .as_ref()
2380 .is_some_and(|m| m.is_enabled())
2381 }
2382
2383 pub async fn init_storage(&self) -> Result<()> {
2384 let _guard = self.storage_init.lock().await;
2388 let mut storage = self.storage.read().clone();
2389 if storage.is_none() && !self.storage_config.is_none() {
2390 let storage_config = self.convert_storage_config();
2391 storage = create_storage(&storage_config).await?;
2392 *self.storage.write() = storage.clone();
2393 }
2394
2395 self.validate_storage_requirements(storage.as_deref())?;
2396 self.complete_facts_init().await;
2397 Ok(())
2398 }
2399
2400 fn validate_storage_requirements(&self, storage: Option<&dyn AgentStorage>) -> Result<()> {
2401 let facts_required = self
2402 .facts_config
2403 .as_ref()
2404 .is_some_and(|config| config.enabled)
2405 || self
2406 .actor_memory_config
2407 .as_ref()
2408 .is_some_and(|config| config.enabled);
2409 let relationships_required = self
2410 .relationship_manager
2411 .as_ref()
2412 .is_some_and(|manager| manager.config().persistence.enabled);
2413
2414 let Some(storage) = storage else {
2415 let mut requirements = Vec::new();
2416 if facts_required {
2417 requirements.push("actor facts or actor memory");
2418 }
2419 if relationships_required {
2420 requirements.push("persistent relationships");
2421 }
2422 if requirements.is_empty() {
2423 return Ok(());
2424 }
2425 return Err(AgentError::Config(format!(
2426 "Storage is required for enabled {} but none is configured or injected",
2427 requirements.join(" and ")
2428 )));
2429 };
2430
2431 if facts_required && !storage.supports(StorageCapability::ActorFacts) {
2435 return Err(AgentError::UnsupportedStorageCapability(
2436 StorageCapability::ActorFacts,
2437 ));
2438 }
2439 if relationships_required && !storage.supports(StorageCapability::ActorRelationships) {
2440 return Err(AgentError::UnsupportedStorageCapability(
2441 StorageCapability::ActorRelationships,
2442 ));
2443 }
2444 Ok(())
2445 }
2446
2447 async fn complete_facts_init(&self) {
2450 if self.fact_store.read().is_some() {
2451 return;
2452 }
2453 let storage = match self.storage.read().clone() {
2454 Some(s) => s,
2455 None => return,
2456 };
2457
2458 let facts_enabled = self
2459 .facts_config
2460 .as_ref()
2461 .map(|f| f.enabled)
2462 .unwrap_or(false);
2463 let actor_memory_enabled = self
2464 .actor_memory_config
2465 .as_ref()
2466 .map(|a| a.enabled)
2467 .unwrap_or(false);
2468
2469 if !facts_enabled && !actor_memory_enabled {
2470 return;
2471 }
2472
2473 let fc = self.facts_config.clone().unwrap_or_default();
2474 let store = Arc::new(ai_agents_facts::FactStore::new(
2475 storage,
2476 self.info.id.clone(),
2477 fc.clone(),
2478 ));
2479
2480 let extractor: Option<Arc<dyn ai_agents_facts::FactExtractor>> = if facts_enabled {
2481 let extractor_llm = fc
2482 .extractor_llm
2483 .as_ref()
2484 .and_then(|alias| self.llm_registry.get(alias).ok())
2485 .or_else(|| self.llm_registry.router().ok())
2486 .or_else(|| self.llm_registry.default().ok());
2487 extractor_llm.map(|llm| {
2488 Arc::new(ai_agents_facts::LLMFactExtractor::new(llm, fc.clone()))
2489 as Arc<dyn ai_agents_facts::FactExtractor>
2490 })
2491 } else {
2492 None
2493 };
2494
2495 *self.fact_store.write() = Some(store);
2496 *self.fact_extractor.write() = extractor;
2497 debug!(
2498 agent = %self.info.id,
2499 facts_enabled,
2500 actor_memory_enabled,
2501 "facts storage initialized"
2502 );
2503 }
2504
2505 fn convert_storage_config(&self) -> StorageStorageConfig {
2506 crate::spec::storage::to_storage_config(&self.storage_config)
2507 }
2508
2509 pub fn storage(&self) -> Option<Arc<dyn AgentStorage>> {
2510 self.storage.read().clone()
2511 }
2512
2513 pub fn storage_config(&self) -> &StorageConfig {
2514 &self.storage_config
2515 }
2516
2517 pub fn spawner(&self) -> Option<&Arc<crate::spawner::AgentSpawner>> {
2519 self.spawner.as_ref()
2520 }
2521
2522 pub fn spawner_registry(&self) -> Option<&Arc<crate::spawner::AgentRegistry>> {
2524 self.spawner_registry.as_ref()
2525 }
2526
2527 pub fn has_spawner(&self) -> bool {
2528 self.spawner_registry.is_some()
2529 }
2530
2531 pub fn with_spawner_handles(
2532 mut self,
2533 spawner: Arc<crate::spawner::AgentSpawner>,
2534 registry: Arc<crate::spawner::AgentRegistry>,
2535 ) -> Self {
2536 self.spawner = Some(spawner);
2537 self.spawner_registry = Some(registry);
2538 self
2539 }
2540
2541 pub fn with_hooks(mut self, hooks: Arc<dyn AgentHooks>) -> Self {
2542 self.hooks = hooks;
2543 self
2544 }
2545
2546 pub fn with_parallel_tools(mut self, config: ParallelToolsConfig) -> Self {
2547 self.parallel_tools = config;
2548 self
2549 }
2550
2551 pub fn with_streaming(mut self, config: StreamingConfig) -> Self {
2552 self.streaming = config;
2553 self
2554 }
2555
2556 pub fn with_hitl(mut self, engine: HITLEngine, handler: Arc<dyn ApprovalHandler>) -> Self {
2557 self.hitl_engine = Some(engine);
2558 self.approval_handler = handler;
2559 self
2560 }
2561
2562 pub fn with_max_context_tokens(mut self, tokens: u32) -> Self {
2563 self.max_context_tokens = tokens;
2564 self
2565 }
2566
2567 pub fn with_memory_token_budget(mut self, budget: MemoryTokenBudget) -> Self {
2568 self.memory_token_budget = Some(budget);
2569 self
2570 }
2571
2572 pub fn with_recovery_manager(mut self, manager: RecoveryManager) -> Self {
2573 self.recovery_manager = manager;
2574 self
2575 }
2576
2577 pub fn with_tool_security(mut self, engine: ToolSecurityEngine) -> Self {
2578 self.tool_security = engine;
2579 self
2580 }
2581
2582 pub fn runtime_control(&self) -> RuntimeControlHandle {
2584 RuntimeControlHandle {
2585 state: Arc::clone(&self.runtime_control),
2586 }
2587 }
2588
2589 pub fn set_question_handler(&self, handler: Option<Arc<dyn QuestionHandler>>) {
2591 self.tools.set_question_handler(handler);
2592 }
2593
2594 pub fn set_diagnostics_provider(&self, provider: Arc<dyn DiagnosticsProvider>) {
2596 self.tools.set_diagnostics_provider(provider);
2597 }
2598
2599 pub fn set_command_runner(&self, runner: Arc<dyn CommandRunner>) {
2601 self.tools.set_command_runner(runner);
2602 }
2603
2604 pub fn set_web_search_provider(&self, provider: Arc<dyn ai_agents_tools::WebSearchProvider>) {
2606 self.tools.set_web_search_provider(provider);
2607 }
2608
2609 pub fn todos(&self) -> Vec<TodoItem> {
2611 self.tools.todos()
2612 }
2613
2614 fn active_tool_security(&self) -> ToolSecurityEngine {
2616 self.runtime_control
2617 .tool_security_override
2618 .read()
2619 .clone()
2620 .unwrap_or_else(|| self.tool_security.clone())
2621 }
2622
2623 fn runtime_safety_snapshot(&self) -> RuntimeSafetySnapshot {
2625 let _guard = self.runtime_control.snapshot_guard.read();
2626 RuntimeSafetySnapshot {
2627 version: self.runtime_control.version.load(Ordering::SeqCst),
2628 emergency_deny: self.runtime_control.emergency_deny.load(Ordering::SeqCst),
2629 tool_security: self
2630 .runtime_control
2631 .tool_security_override
2632 .read()
2633 .clone()
2634 .unwrap_or_else(|| self.tool_security.clone()),
2635 tool_scope_override: self.runtime_control.tool_scope_override.read().clone(),
2636 }
2637 }
2638
2639 fn admit_tool_execution(
2641 &self,
2642 expected_runtime_version: u64,
2643 expected_policy_version: u64,
2644 expected_state_generation: Option<u64>,
2645 canonical_id: &str,
2646 ) -> SecurityCheckResult {
2647 let _guard = self.runtime_control.snapshot_guard.read();
2648 if self.runtime_control.emergency_deny.load(Ordering::SeqCst) {
2649 return SecurityCheckResult::Block {
2650 reason: "runtime emergency deny is enabled".to_string(),
2651 };
2652 }
2653 let runtime_version = self.runtime_control.version.load(Ordering::SeqCst);
2654 let security_engine = self
2655 .runtime_control
2656 .tool_security_override
2657 .read()
2658 .clone()
2659 .unwrap_or_else(|| self.tool_security.clone());
2660 if runtime_version != expected_runtime_version
2661 || security_engine.policy_version() != expected_policy_version
2662 {
2663 return SecurityCheckResult::Block {
2664 reason: "runtime safety controls changed before admission".to_string(),
2665 };
2666 }
2667 let current_state_generation = self
2668 .state_machine
2669 .as_ref()
2670 .map(|state_machine| state_machine.generation());
2671 if current_state_generation != expected_state_generation {
2672 return SecurityCheckResult::Block {
2673 reason: "state scope changed before admission".to_string(),
2674 };
2675 }
2676 security_engine.admit_tool_execution(canonical_id)
2677 }
2678
2679 pub fn with_process_processor(mut self, processor: ProcessProcessor) -> Self {
2680 let processor = processor.with_stage_observer(Arc::new(ObservabilityProcessStageObserver));
2681 self.process_processor = Some(processor);
2682 self
2683 }
2684
2685 pub fn with_state_machine(
2686 mut self,
2687 state_machine: Arc<StateMachine>,
2688 evaluator: Arc<dyn TransitionEvaluator>,
2689 ) -> Self {
2690 self.state_machine = Some(state_machine);
2691 self.transition_evaluator = Some(evaluator);
2692 self
2693 }
2694
2695 pub fn with_context_manager(mut self, manager: Arc<ContextManager>) -> Self {
2696 self.context_manager = manager;
2697 self
2698 }
2699
2700 pub fn register_message_filter(&self, name: impl Into<String>, filter: Arc<dyn MessageFilter>) {
2701 self.message_filters.write().insert(name.into(), filter);
2702 }
2703
2704 pub fn set_context(&self, key: &str, value: Value) -> Result<()> {
2705 self.context_manager.update(key, value)
2706 }
2707
2708 pub fn update_context(&self, path: &str, value: Value) -> Result<()> {
2709 self.context_manager.update(path, value)
2710 }
2711
2712 pub fn get_context(&self) -> HashMap<String, Value> {
2713 self.build_context_with_overlays()
2714 }
2715
2716 pub fn remove_context(&self, key: &str) -> Option<Value> {
2717 self.context_manager.remove(key)
2718 }
2719
2720 pub async fn refresh_context(&self, key: &str) -> Result<()> {
2721 self.context_manager.refresh(key).await
2722 }
2723
2724 pub fn register_context_provider(&self, name: &str, provider: Arc<dyn ContextProvider>) {
2725 self.context_manager.register_provider(name, provider);
2726 }
2727
2728 pub fn current_state(&self) -> Option<String> {
2729 self.state_machine.as_ref().map(|sm| sm.current())
2730 }
2731
2732 async fn invalidate_pending_confirmation(&self, reason: &'static str) {
2734 self.disambiguation_epoch.fetch_add(1, Ordering::SeqCst);
2735 let Some(disambiguator) = self.disambiguation_manager.as_ref() else {
2736 return;
2737 };
2738 if disambiguator.has_pending_confirmation().await {
2739 disambiguator.clear_pending().await;
2740 *self.pending_skill_id.write() = None;
2741 info!(
2742 confirmation_event = "invalidated",
2743 invalidation_reason = reason,
2744 "Runtime invalidated pending confirmation"
2745 );
2746 }
2747 }
2748
2749 async fn admit_disambiguation_redispatch(
2751 &self,
2752 expected_epoch: u64,
2753 expected_state_generation: Option<u64>,
2754 ) -> Result<tokio::sync::RwLockReadGuard<'_, ()>> {
2755 let admission = self.disambiguation_admission.read().await;
2756 let state_generation = self
2757 .state_machine
2758 .as_ref()
2759 .map(|state_machine| state_machine.generation());
2760 if self.disambiguation_epoch.load(Ordering::SeqCst) != expected_epoch
2761 || state_generation != expected_state_generation
2762 {
2763 return Err(AgentError::Other(
2764 "Disambiguation ownership changed before redispatch admission".to_string(),
2765 ));
2766 }
2767 Ok(admission)
2768 }
2769
2770 fn reserve_state_transition(&self) -> Option<StateTransitionReservation<'_>> {
2772 self.state_transition_reserved
2773 .compare_exchange(false, true, Ordering::SeqCst, Ordering::SeqCst)
2774 .ok()
2775 .map(|_| StateTransitionReservation {
2776 reserved: &self.state_transition_reserved,
2777 })
2778 }
2779
2780 async fn admit_optional_disambiguation_ownership(
2782 &self,
2783 ownership: Option<DisambiguationOwnership>,
2784 ) -> Result<Option<tokio::sync::RwLockReadGuard<'_, ()>>> {
2785 match ownership {
2786 Some(ownership) => self
2787 .admit_disambiguation_redispatch(ownership.epoch, ownership.state_generation)
2788 .await
2789 .map(Some),
2790 None => Ok(None),
2791 }
2792 }
2793
2794 pub async fn transition_to(&self, state: &str) -> Result<()> {
2796 let Some(ref sm) = self.state_machine else {
2797 return Ok(());
2798 };
2799 let claim_admission = self.disambiguation_admission.write().await;
2800 let reservation = self.reserve_state_transition().ok_or_else(|| {
2801 AgentError::Other("Another state transition is already in progress".to_string())
2802 })?;
2803 let from_state = sm.current();
2804 let expected_state_generation = sm.generation();
2805 let expected_disambiguation_epoch = self.disambiguation_epoch.load(Ordering::SeqCst);
2806 let history_before = sm.history();
2807 drop(claim_admission);
2808
2809 self.execute_state_exit_actions(&from_state).await;
2810
2811 let admission = self.disambiguation_admission.write().await;
2812 if sm.current() != from_state
2813 || sm.generation() != expected_state_generation
2814 || self.disambiguation_epoch.load(Ordering::SeqCst) != expected_disambiguation_epoch
2815 {
2816 return Err(AgentError::Other(
2817 "State ownership changed during manual transition preparation".to_string(),
2818 ));
2819 }
2820 sm.transition_to(state, "manual transition")?;
2821 self.invalidate_pending_confirmation("state_transition")
2822 .await;
2823 let entered = sm.current();
2824 let is_reentry = Self::state_was_previously_entered(&entered, &from_state, &history_before);
2825 drop(admission);
2826
2827 self.execute_state_enter_actions(&entered, is_reentry).await;
2828 drop(reservation);
2829 info!(to = %entered, "Manual state transition");
2830 Ok(())
2831 }
2832
2833 pub fn state_history(&self) -> Vec<StateTransitionEvent> {
2834 self.state_machine
2835 .as_ref()
2836 .map(|sm| sm.history())
2837 .unwrap_or_default()
2838 }
2839
2840 pub fn session_metadata(&self) -> ai_agents_core::SessionMetadata {
2842 self.session_metadata.read().clone()
2843 }
2844
2845 pub async fn delete_actor_data(&self, actor_id: &str) -> Result<()> {
2848 let allowed = self
2849 .actor_memory_config
2850 .as_ref()
2851 .map(|c| c.privacy.allow_deletion)
2852 .unwrap_or(true);
2853 if !allowed {
2854 return Err(AgentError::Config(
2855 "privacy.allow_deletion is false; actor data deletion is not permitted".into(),
2856 ));
2857 }
2858 let storage = self.storage.read().clone();
2859 if let Some(storage) = storage {
2860 if !storage.supports(StorageCapability::ActorDataDeletion) {
2864 return Err(AgentError::UnsupportedStorageCapability(
2865 StorageCapability::ActorDataDeletion,
2866 ));
2867 }
2868 storage.delete_actor_data(&self.info.id, actor_id).await?;
2869 } else {
2870 let store = { self.fact_store.read().clone() };
2874 if let Some(store) = store {
2875 store.delete_actor_data(actor_id).await?;
2876 }
2877 }
2878 if let Some(manager) = self.relationship_manager.as_ref() {
2879 manager.remove(actor_id);
2880 }
2881 self.actor_facts_cache.write().remove(actor_id);
2882 Ok(())
2883 }
2884
2885 pub fn set_session_metadata(&self, meta: ai_agents_core::SessionMetadata) {
2887 *self.session_metadata.write() = meta;
2888 }
2889
2890 pub async fn cleanup_expired_sessions(&self) -> Result<usize> {
2892 let storage = self.storage.read().clone();
2893 match storage {
2894 Some(s) => {
2895 let count = s.cleanup_expired().await?;
2896 if count > 0 {
2897 self.hooks.on_sessions_expired(count).await;
2898 }
2899 Ok(count)
2900 }
2901 None => Err(AgentError::Config(
2902 "No storage configured. Use with_storage_config() or with_storage() first".into(),
2903 )),
2904 }
2905 }
2906
2907 pub async fn list_sessions_filtered(
2909 &self,
2910 filter: &ai_agents_core::SessionFilter,
2911 ) -> Result<Vec<ai_agents_core::SessionSummary>> {
2912 let storage = self.storage.read().clone();
2913 match storage {
2914 Some(s) => s.list_sessions_filtered(filter).await,
2915 None => Err(AgentError::Config(
2916 "No storage configured. Use with_storage_config() or with_storage() first".into(),
2917 )),
2918 }
2919 }
2920
2921 pub async fn save_state(&self) -> Result<AgentSnapshot> {
2922 let memory_snapshot = self.memory.snapshot().await?;
2923 let state_machine_snapshot = self.state_machine.as_ref().map(|sm| sm.snapshot());
2924 let context_snapshot = self.context_manager.snapshot();
2925
2926 let mut snapshot = AgentSnapshot::new(self.info.id.clone())
2927 .with_memory(memory_snapshot)
2928 .with_context(context_snapshot)
2929 .with_state_machine(
2930 state_machine_snapshot.unwrap_or_else(|| StateMachineSnapshot {
2931 current_state: String::new(),
2932 previous_state: None,
2933 turn_count: 0,
2934 no_transition_count: 0,
2935 history: vec![],
2936 }),
2937 );
2938
2939 if let Some(ref persona) = self.persona_manager {
2940 snapshot.persona = Some(persona.snapshot_as_value()?);
2941 }
2942
2943 if let Some(ref relationships) = self.relationship_manager {
2944 snapshot.relationships = Some(relationships.snapshot_as_value()?);
2945 }
2946
2947 Ok(snapshot)
2948 }
2949
2950 pub async fn save_state_full(&self) -> Result<AgentSnapshot> {
2952 let mut snapshot = self.save_state().await?;
2953 if let Some(ref registry) = self.spawner_registry {
2954 let entries = registry.list_with_specs();
2955 if !entries.is_empty() {
2956 snapshot = snapshot.with_spawned_agents(entries);
2957 }
2958 }
2959 Ok(snapshot)
2960 }
2961
2962 pub async fn restore_state(&self, snapshot: AgentSnapshot) -> Result<()> {
2964 let _admission = self.disambiguation_admission.write().await;
2965 if self.state_transition_reserved.load(Ordering::SeqCst) {
2966 return Err(AgentError::Other(
2967 "Cannot restore state while a state transition is in progress".to_string(),
2968 ));
2969 }
2970 self.invalidate_pending_confirmation("state_restore").await;
2971 *self.pending_skill_id.write() = None;
2972 if let Some(disambiguator) = self.disambiguation_manager.as_ref() {
2973 disambiguator.clear_pending().await;
2974 }
2975 self.memory.restore(snapshot.memory).await?;
2976 self.active_native_exchanges.write().clear();
2977
2978 if let (Some(sm), Some(sm_snapshot)) = (&self.state_machine, snapshot.state_machine)
2979 && !sm_snapshot.current_state.is_empty()
2980 {
2981 sm.restore(sm_snapshot)?;
2982 }
2983
2984 self.context_manager.restore(snapshot.context);
2985
2986 if let (Some(persona_value), Some(persona_manager)) =
2987 (snapshot.persona, &self.persona_manager)
2988 {
2989 persona_manager.restore_from_value(persona_value)?;
2990 }
2991
2992 if let (Some(relationship_value), Some(relationship_manager)) =
2993 (snapshot.relationships, &self.relationship_manager)
2994 {
2995 relationship_manager.restore_from_value(relationship_value)?;
2996 }
2997
2998 info!(agent_id = %snapshot.agent_id, "State restored");
2999 Ok(())
3000 }
3001
3002 pub async fn save_to(&self, storage: &dyn AgentStorage, session_id: &str) -> Result<()> {
3003 let snapshot = self.save_state().await?;
3004 storage.save(session_id, &snapshot).await
3005 }
3006
3007 async fn load_session_restore(
3008 storage: &dyn AgentStorage,
3009 session_id: &str,
3010 ) -> Result<Option<StoredSessionRestore>> {
3011 let Some(snapshot) = storage.load(session_id).await? else {
3012 return Ok(None);
3013 };
3014 let metadata = if storage.supports(StorageCapability::SessionMetadata) {
3018 storage.load_metadata(session_id).await?
3019 } else {
3020 None
3021 };
3022 Ok(Some(StoredSessionRestore { snapshot, metadata }))
3023 }
3024
3025 async fn capture_session_restore_point(&self) -> Result<RuntimeSessionRestorePoint> {
3026 Ok(RuntimeSessionRestorePoint {
3027 snapshot: self.save_state().await?,
3028 metadata: self.session_metadata(),
3029 actor_id: self.actor_id(),
3030 session_id: self.current_session_id.read().clone(),
3031 })
3032 }
3033
3034 async fn apply_session_restore_unchecked(
3035 &self,
3036 session_id: &str,
3037 stored: StoredSessionRestore,
3038 ) -> Result<()> {
3039 self.restore_state(stored.snapshot).await?;
3040 let metadata = stored.metadata.unwrap_or_default();
3041 if let Some(actor_id) = metadata.actor_id.as_deref() {
3042 self.set_actor_id(actor_id)?;
3043 } else {
3044 self.clear_actor_id();
3045 }
3046 self.set_session_metadata(metadata);
3047 *self.current_session_id.write() = Some(session_id.to_string());
3048 Ok(())
3049 }
3050
3051 async fn restore_session_restore_point(
3052 &self,
3053 restore_point: &RuntimeSessionRestorePoint,
3054 ) -> Result<()> {
3055 self.restore_state(restore_point.snapshot.clone()).await?;
3056 if let Some(actor_id) = restore_point.actor_id.as_deref() {
3057 self.set_actor_id(actor_id)?;
3058 } else {
3059 self.clear_actor_id();
3060 }
3061 self.set_session_metadata(restore_point.metadata.clone());
3062 *self.current_session_id.write() = restore_point.session_id.clone();
3063 Ok(())
3064 }
3065
3066 async fn apply_session_restore(
3067 &self,
3068 session_id: &str,
3069 stored: StoredSessionRestore,
3070 ) -> Result<()> {
3071 let before = self.capture_session_restore_point().await?;
3072 if let Err(error) = self
3073 .apply_session_restore_unchecked(session_id, stored)
3074 .await
3075 {
3076 return match self.restore_session_restore_point(&before).await {
3077 Ok(()) => Err(error),
3078 Err(rollback_error) => Err(AgentError::Other(format!(
3079 "Session restore failed: {error}; rollback failed: {rollback_error}"
3080 ))),
3081 };
3082 }
3083 Ok(())
3084 }
3085
3086 async fn rollback_session_restore_set(
3087 parent: Option<(&RuntimeAgent, &RuntimeSessionRestorePoint)>,
3088 children: &[(String, Arc<RuntimeAgent>, RuntimeSessionRestorePoint)],
3089 ) -> Vec<String> {
3090 let mut errors = Vec::new();
3091 if let Some((agent, restore_point)) = parent
3092 && let Err(error) = agent.restore_session_restore_point(restore_point).await
3093 {
3094 errors.push(format!("parent: {error}"));
3095 }
3096 for (id, agent, restore_point) in children {
3097 if let Err(error) = agent.restore_session_restore_point(restore_point).await {
3098 errors.push(format!("child '{id}': {error}"));
3099 }
3100 }
3101 errors
3102 }
3103
3104 fn restore_failure(error: impl std::fmt::Display, rollback_errors: Vec<String>) -> AgentError {
3105 if rollback_errors.is_empty() {
3106 AgentError::Other(format!(
3107 "Session restore failed: {error}; runtime state was rolled back"
3108 ))
3109 } else {
3110 AgentError::Other(format!(
3111 "Session restore failed: {error}; rollback also failed for {}",
3112 rollback_errors.join(", ")
3113 ))
3114 }
3115 }
3116
3117 pub async fn load_from(&self, storage: &dyn AgentStorage, session_id: &str) -> Result<bool> {
3118 let Some(stored) = Self::load_session_restore(storage, session_id).await? else {
3119 return Ok(false);
3120 };
3121 self.apply_session_restore(session_id, stored).await?;
3122 Ok(true)
3123 }
3124
3125 pub async fn save_session(&self, session_id: &str) -> Result<()> {
3126 let storage = self.storage.read().clone();
3127 match storage {
3128 Some(s) => {
3129 let is_new = {
3131 let cur = self.current_session_id.read().clone();
3132 cur.as_deref() != Some(session_id)
3133 };
3134 if is_new {
3135 *self.current_session_id.write() = Some(session_id.to_string());
3136 self.hooks.on_session_created(session_id).await;
3137 }
3138
3139 {
3141 let now = chrono::Utc::now();
3142 let msg_count = self
3143 .memory
3144 .get_messages(None)
3145 .await
3146 .map(|v| v.len())
3147 .unwrap_or(0);
3148 let mut meta = self.session_metadata.write();
3149 meta.last_active = now;
3150 meta.message_count = msg_count;
3151 if meta.actor_id.is_none() {
3152 meta.actor_id = self.actor_id.read().clone();
3153 }
3154 }
3155
3156 let snapshot = self.save_state().await?;
3157 if s.supports(StorageCapability::SessionMetadata) {
3161 let metadata = self.session_metadata.read().clone();
3162 s.save_snapshot_with_metadata(session_id, &snapshot, &metadata)
3163 .await
3164 } else {
3165 s.save(session_id, &snapshot).await
3166 }
3167 }
3168 None => Err(AgentError::Config(
3169 "No storage configured. Use with_storage_config() or with_storage() first".into(),
3170 )),
3171 }
3172 }
3173
3174 pub async fn load_session(&self, session_id: &str) -> Result<bool> {
3175 let storage = self.storage.read().clone();
3176 match storage {
3177 Some(storage) => self.load_from(storage.as_ref(), session_id).await,
3178 None => Err(AgentError::Config(
3179 "No storage configured. Use with_storage_config() or with_storage() first".into(),
3180 )),
3181 }
3182 }
3183
3184 pub async fn restore_session_full(&self, session_id: &str) -> Result<usize> {
3186 self.init_storage().await?;
3187 let storage = self.storage.read().clone().ok_or_else(|| {
3188 AgentError::Config(
3189 "No storage configured. Use with_storage_config() or with_storage() first".into(),
3190 )
3191 })?;
3192 let target_parent = Self::load_session_restore(storage.as_ref(), session_id)
3193 .await?
3194 .ok_or_else(|| AgentError::Persistence(format!("Session not found: {session_id}")))?;
3195 let manifest = target_parent
3196 .snapshot
3197 .spawned_agents
3198 .clone()
3199 .unwrap_or_default();
3200
3201 let registry = self.spawner_registry.as_ref().cloned();
3202 let spawner = if manifest.is_empty() {
3203 self.spawner.as_ref().cloned()
3204 } else {
3205 Some(self.spawner.as_ref().cloned().ok_or_else(|| {
3206 AgentError::Config(
3207 "Saved session contains child agents but this runtime has no spawner".into(),
3208 )
3209 })?)
3210 };
3211 let registry = if manifest.is_empty() {
3212 registry
3213 } else {
3214 Some(registry.ok_or_else(|| {
3215 AgentError::Config(
3216 "Saved session contains child agents but this runtime has no registry".into(),
3217 )
3218 })?)
3219 };
3220
3221 let mut target_ids = HashSet::with_capacity(manifest.len());
3222 let mut prepared = Vec::with_capacity(manifest.len());
3223 for entry in manifest {
3224 if !target_ids.insert(entry.id.clone()) {
3225 return Err(AgentError::InvalidSpec(format!(
3226 "Saved child manifest contains duplicate ID: {}",
3227 entry.id
3228 )));
3229 }
3230 let spec = crate::spec::AgentSpec::from_yaml_strict(&entry.spec_yaml)?;
3231 spawner
3232 .as_ref()
3233 .expect("non-empty manifests require a spawner")
3234 .validate_explicit_child(&entry.id, &spec)?;
3235 prepared.push((entry.id, spec));
3236 }
3237
3238 let current_ids = registry
3239 .as_ref()
3240 .map(|registry| {
3241 registry
3242 .list()
3243 .into_iter()
3244 .map(|info| info.id)
3245 .collect::<HashSet<_>>()
3246 })
3247 .unwrap_or_default();
3248 let removal_count = current_ids.difference(&target_ids).count();
3249 let additions = prepared
3250 .iter()
3251 .filter(|(id, _)| !current_ids.contains(id))
3252 .cloned()
3253 .collect::<Vec<_>>();
3254
3255 let mut existing = Vec::new();
3256 if let Some(registry) = registry.as_ref() {
3257 for (id, _) in prepared.iter().filter(|(id, _)| current_ids.contains(id)) {
3258 let agent = registry.get(id).ok_or_else(|| {
3259 AgentError::Config(format!("Retained child disappeared during restore: {id}"))
3260 })?;
3261 let child_storage = agent.storage().ok_or_else(|| {
3262 AgentError::Config(format!("Child '{id}' has no storage for session restore"))
3263 })?;
3264 let stored = Self::load_session_restore(child_storage.as_ref(), session_id)
3265 .await?
3266 .ok_or_else(|| {
3267 AgentError::Persistence(format!(
3268 "Child '{id}' has no saved session '{session_id}'"
3269 ))
3270 })?;
3271 existing.push((id.clone(), agent, stored));
3272 }
3273 }
3274
3275 let mut staged = Vec::with_capacity(additions.len());
3276 if !additions.is_empty() {
3277 let spawner = spawner
3278 .as_ref()
3279 .expect("restored additions require a spawner");
3280 let reservations = spawner.reserve_restore_capacity(additions.len(), removal_count)?;
3281 for ((id, spec), reservation) in additions.into_iter().zip(reservations) {
3282 let spawned = spawner
3283 .spawn_with_reserved_capacity(id.clone(), spec, reservation)
3284 .await?;
3285 let child_storage = spawned.agent.storage().ok_or_else(|| {
3286 AgentError::Config(format!("Child '{id}' has no storage for session restore"))
3287 })?;
3288 let stored = Self::load_session_restore(child_storage.as_ref(), session_id)
3289 .await?
3290 .ok_or_else(|| {
3291 AgentError::Persistence(format!(
3292 "Child '{id}' has no saved session '{session_id}'"
3293 ))
3294 })?;
3295 staged.push((spawned, stored));
3296 }
3297 } else if let Some(spawner) = spawner.as_ref() {
3298 spawner.reserve_restore_capacity(0, removal_count)?;
3299 }
3300
3301 let parent_before = self.capture_session_restore_point().await?;
3302 let mut existing_before = Vec::with_capacity(existing.len());
3303 for (id, agent, _) in &existing {
3304 existing_before.push((
3305 id.clone(),
3306 Arc::clone(agent),
3307 agent.capture_session_restore_point().await?,
3308 ));
3309 }
3310
3311 for (_, agent, stored) in &existing {
3315 if let Err(error) = agent
3316 .apply_session_restore_unchecked(session_id, stored.clone())
3317 .await
3318 {
3319 drop(staged);
3320 let rollback_errors =
3321 Self::rollback_session_restore_set(None, &existing_before).await;
3322 return Err(Self::restore_failure(error, rollback_errors));
3323 }
3324 }
3325 for (spawned, stored) in &staged {
3326 if let Err(error) = spawned
3327 .agent
3328 .apply_session_restore_unchecked(session_id, stored.clone())
3329 .await
3330 {
3331 drop(staged);
3332 let rollback_errors =
3333 Self::rollback_session_restore_set(None, &existing_before).await;
3334 return Err(Self::restore_failure(error, rollback_errors));
3335 }
3336 }
3337 if let Err(error) = self
3338 .apply_session_restore_unchecked(session_id, target_parent)
3339 .await
3340 {
3341 drop(staged);
3342 let rollback_errors =
3343 Self::rollback_session_restore_set(Some((self, &parent_before)), &existing_before)
3344 .await;
3345 return Err(Self::restore_failure(error, rollback_errors));
3346 }
3347
3348 if let Some(registry) = registry.as_ref()
3349 && let Err(error) = registry
3350 .reconcile(
3351 &target_ids,
3352 staged.into_iter().map(|(spawned, _)| spawned).collect(),
3353 )
3354 .await
3355 {
3356 let rollback_errors =
3357 Self::rollback_session_restore_set(Some((self, &parent_before)), &existing_before)
3358 .await;
3359 return Err(Self::restore_failure(error, rollback_errors));
3360 }
3361
3362 Ok(target_ids.len())
3363 }
3364
3365 pub async fn delete_session(&self, session_id: &str) -> Result<()> {
3366 let storage = self.storage.read().clone();
3367 match storage {
3368 Some(s) => s.delete(session_id).await,
3369 None => Err(AgentError::Config(
3370 "No storage configured. Use with_storage_config() or with_storage() first".into(),
3371 )),
3372 }
3373 }
3374
3375 pub async fn list_sessions(&self) -> Result<Vec<String>> {
3376 let storage = self.storage.read().clone();
3377 match storage {
3378 Some(s) => s.list_sessions().await,
3379 None => Err(AgentError::Config(
3380 "No storage configured. Use with_storage_config() or with_storage() first".into(),
3381 )),
3382 }
3383 }
3384
3385 fn estimate_tokens(&self, text: &str) -> u32 {
3386 (text.len() as f32 / 4.0).ceil() as u32
3387 }
3388
3389 fn estimate_total_tokens(&self, messages: &[ChatMessage]) -> u32 {
3390 messages
3391 .iter()
3392 .map(|m| self.estimate_tokens(&m.content))
3393 .sum()
3394 }
3395
3396 fn native_safe_prefix_at_least(messages: &[ChatMessage], required: usize) -> Result<usize> {
3398 let inspection =
3399 inspect_native_history(messages).map_err(|error| AgentError::LLM(error.to_string()))?;
3400 let has_signed_history = !inspection.exchanges().is_empty();
3401 for count in required.min(messages.len())..=messages.len() {
3402 if has_signed_history
3403 && count < messages.len()
3404 && messages[count].role != ai_agents_core::Role::User
3405 {
3406 continue;
3407 }
3408 if inspection.is_safe_prefix_len(count)
3409 && inspect_native_history(&messages[count..]).is_ok()
3410 {
3411 return Ok(count);
3412 }
3413 }
3414 Err(AgentError::LLM(
3415 "Context limits cannot remove a complete native history prefix".to_string(),
3416 ))
3417 }
3418
3419 fn truncate_context(&self, messages: &mut Vec<ChatMessage>, keep_recent: usize) -> Result<()> {
3421 if messages.len() <= keep_recent + 1 {
3422 return Ok(());
3423 }
3424 let system_msg = messages.remove(0);
3425 let required = messages.len().saturating_sub(keep_recent);
3426 let to_remove = Self::native_safe_prefix_at_least(messages, required)?;
3427 messages.drain(..to_remove);
3428 messages.insert(0, system_msg);
3429 Ok(())
3430 }
3431
3432 fn get_filter(&self, config: &FilterConfig) -> Arc<dyn MessageFilter> {
3433 match config {
3434 FilterConfig::KeepRecent(n) => Arc::new(KeepRecentFilter::new(*n)),
3435 FilterConfig::ByRole { keep_roles } => Arc::new(ByRoleFilter::new(keep_roles.clone())),
3436 FilterConfig::SkipPattern { skip_if_contains } => {
3437 Arc::new(SkipPatternFilter::new(skip_if_contains.clone()))
3438 }
3439 FilterConfig::Custom { name } => {
3440 let filters = self.message_filters.read();
3441 filters
3442 .get(name)
3443 .cloned()
3444 .unwrap_or_else(|| Arc::new(KeepRecentFilter::new(10)))
3445 }
3446 }
3447 }
3448
3449 async fn summarize_context(
3450 &self,
3451 messages: &mut Vec<ChatMessage>,
3452 summarizer_llm: Option<&str>,
3453 max_summary_tokens: u32,
3454 custom_prompt: Option<&str>,
3455 keep_recent: usize,
3456 filter: Option<&FilterConfig>,
3457 ) -> Result<()> {
3458 let system_msg = messages.remove(0);
3459
3460 let required = messages.len().saturating_sub(keep_recent);
3461 if required == 0 {
3462 messages.insert(0, system_msg);
3463 return Ok(());
3464 }
3465 let to_summarize_count = Self::native_safe_prefix_at_least(messages, required)?;
3466
3467 let recent_msgs: Vec<ChatMessage> = messages.drain(to_summarize_count..).collect();
3468 let mut to_summarize = std::mem::take(messages);
3469
3470 if let Some(filter_config) = filter {
3471 let filter = self.get_filter(filter_config);
3472 to_summarize = filter.filter(to_summarize);
3473 }
3474
3475 if to_summarize.is_empty() {
3476 *messages = recent_msgs;
3477 messages.insert(0, system_msg);
3478 return Ok(());
3479 }
3480
3481 let to_summarize = Self::readable_native_messages(to_summarize)?;
3482 let conversation_text = to_summarize
3483 .iter()
3484 .map(|m| format!("{:?}: {}", m.role, m.content))
3485 .collect::<Vec<_>>()
3486 .join("\n");
3487
3488 let default_prompt = format!(
3489 "Summarize the following conversation in under {} tokens, preserving key information:\n\n{}",
3490 max_summary_tokens, conversation_text
3491 );
3492
3493 let summary_prompt = custom_prompt
3494 .map(|p| format!("{}\n\n{}", p, conversation_text))
3495 .unwrap_or(default_prompt);
3496
3497 let summarizer = if let Some(alias) = summarizer_llm {
3498 self.llm_registry
3499 .get(alias)
3500 .map_err(|e| AgentError::Config(e.to_string()))?
3501 } else {
3502 self.llm_registry
3503 .router()
3504 .or_else(|_| self.llm_registry.default())
3505 .map_err(|e| AgentError::Config(e.to_string()))?
3506 };
3507
3508 let summary_msgs = vec![ChatMessage::user(&summary_prompt)];
3509 let response = self
3510 .observe_purpose(
3511 ObservationPurpose::Summarization,
3512 summarizer.complete(&summary_msgs, None),
3513 )
3514 .await?;
3515
3516 let summary_message = ChatMessage::system(format!(
3517 "[Previous conversation summary]\n{}",
3518 response.content
3519 ));
3520
3521 *messages = vec![system_msg, summary_message];
3522 messages.extend(recent_msgs);
3523
3524 debug!(
3525 summarized_count = to_summarize_count,
3526 kept_recent = keep_recent,
3527 "Context summarized"
3528 );
3529
3530 Ok(())
3531 }
3532
3533 fn render_system_prompt(&self) -> Result<String> {
3534 let mut context = self.build_context_with_overlays();
3535
3536 let facts_text = self.format_actor_facts_for_context();
3538 if !facts_text.is_empty() {
3539 context.insert(
3540 "actor_facts".to_string(),
3541 serde_json::Value::String(facts_text),
3542 );
3543 }
3544
3545 if let Some((key, text)) = self.format_relationship_for_context() {
3546 context.insert(key, serde_json::Value::String(text));
3547 }
3548
3549 self.template_renderer
3550 .render(&self.base_system_prompt, &context)
3551 }
3552
3553 fn canonical_unique_tool_ids(&self, ids: &[String]) -> Vec<String> {
3555 let mut seen = HashSet::new();
3556 ids.iter()
3557 .filter_map(|id| self.tools.canonical_id(id))
3558 .filter(|canonical_id| seen.insert(canonical_id.clone()))
3559 .collect()
3560 }
3561
3562 fn get_top_level_tool_ids_for_scope(&self, scope_override: Option<&[String]>) -> Vec<String> {
3564 let Some(declared) = self.declared_tool_ids.as_deref() else {
3565 return Vec::new();
3566 };
3567 let mut effective = self.canonical_unique_tool_ids(declared);
3568 if let Some(scope) = scope_override {
3569 let scope: HashSet<String> =
3570 self.canonical_unique_tool_ids(scope).into_iter().collect();
3571 effective.retain(|canonical_id| scope.contains(canonical_id));
3572 }
3573 effective
3574 }
3575
3576 async fn get_available_tool_ids(&self) -> Result<Vec<String>> {
3578 Ok(self.get_available_tool_ids_snapshot().await?.tool_ids)
3579 }
3580
3581 async fn get_available_tool_ids_snapshot(&self) -> Result<AvailableToolIdsSnapshot> {
3583 let scope_override = self.runtime_control.tool_scope_override.read().clone();
3584 self.get_available_tool_ids_snapshot_for_scope(scope_override.as_deref())
3585 .await
3586 }
3587
3588 async fn get_available_tool_ids_snapshot_for_scope(
3590 &self,
3591 scope_override: Option<&[String]>,
3592 ) -> Result<AvailableToolIdsSnapshot> {
3593 let mut available = self.get_top_level_tool_ids_for_scope(scope_override);
3594 let (state_generation, state_scopes) = self
3595 .state_machine
3596 .as_ref()
3597 .map(|state_machine| {
3598 let (generation, scopes) = state_machine.current_tool_scope_snapshot();
3599 (Some(generation), scopes)
3600 })
3601 .unwrap_or((None, Vec::new()));
3602
3603 if available.is_empty() || state_scopes.is_empty() {
3604 return Ok(AvailableToolIdsSnapshot {
3605 tool_ids: available,
3606 state_generation,
3607 });
3608 }
3609
3610 let eval_ctx = self.build_evaluation_context().await?;
3611 let llm_getter = RegistryLLMGetter {
3612 registry: self.llm_registry.clone(),
3613 };
3614 let evaluator = ConditionEvaluator::new(llm_getter);
3615
3616 for state_scope in state_scopes {
3617 if state_scope.is_empty() {
3618 available.clear();
3619 break;
3620 }
3621
3622 let mut allowed = HashSet::new();
3623 for tool_ref in &state_scope {
3624 let tool_id = tool_ref.id();
3625 let Some(canonical_id) = self.tools.canonical_id(tool_id) else {
3626 continue;
3627 };
3628 let condition_matches = if let Some(condition) = tool_ref.condition() {
3629 match evaluator.evaluate(condition, &eval_ctx).await {
3630 Ok(matches) => matches,
3631 Err(error) => {
3632 warn!(tool = tool_id, error = %error, "Error evaluating tool condition");
3633 false
3634 }
3635 }
3636 } else {
3637 true
3638 };
3639 if condition_matches {
3640 allowed.insert(canonical_id);
3641 } else {
3642 debug!(tool = tool_id, "Tool condition not met, skipping");
3643 }
3644 }
3645 available.retain(|canonical_id| allowed.contains(canonical_id));
3646 if available.is_empty() {
3647 break;
3648 }
3649 }
3650
3651 Ok(AvailableToolIdsSnapshot {
3652 tool_ids: available,
3653 state_generation,
3654 })
3655 }
3656
3657 async fn build_evaluation_context(&self) -> Result<EvaluationContext> {
3658 let context = self.build_context_with_overlays();
3659 let messages = Self::readable_native_messages(self.memory.get_messages(Some(10)).await?)?;
3660 let tool_history = self.tool_call_history.read().clone();
3661
3662 let (state_name, turn_count, previous_state) = if let Some(ref sm) = self.state_machine {
3663 (Some(sm.current()), sm.turn_count(), sm.previous())
3664 } else {
3665 (None, 0, None)
3666 };
3667
3668 Ok(EvaluationContext::default()
3669 .with_context(context)
3670 .with_state(state_name, turn_count, previous_state)
3671 .with_called_tools(tool_history)
3672 .with_messages(messages))
3673 }
3674
3675 fn record_tool_call(&self, tool_id: &str, result: Value) {
3676 self.tool_call_history.write().push(ToolCallRecord {
3677 tool_id: tool_id.to_string(),
3678 result,
3679 timestamp: chrono::Utc::now(),
3680 });
3681 }
3682
3683 async fn get_effective_system_prompt_with_persona_hooks(
3684 &self,
3685 fire_persona_hooks: bool,
3686 include_tool_prompt: bool,
3687 ) -> Result<String> {
3688 let rendered_base = self.render_system_prompt()?;
3689
3690 let persona_prefix = if let Some(ref persona) = self.persona_manager {
3691 let context = self.build_context_with_overlays();
3692 if fire_persona_hooks {
3693 let render_result = persona.render_prompt(&context)?;
3694 for content in &render_result.newly_revealed {
3695 self.hooks.on_secret_revealed(content).await;
3696 }
3697 render_result.prompt
3698 } else {
3699 persona.render_prompt_preview(&context)?
3700 }
3701 } else {
3702 String::new()
3703 };
3704
3705 if let Some(ref sm) = self.state_machine
3706 && let Some(state_def) = sm.current_definition()
3707 {
3708 let state_prompt = if let Some(ref prompt) = state_def.prompt {
3709 let context = self.build_context_with_overlays();
3710 self.template_renderer.render_with_state(
3711 prompt,
3712 &context,
3713 &sm.current(),
3714 sm.previous().as_deref(),
3715 sm.turn_count(),
3716 state_def.max_turns,
3717 )?
3718 } else {
3719 String::new()
3720 };
3721
3722 let combined = match state_def.prompt_mode {
3723 PromptMode::Append => {
3724 if state_prompt.is_empty() {
3725 rendered_base
3726 } else {
3727 format!(
3728 "{}\n\n[Current State: {}]\n{}",
3729 rendered_base,
3730 sm.current(),
3731 state_prompt
3732 )
3733 }
3734 }
3735 PromptMode::Replace => {
3736 if state_prompt.is_empty() {
3737 rendered_base
3738 } else {
3739 state_prompt
3740 }
3741 }
3742 PromptMode::Prepend => {
3743 if state_prompt.is_empty() {
3744 rendered_base
3745 } else {
3746 format!("{}\n\n{}", state_prompt, rendered_base)
3747 }
3748 }
3749 };
3750
3751 let with_persona = if persona_prefix.is_empty() {
3753 combined
3754 } else {
3755 format!("{}\n\n{}", persona_prefix, combined)
3756 };
3757
3758 if include_tool_prompt {
3759 let available_tool_ids = self.get_available_tool_ids().await?;
3760 if !available_tool_ids.is_empty() {
3761 let tools_prompt = self.tools.generate_scoped_prompt_with_mode(
3762 &available_tool_ids,
3763 None,
3764 self.parallel_tools.enabled,
3765 self.runtime_config.tool_schema_prompt_mode,
3766 );
3767 if !tools_prompt.is_empty() {
3768 return Ok(format!("{}\n\n{}", with_persona, tools_prompt));
3769 }
3770 }
3771 }
3772 return Ok(with_persona);
3773 }
3774
3775 let with_persona = if persona_prefix.is_empty() {
3777 rendered_base
3778 } else {
3779 format!("{}\n\n{}", persona_prefix, rendered_base)
3780 };
3781
3782 if include_tool_prompt {
3783 let available_tool_ids = self.get_available_tool_ids().await?;
3784 let tools_prompt = self.tools.generate_scoped_prompt_with_mode(
3785 &available_tool_ids,
3786 None,
3787 self.parallel_tools.enabled,
3788 self.runtime_config.tool_schema_prompt_mode,
3789 );
3790 if !tools_prompt.is_empty() {
3791 return Ok(format!("{}\n\n{}", with_persona, tools_prompt));
3792 }
3793 }
3794 Ok(with_persona)
3795 }
3796
3797 fn get_state_llm(&self) -> Result<Arc<dyn LLMProvider>> {
3798 if let Some(ref sm) = self.state_machine
3799 && let Some(state_def) = sm.current_definition()
3800 && let Some(ref llm_alias) = state_def.llm
3801 {
3802 return self
3803 .llm_registry
3804 .get(llm_alias)
3805 .map_err(|e| AgentError::Config(e.to_string()));
3806 }
3807 self.llm_registry
3808 .default()
3809 .map_err(|e| AgentError::Config(e.to_string()))
3810 }
3811
3812 fn get_effective_reasoning_config(&self) -> ReasoningConfig {
3813 if let Some(ref sm) = self.state_machine
3814 && let Some(state_def) = sm.current_definition()
3815 && let Some(ref state_reasoning) = state_def.reasoning
3816 {
3817 return state_reasoning.clone();
3818 }
3819 self.reasoning_config.clone()
3820 }
3821
3822 fn get_effective_reflection_config(&self) -> ReflectionConfig {
3823 if let Some(ref sm) = self.state_machine
3824 && let Some(state_def) = sm.current_definition()
3825 && let Some(ref state_reflection) = state_def.reflection
3826 {
3827 return state_reflection.clone();
3828 }
3829 self.reflection_config.clone()
3830 }
3831
3832 fn get_skill_reasoning_config(&self, skill: &SkillDefinition) -> ReasoningConfig {
3833 skill
3834 .reasoning
3835 .clone()
3836 .unwrap_or_else(|| self.get_effective_reasoning_config())
3837 }
3838
3839 fn get_skill_reflection_config(&self, skill: &SkillDefinition) -> ReflectionConfig {
3840 skill
3841 .reflection
3842 .clone()
3843 .unwrap_or_else(|| self.get_effective_reflection_config())
3844 }
3845
3846 async fn build_disambiguation_context(&self) -> Result<DisambiguationContext> {
3847 let recent_messages: Vec<String> =
3848 Self::readable_native_messages(self.memory.get_messages(Some(5)).await?)?
3849 .iter()
3850 .rev()
3851 .map(|m| format!("{:?}: {}", m.role, m.content))
3852 .collect();
3853
3854 let current_state = self.current_state().map(|s| s.to_string());
3855
3856 let state_prompt: Option<String> = self
3859 .state_machine
3860 .as_ref()
3861 .and_then(|sm| sm.current_definition())
3862 .and_then(|def| def.prompt.clone());
3863
3864 let available_tools: Vec<String> = self
3865 .get_available_tool_ids()
3866 .await
3867 .unwrap_or_else(|_| self.tools.list_ids());
3868
3869 let available_skills: Vec<String> = self.skills.iter().map(|s| s.id.clone()).collect();
3870
3871 let mut user_context = self.build_context_with_overlays();
3872 user_context.remove(DISAMBIGUATION_STATE_GENERATION_KEY);
3873 if let Some(state_generation) = self
3874 .state_machine
3875 .as_ref()
3876 .map(|state_machine| state_machine.generation())
3877 {
3878 user_context.insert(
3879 DISAMBIGUATION_STATE_GENERATION_KEY.to_string(),
3880 serde_json::json!(state_generation),
3881 );
3882 }
3883
3884 let available_intents: Vec<String> = if let Some(ref sm) = self.state_machine {
3886 sm.current_definition()
3887 .map(|def| {
3888 def.transitions
3889 .iter()
3890 .filter_map(|t| t.intent.clone())
3891 .collect()
3892 })
3893 .unwrap_or_default()
3894 } else {
3895 Vec::new()
3896 };
3897
3898 Ok(DisambiguationContext::from_agent_state(
3899 recent_messages,
3900 current_state,
3901 state_prompt,
3902 available_tools,
3903 available_skills,
3904 available_intents,
3905 user_context,
3906 ))
3907 }
3908
3909 fn get_available_skills(&self) -> Vec<&SkillDefinition> {
3910 if let Some(ref sm) = self.state_machine
3911 && let Some(state_def) = sm.current_definition()
3912 {
3913 let parent_def = sm.get_parent_definition();
3914 let effective_skills = state_def.get_effective_skills(parent_def.as_ref());
3915 if !effective_skills.is_empty() {
3916 return self
3917 .skills
3918 .iter()
3919 .filter(|s| effective_skills.contains(&&s.id))
3920 .collect();
3921 }
3922 }
3923 self.skills.iter().collect()
3924 }
3925
3926 async fn build_messages(&self) -> Result<Vec<ChatMessage>> {
3927 self.build_messages_internal(true, None, true).await
3928 }
3929
3930 async fn build_messages_for_draft(&self, user_message: &str) -> Result<Vec<ChatMessage>> {
3931 self.build_messages_internal(false, Some(user_message), true)
3932 .await
3933 }
3934
3935 async fn build_messages_internal(
3936 &self,
3937 fire_persona_hooks: bool,
3938 ephemeral_user_message: Option<&str>,
3939 include_tool_prompt: bool,
3940 ) -> Result<Vec<ChatMessage>> {
3941 let system_prompt = self
3942 .get_effective_system_prompt_with_persona_hooks(fire_persona_hooks, include_tool_prompt)
3943 .await?;
3944 let mut messages = vec![ChatMessage::system(&system_prompt)];
3945
3946 let context = self.memory.get_context().await?;
3947 let history = if let Some(ref budget) = self.memory_token_budget {
3948 context.to_llm_messages_with_allocation(&budget.allocation)
3949 } else {
3950 context.to_llm_messages()
3951 };
3952 messages.extend(history);
3953 if let Some(user_message) = ephemeral_user_message {
3954 messages.push(ChatMessage::user(user_message));
3955 }
3956
3957 let total_tokens = self.estimate_total_tokens(&messages);
3958
3959 if total_tokens > self.max_context_tokens {
3960 debug!(
3961 total = total_tokens,
3962 limit = self.max_context_tokens,
3963 "Context overflow"
3964 );
3965
3966 match &self.recovery_manager.config().llm.on_context_overflow {
3967 ContextOverflowAction::Error => {
3968 return Err(AgentError::LLM(format!(
3969 "Context overflow: {} tokens > {} limit",
3970 total_tokens, self.max_context_tokens
3971 )));
3972 }
3973 ContextOverflowAction::Truncate { keep_recent } => {
3974 self.truncate_context(&mut messages, *keep_recent)?;
3975 }
3976 ContextOverflowAction::Summarize {
3977 summarizer_llm,
3978 max_summary_tokens,
3979 custom_prompt,
3980 keep_recent,
3981 filter,
3982 } => {
3983 self.summarize_context(
3984 &mut messages,
3985 summarizer_llm.as_deref(),
3986 *max_summary_tokens,
3987 custom_prompt.as_deref(),
3988 *keep_recent,
3989 filter.as_ref(),
3990 )
3991 .await?;
3992 }
3993 }
3994 }
3995
3996 self.validate_active_native_history(&messages, true)?;
3997 Ok(messages)
3998 }
3999
4000 async fn main_tool_protocol(
4001 &self,
4002 llm: &dyn LLMProvider,
4003 ephemeral_new_turn: bool,
4004 ) -> Result<MainToolProtocol> {
4005 let mut choice = llm.configured_tool_choice();
4006 if matches!(choice.as_ref(), Some(ToolChoice::None)) {
4007 return Ok(MainToolProtocol {
4008 choice,
4009 tool_ids: Vec::new(),
4010 definitions: Vec::new(),
4011 });
4012 }
4013
4014 let mut tool_ids = self.get_available_tool_ids().await?;
4015 tool_ids.sort();
4016 tool_ids.dedup();
4017 if let Some(ToolChoice::Specific(expected)) = choice.as_ref() {
4018 let canonical = self.tools.canonical_id(expected).ok_or_else(|| {
4019 AgentError::Config(format!(
4020 "specific tool choice '{expected}' is not registered"
4021 ))
4022 })?;
4023 if canonical != *expected {
4024 return Err(AgentError::Config(format!(
4025 "specific tool choice must use canonical ID '{canonical}', not '{expected}'"
4026 )));
4027 }
4028 if !tool_ids.iter().any(|tool_id| tool_id == expected) {
4029 return Err(AgentError::Config(format!(
4030 "specific tool choice '{expected}' is outside the effective tool grant"
4031 )));
4032 }
4033 }
4034 if matches!(
4035 choice.as_ref(),
4036 Some(ToolChoice::Required | ToolChoice::Specific(_))
4037 ) && tool_ids.is_empty()
4038 {
4039 return Err(AgentError::Config(
4040 "required tool choice has no tool inside the effective grant".to_string(),
4041 ));
4042 }
4043 if !ephemeral_new_turn
4044 && let Some(configured_choice) = choice.as_ref()
4045 && matches!(
4046 configured_choice,
4047 ToolChoice::Required | ToolChoice::Specific(_)
4048 )
4049 && self
4050 .tool_choice_satisfied_in_current_turn(configured_choice, &tool_ids)
4051 .await?
4052 {
4053 choice = Some(ToolChoice::Auto);
4054 }
4055 if let Some(ToolChoice::Specific(expected)) = choice.as_ref() {
4056 tool_ids.retain(|tool_id| tool_id == expected);
4057 }
4058
4059 let definitions = tool_ids
4060 .iter()
4061 .map(|tool_id| {
4062 let tool = self.tools.get(tool_id).ok_or_else(|| {
4063 AgentError::Config(format!(
4064 "effective tool '{tool_id}' disappeared before provider exposure"
4065 ))
4066 })?;
4067 Ok(LLMToolDefinition {
4068 name: tool_id.clone(),
4069 description: tool.description().to_string(),
4070 input_schema: tool.input_schema(),
4071 })
4072 })
4073 .collect::<Result<Vec<_>>>()?;
4074
4075 Ok(MainToolProtocol {
4079 choice,
4080 tool_ids,
4081 definitions,
4082 })
4083 }
4084
4085 async fn tool_choice_satisfied_in_current_turn(
4086 &self,
4087 choice: &ToolChoice,
4088 effective_tool_ids: &[String],
4089 ) -> Result<bool> {
4090 let messages = self.memory.get_messages(None).await?;
4091 let mut saw_tool_result = false;
4092 for message in messages.iter().rev() {
4093 match message.role {
4094 ai_agents_core::Role::Tool | ai_agents_core::Role::Function => {
4095 saw_tool_result = true;
4096 }
4097 ai_agents_core::Role::Assistant if saw_tool_result => {
4098 let Some(calls) = self.parse_tool_calls(&message.content)? else {
4099 continue;
4100 };
4101 let calls_are_effective = !calls.is_empty()
4102 && calls.iter().all(|call| {
4103 self.tools
4104 .canonical_id(&call.name)
4105 .is_some_and(|canonical| effective_tool_ids.contains(&canonical))
4106 });
4107 return Ok(calls_are_effective
4108 && match choice {
4109 ToolChoice::Required => true,
4110 ToolChoice::Specific(expected) => calls.iter().all(|call| {
4111 self.tools.canonical_id(&call.name).as_deref()
4112 == Some(expected.as_str())
4113 }),
4114 _ => false,
4115 });
4116 }
4117 ai_agents_core::Role::User => return Ok(false),
4118 _ => {}
4119 }
4120 }
4121 Ok(false)
4122 }
4123
4124 fn provider_can_use_native_tools(
4125 &self,
4126 llm: &dyn LLMProvider,
4127 protocol: &MainToolProtocol,
4128 ) -> bool {
4129 let Some(choice) = protocol.choice.as_ref() else {
4130 return false;
4131 };
4132 if matches!(choice, ToolChoice::None) || protocol.definitions.is_empty() {
4133 return false;
4134 }
4135 llm.supports_tool_choice(choice)
4136 && protocol.definitions.iter().all(|definition| {
4137 !definition.name.is_empty()
4138 && definition.name.len() <= 64
4139 && definition
4140 .name
4141 .bytes()
4142 .all(|byte| byte.is_ascii_alphanumeric() || matches!(byte, b'_' | b'-'))
4143 })
4144 }
4145
4146 fn prompt_messages_for_tool_protocol(
4147 &self,
4148 messages: &[ChatMessage],
4149 protocol: &MainToolProtocol,
4150 corrective: bool,
4151 ) -> Vec<ChatMessage> {
4152 let mut messages = messages.to_vec();
4153 let Some(choice) = protocol.choice.as_ref() else {
4154 return messages;
4155 };
4156 if matches!(choice, ToolChoice::None) || protocol.tool_ids.is_empty() {
4157 return messages;
4158 }
4159
4160 let mut tool_prompt = self.tools.generate_scoped_prompt_with_mode(
4161 &protocol.tool_ids,
4162 None,
4163 self.parallel_tools.enabled,
4164 self.runtime_config.tool_schema_prompt_mode,
4165 );
4166 match choice {
4167 ToolChoice::Required => tool_prompt.push_str(
4168 "\n\nYou must call at least one listed tool before giving a final answer.",
4169 ),
4170 ToolChoice::Specific(tool_id) => tool_prompt.push_str(&format!(
4171 "\n\nYou must call the '{tool_id}' tool before giving a final answer."
4172 )),
4173 ToolChoice::Auto => {}
4174 ToolChoice::None => return messages,
4175 _ => return messages,
4176 }
4177 if let Some(system) = messages
4178 .iter_mut()
4179 .find(|message| message.role == ai_agents_core::Role::System)
4180 {
4181 system.content.push_str("\n\n");
4182 system.content.push_str(&tool_prompt);
4183 } else {
4184 messages.insert(0, ChatMessage::system(tool_prompt));
4185 }
4186 if corrective {
4187 let instruction = match choice {
4188 ToolChoice::Required => {
4189 "Your previous response did not call a required tool. Call at least one listed tool now and return only the JSON tool call."
4190 }
4191 ToolChoice::Specific(tool_id) => {
4192 messages.push(ChatMessage::user(format!(
4193 "Your previous response did not call the required '{tool_id}' tool. Call it now and return only the JSON tool call."
4194 )));
4195 return messages;
4196 }
4197 _ => return messages,
4198 };
4199 messages.push(ChatMessage::user(instruction));
4200 }
4201 messages
4202 }
4203
4204 async fn invoke_main_provider(
4205 &self,
4206 llm: Arc<dyn LLMProvider>,
4207 messages: &[ChatMessage],
4208 protocol: &MainToolProtocol,
4209 corrective: bool,
4210 ) -> std::result::Result<MainProviderResponse, LLMError> {
4211 let use_native = self.provider_can_use_native_tools(llm.as_ref(), protocol);
4212 let response = if use_native {
4213 let request = LLMToolRequest {
4214 tools: protocol.definitions.clone(),
4215 choice: protocol
4216 .choice
4217 .clone()
4218 .expect("native tool requests require an explicit choice"),
4219 };
4220 self.observe_purpose(
4221 ObservationPurpose::MainResponse,
4222 llm.complete_with_tools(messages, None, &request),
4223 )
4224 .await?
4225 } else {
4226 let prompt_messages =
4227 self.prompt_messages_for_tool_protocol(messages, protocol, corrective);
4228 self.observe_purpose(
4229 ObservationPurpose::MainResponse,
4230 llm.complete(&prompt_messages, None),
4231 )
4232 .await?
4233 };
4234 Ok(MainProviderResponse {
4235 response,
4236 used_native_tools: use_native,
4237 })
4238 }
4239
4240 async fn complete_main_attempt_with_recovery(
4241 &self,
4242 llm: Arc<dyn LLMProvider>,
4243 messages: &[ChatMessage],
4244 protocol: &MainToolProtocol,
4245 corrective: bool,
4246 ) -> Result<MainProviderResponse> {
4247 let primary_result = self
4249 .recovery_manager
4250 .with_llm_retry(
4251 "llm_call",
4252 None,
4253 || {
4254 let llm = Arc::clone(&llm);
4255 async move {
4256 self.invoke_main_provider(llm, messages, protocol, corrective)
4257 .await
4258 }
4259 },
4260 |error| llm.is_terminal_error(error),
4261 )
4262 .await;
4263
4264 match primary_result {
4265 Ok(response) => Ok(response),
4266 Err(ai_agents_recovery::RetryFailure::Terminal { error, .. }) => {
4267 Err(AgentError::LLM(error.to_string()))
4268 }
4269 Err(failure) => {
4270 let primary_error = AgentError::LLM(failure.into_error().to_string());
4271 match &self.recovery_manager.config().llm.on_failure {
4272 LLMFailureAction::FallbackLlm { fallback_llm } => {
4273 let fallback = self.llm_registry.get(fallback_llm).map_err(|error| {
4274 AgentError::Config(format!(
4275 "Fallback LLM '{fallback_llm}' not found: {error}"
4276 ))
4277 })?;
4278 self.invoke_main_provider(fallback, messages, protocol, corrective)
4279 .await
4280 .map_err(|error| AgentError::LLM(error.to_string()))
4281 }
4282 LLMFailureAction::FallbackResponse { message } => {
4283 if matches!(
4284 protocol.choice.as_ref(),
4285 Some(ToolChoice::Required | ToolChoice::Specific(_))
4286 ) {
4287 Err(AgentError::LLM(format!(
4288 "Required tool selection failed and cannot be satisfied by a static fallback response: {primary_error}"
4289 )))
4290 } else {
4291 Ok(MainProviderResponse {
4292 response: LLMResponse::new(message.clone(), FinishReason::Stop),
4293 used_native_tools: false,
4294 })
4295 }
4296 }
4297 LLMFailureAction::Error => Err(primary_error),
4298 }
4299 }
4300 }
4301 }
4302
4303 fn normalize_main_provider_response(
4304 &self,
4305 mut response: LLMResponse,
4306 protocol: &MainToolProtocol,
4307 ) -> Result<(LLMResponse, bool)> {
4308 let provider_state = response
4309 .take_provider_state()
4310 .map_err(|error| AgentError::LLM(error.to_string()))?;
4311 let native_calls = response
4312 .tool_calls()
4313 .map_err(|error| AgentError::LLM(error.to_string()))?;
4314 let calls = match native_calls {
4315 Some(calls) => {
4316 response.content = encode_native_tool_call_markers(&calls, provider_state.as_ref())
4317 .map_err(|error| AgentError::LLM(error.to_string()))?;
4318 Some(calls)
4319 }
4320 None if provider_state.is_some() => {
4321 return Err(AgentError::LLM(
4322 "Provider returned replay state without native tool calls".to_string(),
4323 ));
4324 }
4325 None if !matches!(protocol.choice.as_ref(), Some(ToolChoice::None)) => {
4326 self.parse_tool_calls(response.content.trim())?
4327 }
4328 None => None,
4329 };
4330
4331 if protocol.choice.is_some()
4332 && let Some(calls) = calls.as_ref()
4333 && calls.iter().any(|call| {
4334 self.tools
4335 .canonical_id(&call.name)
4336 .is_none_or(|canonical| !protocol.tool_ids.contains(&canonical))
4337 })
4338 {
4339 return Err(AgentError::LLM(
4340 "Provider returned a tool call outside the effective grant".to_string(),
4341 ));
4342 }
4343
4344 let compliant = match protocol.choice.as_ref() {
4345 Some(ToolChoice::Required) => calls.as_ref().is_some_and(|calls| !calls.is_empty()),
4346 Some(ToolChoice::Specific(expected)) => calls.as_ref().is_some_and(|calls| {
4347 !calls.is_empty()
4348 && calls.iter().all(|call| {
4349 self.tools.canonical_id(&call.name).as_deref() == Some(expected.as_str())
4350 })
4351 }),
4352 _ => true,
4353 };
4354 Ok((response, compliant))
4355 }
4356
4357 async fn complete_main_llm_with_recovery(
4358 &self,
4359 llm: Arc<dyn LLMProvider>,
4360 messages: &[ChatMessage],
4361 protocol: &MainToolProtocol,
4362 ) -> Result<LLMResponse> {
4363 let first = self
4364 .complete_main_attempt_with_recovery(Arc::clone(&llm), messages, protocol, false)
4365 .await?;
4366 let (response, compliant) =
4367 self.normalize_main_provider_response(first.response, protocol)?;
4368 if compliant {
4369 return Ok(response);
4370 }
4371 if first.used_native_tools {
4372 return Err(AgentError::LLM(
4373 "Provider returned no compliant native call for required tool choice".to_string(),
4374 ));
4375 }
4376
4377 let corrected = self
4378 .complete_main_attempt_with_recovery(llm, messages, protocol, true)
4379 .await?;
4380 let (response, compliant) =
4381 self.normalize_main_provider_response(corrected.response, protocol)?;
4382 if compliant {
4383 return Ok(response);
4384 }
4385 Err(AgentError::LLM(
4386 "Provider returned no compliant tool call after one corrective retry".to_string(),
4387 ))
4388 }
4389
4390 async fn open_main_stream_with_recovery(
4397 &self,
4398 llm: Arc<dyn LLMProvider>,
4399 messages: &[ChatMessage],
4400 protocol: &MainToolProtocol,
4401 ) -> Result<MainStreamSource> {
4402 debug_assert!(
4403 protocol.choice.is_none(),
4404 "streaming raw path must not run with explicit tool choice"
4405 );
4406 let primary = self
4407 .recovery_manager
4408 .with_llm_retry(
4409 "llm_stream_open",
4410 None,
4411 || {
4412 let llm = Arc::clone(&llm);
4413 async move {
4414 self.observe_purpose(
4415 ObservationPurpose::MainResponse,
4416 llm.complete_stream(messages, None),
4417 )
4418 .await
4419 }
4420 },
4421 |error| llm.is_terminal_error(error),
4422 )
4423 .await;
4424
4425 match primary {
4426 Ok(stream) => Ok(MainStreamSource::Stream(stream)),
4427 Err(ai_agents_recovery::RetryFailure::Terminal { error, .. }) => {
4428 Err(AgentError::LLM(error.to_string()))
4429 }
4430 Err(failure) => {
4431 let primary_error = AgentError::LLM(failure.into_error().to_string());
4432 match &self.recovery_manager.config().llm.on_failure {
4433 LLMFailureAction::FallbackLlm { fallback_llm } => {
4434 let fallback = self.llm_registry.get(fallback_llm).map_err(|error| {
4435 AgentError::Config(format!(
4436 "Fallback LLM '{fallback_llm}' not found: {error}"
4437 ))
4438 })?;
4439 if fallback.supports(LLMFeature::Streaming) {
4440 let stream = self
4441 .observe_purpose(
4442 ObservationPurpose::MainResponse,
4443 fallback.complete_stream(messages, None),
4444 )
4445 .await
4446 .map_err(|error| AgentError::LLM(error.to_string()))?;
4447 Ok(MainStreamSource::Stream(stream))
4448 } else {
4449 let response = self
4450 .observe_purpose(
4451 ObservationPurpose::MainResponse,
4452 fallback.complete(messages, None),
4453 )
4454 .await
4455 .map_err(|error| AgentError::LLM(error.to_string()))?;
4456 Ok(MainStreamSource::StaticResponse(response.content))
4457 }
4458 }
4459 LLMFailureAction::FallbackResponse { message } => {
4460 Ok(MainStreamSource::StaticResponse(message.clone()))
4461 }
4462 LLMFailureAction::Error => Err(primary_error),
4463 }
4464 }
4465 }
4466 }
4467
4468 fn is_native_tool_call_content(content: &str) -> Result<bool> {
4470 decode_native_tool_call_markers(content)
4471 .map(|batch| batch.is_some())
4472 .map_err(|error| AgentError::LLM(error.to_string()))
4473 }
4474
4475 fn tool_result_message(
4477 tool_call: &ToolCall,
4478 output: &str,
4479 native_tool_call: bool,
4480 ) -> Result<ChatMessage> {
4481 if !native_tool_call {
4482 return Ok(ChatMessage::function(&tool_call.name, output));
4483 }
4484 let output = serde_json::from_str::<serde_json::Value>(output)
4485 .unwrap_or_else(|_| serde_json::Value::String(output.to_string()));
4486 let content = encode_native_tool_result_marker(tool_call, output)
4487 .map_err(|error| AgentError::LLM(error.to_string()))?;
4488 Ok(ChatMessage::function(&tool_call.name, content))
4489 }
4490
4491 fn remember_active_native_exchange(&self, content: &str) -> Result<()> {
4493 let Some(batch) = decode_native_tool_call_markers(content)
4494 .map_err(|error| AgentError::LLM(error.to_string()))?
4495 else {
4496 return Ok(());
4497 };
4498 let Some(state) = batch.provider_state() else {
4499 return Ok(());
4500 };
4501 let expected = ActiveNativeExchange {
4502 exchange_id: state.exchange_id().to_string(),
4503 call_ids: batch.calls().iter().map(|call| call.id.clone()).collect(),
4504 };
4505 let mut active = self.active_native_exchanges.write();
4506 if let Some(existing) = active
4507 .iter()
4508 .find(|existing| existing.exchange_id == expected.exchange_id)
4509 {
4510 if existing.call_ids != expected.call_ids {
4511 return Err(AgentError::LLM(format!(
4512 "Active native exchange '{}' changed its call identities",
4513 expected.exchange_id
4514 )));
4515 }
4516 } else {
4517 active.push(expected);
4518 }
4519 Ok(())
4520 }
4521
4522 fn validate_active_native_history(
4524 &self,
4525 messages: &[ChatMessage],
4526 require_complete: bool,
4527 ) -> Result<()> {
4528 let expected = self.active_native_exchanges.read().clone();
4529 if expected.is_empty() {
4530 return Ok(());
4531 }
4532 let inspection =
4533 inspect_native_history(messages).map_err(|error| AgentError::LLM(error.to_string()))?;
4534 let expected_count = expected.len();
4535 for (index, expected) in expected.iter().enumerate() {
4536 let Some(exchange) = inspection
4537 .exchanges()
4538 .iter()
4539 .find(|exchange| exchange.state().exchange_id() == expected.exchange_id)
4540 else {
4541 return Err(AgentError::LLM(format!(
4542 "Active native exchange '{}' was removed before provider continuation",
4543 expected.exchange_id
4544 )));
4545 };
4546 let must_be_complete = require_complete || index + 1 < expected_count;
4547 if exchange.call_ids() != expected.call_ids
4548 || (must_be_complete && !exchange.is_complete())
4549 {
4550 return Err(AgentError::LLM(format!(
4551 "Active native exchange '{}' is incomplete before provider continuation",
4552 expected.exchange_id
4553 )));
4554 }
4555 }
4556 Ok(())
4557 }
4558
4559 async fn remember_committed_native_exchange(&self, content: &str) -> Result<()> {
4561 self.remember_active_native_exchange(content)?;
4562 if !self.active_native_exchanges.read().is_empty() {
4563 let messages = self.memory.get_messages(None).await?;
4564 self.validate_active_native_history(&messages, false)?;
4565 }
4566 Ok(())
4567 }
4568
4569 fn readable_native_messages(mut messages: Vec<ChatMessage>) -> Result<Vec<ChatMessage>> {
4571 for message in &mut messages {
4572 if matches!(
4573 message.role,
4574 ai_agents_core::Role::Assistant
4575 | ai_agents_core::Role::Tool
4576 | ai_agents_core::Role::Function
4577 ) {
4578 message.content = native_readable_projection(&message.content)
4579 .map_err(|error| AgentError::LLM(error.to_string()))?;
4580 }
4581 }
4582 Ok(messages)
4583 }
4584
4585 fn parse_main_tool_calls(
4587 &self,
4588 content: &str,
4589 protocol: &MainToolProtocol,
4590 ) -> Result<Option<Vec<ToolCall>>> {
4591 if matches!(protocol.choice.as_ref(), Some(ToolChoice::None)) {
4592 Ok(None)
4593 } else {
4594 self.parse_tool_calls(content)
4595 }
4596 }
4597
4598 fn parse_tool_calls(&self, content: &str) -> Result<Option<Vec<ToolCall>>> {
4600 if let Some(batch) = decode_native_tool_call_markers(content)
4601 .map_err(|error| AgentError::LLM(error.to_string()))?
4602 {
4603 return Ok(Some(batch.into_parts().0));
4604 }
4605 if let Ok(parsed) = serde_json::from_str::<serde_json::Value>(content) {
4607 if let Some(arr) = parsed.as_array() {
4609 let calls: Vec<ToolCall> = arr
4610 .iter()
4611 .filter_map(|v| self.extract_tool_call_from_value(v))
4612 .collect();
4613 if !calls.is_empty() {
4614 return Ok(Some(calls));
4615 }
4616 }
4617 if let Some(tool_call) = self.extract_tool_call_from_value(&parsed) {
4619 return Ok(Some(vec![tool_call]));
4620 }
4621 }
4622
4623 if let Some(json_str) = self.extract_json_from_content(content)
4625 && let Ok(parsed) = serde_json::from_str::<serde_json::Value>(&json_str)
4626 {
4627 if let Some(arr) = parsed.as_array() {
4629 let calls: Vec<ToolCall> = arr
4630 .iter()
4631 .filter_map(|v| self.extract_tool_call_from_value(v))
4632 .collect();
4633 if !calls.is_empty() {
4634 return Ok(Some(calls));
4635 }
4636 }
4637 if let Some(tool_call) = self.extract_tool_call_from_value(&parsed) {
4639 return Ok(Some(vec![tool_call]));
4640 }
4641 }
4642
4643 Ok(None)
4644 }
4645
4646 fn extract_tool_call_from_value(&self, parsed: &serde_json::Value) -> Option<ToolCall> {
4647 if let Some(tool_name) = parsed.get("tool").and_then(|v| v.as_str()) {
4648 let arguments = parsed
4649 .get("arguments")
4650 .cloned()
4651 .unwrap_or(serde_json::json!({}));
4652 return Some(ToolCall {
4653 id: parsed
4654 .get("id")
4655 .and_then(|value| value.as_str())
4656 .filter(|id| !id.is_empty())
4657 .map(str::to_string)
4658 .unwrap_or_else(|| uuid::Uuid::new_v4().to_string()),
4659 name: tool_name.to_string(),
4660 arguments,
4661 });
4662 }
4663 None
4664 }
4665
4666 fn extract_json_from_content(&self, content: &str) -> Option<String> {
4668 if let Some(result) = self.extract_json_array_from_content(content) {
4670 return Some(result);
4671 }
4672 self.extract_json_object_from_content(content)
4673 }
4674
4675 fn extract_json_array_from_content(&self, content: &str) -> Option<String> {
4677 let start = content.find('[')?;
4678 let content_from_start = &content[start..];
4679
4680 let mut depth = 0;
4681 let mut end = 0;
4682 for (i, ch) in content_from_start.char_indices() {
4683 match ch {
4684 '[' => depth += 1,
4685 ']' => {
4686 depth -= 1;
4687 if depth == 0 {
4688 end = i + 1;
4689 break;
4690 }
4691 }
4692 _ => {}
4693 }
4694 }
4695
4696 if end > 0 {
4697 let json_str = &content_from_start[..end];
4698 if json_str.contains("\"tool\"") {
4700 return Some(json_str.to_string());
4701 }
4702 }
4703
4704 None
4705 }
4706
4707 fn extract_json_object_from_content(&self, content: &str) -> Option<String> {
4709 let start = content.find('{')?;
4710 let content_from_start = &content[start..];
4711
4712 let mut depth = 0;
4714 let mut end = 0;
4715 for (i, ch) in content_from_start.char_indices() {
4716 match ch {
4717 '{' => depth += 1,
4718 '}' => {
4719 depth -= 1;
4720 if depth == 0 {
4721 end = i + 1;
4722 break;
4723 }
4724 }
4725 _ => {}
4726 }
4727 }
4728
4729 if end > 0 {
4730 let json_str = &content_from_start[..end];
4731 if json_str.contains("\"tool\"") {
4733 return Some(json_str.to_string());
4734 }
4735 }
4736
4737 None
4738 }
4739
4740 #[allow(clippy::too_many_arguments)]
4744 fn record_from_parts(
4745 &self,
4746 request: &ToolExecutionRequest,
4747 canonical_id: String,
4748 executed_arguments: Value,
4749 started_at: chrono::DateTime<chrono::Utc>,
4750 start: Instant,
4751 executed: bool,
4752 success: bool,
4753 output: String,
4754 metadata: HashMap<String, Value>,
4755 policy: ToolPolicyDecisionRecord,
4756 approval: Option<ToolApprovalRecord>,
4757 timed_out: bool,
4758 output_truncated: bool,
4759 ) -> ToolExecutionRecord {
4760 let versions = ToolDecisionVersions {
4761 policy: self.active_tool_security().policy_version(),
4762 registry: self.tools.version(),
4763 runtime_control: self.runtime_control.version.load(Ordering::SeqCst),
4764 state: self
4765 .state_machine
4766 .as_ref()
4767 .map(|state_machine| state_machine.generation()),
4768 };
4769 self.record_from_parts_at(
4770 request,
4771 canonical_id,
4772 executed_arguments,
4773 started_at,
4774 start,
4775 executed,
4776 success,
4777 output,
4778 metadata,
4779 policy,
4780 approval,
4781 timed_out,
4782 output_truncated,
4783 versions,
4784 )
4785 }
4786
4787 #[allow(clippy::too_many_arguments)]
4789 fn record_from_parts_at(
4790 &self,
4791 request: &ToolExecutionRequest,
4792 canonical_id: String,
4793 executed_arguments: Value,
4794 started_at: chrono::DateTime<chrono::Utc>,
4795 start: Instant,
4796 executed: bool,
4797 success: bool,
4798 output: String,
4799 metadata: HashMap<String, Value>,
4800 policy: ToolPolicyDecisionRecord,
4801 approval: Option<ToolApprovalRecord>,
4802 timed_out: bool,
4803 output_truncated: bool,
4804 versions: ToolDecisionVersions,
4805 ) -> ToolExecutionRecord {
4806 ToolExecutionRecord {
4807 call_id: request.call_id.clone(),
4808 requested_name: request.requested_name.clone(),
4809 canonical_id,
4810 source: request.source.clone(),
4811 arguments: request.arguments.clone(),
4812 executed_arguments,
4813 policy_version: versions.policy,
4814 registry_version: versions.registry,
4815 runtime_config_version: versions.runtime_control,
4816 executed,
4817 success,
4818 output,
4819 metadata,
4820 policy,
4821 approval,
4822 started_at,
4823 duration_ms: start.elapsed().as_millis() as u64,
4824 timed_out,
4825 cancelled: false,
4826 cancellation_reason: None,
4827 output_truncated,
4828 }
4829 }
4830
4831 async fn finish_tool_record(&self, record: &ToolExecutionRecord) {
4833 let result = ToolResult {
4834 success: record.success,
4835 output: record.model_output_string(),
4836 metadata: if record.metadata.is_empty() {
4837 None
4838 } else {
4839 Some(record.metadata.clone())
4840 },
4841 };
4842 self.hooks
4843 .on_tool_complete(&record.canonical_id, &result, record.duration_ms)
4844 .await;
4845 self.hooks.on_tool_execution_record(record).await;
4846 self.record_tool_call(&record.canonical_id, record.model_output_value());
4847 if !record.success {
4848 self.hooks
4849 .on_error(&AgentError::Tool(record.output.clone()))
4850 .await;
4851 }
4852 }
4853
4854 async fn finish_tool_record_after_resource_guards(
4856 &self,
4857 resource_guards: ToolResourceGuards,
4858 record: &ToolExecutionRecord,
4859 ) {
4860 drop(resource_guards);
4861 self.finish_tool_record(record).await;
4862 }
4863
4864 fn validated_tool_timeout(timeout_ms: u64) -> Result<ValidatedToolTimeout> {
4868 if timeout_ms > MAX_TOOL_TIMEOUT_MS {
4869 return Err(AgentError::Config(format!(
4870 "effective tool timeout_ms must be no greater than {MAX_TOOL_TIMEOUT_MS} milliseconds"
4871 )));
4872 }
4873 let timer = Duration::from_millis(timeout_ms);
4874 let deadline_delta = chrono::Duration::from_std(timer).map_err(|_| {
4875 AgentError::Config(format!(
4876 "effective tool timeout_ms cannot be represented as a UTC deadline: {timeout_ms}"
4877 ))
4878 })?;
4879 Ok(ValidatedToolTimeout {
4880 timer,
4881 deadline_delta,
4882 })
4883 }
4884
4885 fn effective_tool_limits(
4889 security_engine: &ToolSecurityEngine,
4890 canonical_id: &str,
4891 safety: &ToolSafetyMetadata,
4892 classification: &ToolCallClassification,
4893 recovery_timeout_ms: Option<u64>,
4894 ) -> Result<(ToolExecutionLimits, ValidatedToolTimeout)> {
4895 if let Some(timeout_ms) = classification.timeout_ms {
4896 Self::validated_tool_timeout(timeout_ms)?;
4897 }
4898 if let Some(timeout_ms) = recovery_timeout_ms {
4899 Self::validated_tool_timeout(timeout_ms)?;
4900 }
4901
4902 let mut limits = security_engine.effective_limits(canonical_id, safety, classification);
4903 if let Some(recovery_timeout_ms) = recovery_timeout_ms {
4904 limits.timeout_ms = Some(limits.timeout_ms.map_or(recovery_timeout_ms, |timeout_ms| {
4905 timeout_ms.min(recovery_timeout_ms)
4906 }));
4907 }
4908 let timeout_ms = limits
4909 .timeout_ms
4910 .unwrap_or_else(|| security_engine.get_tool_timeout(canonical_id));
4911 let timeout = Self::validated_tool_timeout(timeout_ms)?;
4912 Ok((limits, timeout))
4913 }
4914
4915 async fn execute_resolved_tool_once(
4917 &self,
4918 tool: Arc<dyn ai_agents_core::Tool>,
4919 args: Value,
4920 mut ctx: ToolExecutionContext,
4921 timeout: ValidatedToolTimeout,
4922 ) -> Result<(ToolResult, bool, bool, bool)> {
4923 if self.runtime_control.emergency_deny.load(Ordering::SeqCst) {
4924 return Ok((
4925 ToolResult::error("Tool execution cancelled by runtime control"),
4926 false,
4927 true,
4928 false,
4929 ));
4930 }
4931 ctx.deadline = Some(
4936 chrono::Utc::now()
4937 .checked_add_signed(timeout.deadline_delta)
4938 .ok_or_else(|| {
4939 AgentError::Config(
4940 "effective tool timeout_ms exceeds the current UTC deadline range"
4941 .to_string(),
4942 )
4943 })?,
4944 );
4945 let invoked = Arc::new(AtomicBool::new(false));
4949 let invoked_by_future = Arc::clone(&invoked);
4950 let actor_context = current_turn_actor_context();
4951 let future = async move {
4952 invoked_by_future.store(true, Ordering::SeqCst);
4953 if let Some(actor_context) = actor_context {
4954 scope_actor_context(actor_context, tool.execute(args, ctx)).await
4955 } else {
4956 tool.execute(args, ctx).await
4957 }
4958 };
4959 tokio::pin!(future);
4960 let timer = tokio::time::sleep(timeout.timer);
4961 tokio::pin!(timer);
4962 let mut cancel_tick = tokio::time::interval(std::time::Duration::from_millis(50));
4963
4964 loop {
4965 tokio::select! {
4966 result = &mut future => return Ok((result, false, false, true)),
4967 _ = &mut timer => {
4968 return Ok((
4969 ToolResult::error("Tool execution timed out"),
4970 true,
4971 false,
4972 invoked.load(Ordering::SeqCst),
4973 ));
4974 }
4975 _ = cancel_tick.tick() => {
4976 if self.runtime_control.emergency_deny.load(Ordering::SeqCst) {
4977 return Ok((
4978 ToolResult::error("Tool execution cancelled by runtime control"),
4979 false,
4980 true,
4981 invoked.load(Ordering::SeqCst),
4982 ));
4983 }
4984 }
4985 }
4986 }
4987 }
4988
4989 fn truncate_tool_output(output: String, max_chars: Option<usize>) -> (String, bool) {
4991 let Some(max_chars) = max_chars else {
4992 return (output, false);
4993 };
4994 let mut chars = output.chars();
4995 let truncated: String = chars.by_ref().take(max_chars).collect();
4996 if chars.next().is_some() {
4997 (truncated, true)
4998 } else {
4999 (output, false)
5000 }
5001 }
5002
5003 async fn acquire_tool_resource_locks(&self, keys: &[String]) -> Option<ToolResourceGuards> {
5005 let locks = {
5006 let mut table = self.resource_locks.write();
5007 table.retain(|_, lock| lock.strong_count() > 0);
5008 keys.iter()
5009 .map(|key| {
5010 if let Some(lock) = table.get(key).and_then(Weak::upgrade) {
5011 lock
5012 } else {
5013 let lock = Arc::new(tokio::sync::Mutex::new(()));
5014 table.insert(key.clone(), Arc::downgrade(&lock));
5015 lock
5016 }
5017 })
5018 .collect::<Vec<_>>()
5019 };
5020 let mut resource_guards = ToolResourceGuards {
5021 guards: Vec::with_capacity(locks.len()),
5022 locks: Arc::clone(&self.resource_locks),
5023 };
5024 let mut locks = locks.into_iter();
5025 while let Some(lock) = locks.next() {
5026 let mut lock = Box::pin(lock.lock_owned());
5027 loop {
5028 tokio::select! {
5029 guard = &mut lock => {
5030 resource_guards.guards.push(guard);
5031 break;
5032 }
5033 _ = tokio::time::sleep(std::time::Duration::from_millis(10)) => {
5034 if self.runtime_control.emergency_deny.load(Ordering::SeqCst) {
5035 drop(lock);
5036 drop(locks);
5037 drop(resource_guards);
5038 return None;
5039 }
5040 }
5041 }
5042 }
5043 }
5044 Some(resource_guards)
5045 }
5046
5047 async fn run_tool_with_retries(
5051 &self,
5052 canonical_id: &str,
5053 tool: Arc<dyn ai_agents_core::Tool>,
5054 args: Value,
5055 ctx: ToolExecutionContext,
5056 timeout: ValidatedToolTimeout,
5057 max_retries: u32,
5058 ) -> Result<(ToolResult, bool, bool, bool)> {
5059 let max_retries = if ctx.classification.safely_retryable {
5060 max_retries
5061 } else {
5062 0
5063 };
5064 let mut attempts = 0;
5065 let mut invoked = false;
5066 loop {
5067 let (result, timed_out, cancelled, attempt_invoked) = self
5068 .execute_resolved_tool_once(tool.clone(), args.clone(), ctx.clone(), timeout)
5069 .await?;
5070 invoked |= attempt_invoked;
5071 if result.success || timed_out || cancelled || attempts >= max_retries {
5072 return Ok((result, timed_out, cancelled, invoked));
5073 }
5074 attempts += 1;
5075 warn!(tool = %canonical_id, attempt = attempts, error = %result.output, "Retrying failed tool call");
5076 }
5077 }
5078
5079 fn host_tool_unavailability(&self, canonical_id: &str) -> Option<(&'static str, &'static str)> {
5081 match canonical_id {
5082 "command" if !self.tools.command_runner_available() => Some((
5083 "Command runner is unavailable",
5084 "command runner is unavailable",
5085 )),
5086 "diagnostics" if !self.tools.diagnostics_available() => Some((
5087 "Diagnostics provider is unavailable",
5088 "diagnostics provider is unavailable",
5089 )),
5090 "web_search" if !self.tools.web_search_available() => Some((
5091 "Web search provider is unavailable",
5092 "web search provider is unavailable",
5093 )),
5094 _ => None,
5095 }
5096 }
5097
5098 fn execute_tool_record(
5100 &self,
5101 request: ToolExecutionRequest,
5102 ) -> Pin<Box<dyn Future<Output = Result<ToolExecutionRecord>> + Send + '_>> {
5103 Box::pin(self.execute_tool_record_inner(request, ToolFallbackState::default()))
5104 }
5105
5106 async fn execute_tool_record_inner(
5110 &self,
5111 request: ToolExecutionRequest,
5112 fallback_state: ToolFallbackState,
5113 ) -> Result<ToolExecutionRecord> {
5114 let started_at = chrono::Utc::now();
5115 let start = Instant::now();
5116 info!(tool = %request.requested_name, args = %request.arguments, "Executing tool");
5117
5118 if self.runtime_control.emergency_deny.load(Ordering::SeqCst) {
5119 let record = self.record_from_parts(
5120 &request,
5121 request.requested_name.clone(),
5122 request.arguments.clone(),
5123 started_at,
5124 start,
5125 false,
5126 false,
5127 "Tool execution is disabled by runtime control".to_string(),
5128 HashMap::new(),
5129 ToolPolicyDecisionRecord::deny("runtime emergency deny is enabled"),
5130 None,
5131 false,
5132 false,
5133 );
5134 self.finish_tool_record(&record).await;
5135 return Ok(record);
5136 }
5137
5138 let Some(resolved) = self.tools.resolve(&request.requested_name) else {
5139 let record = self.record_from_parts(
5140 &request,
5141 request.requested_name.clone(),
5142 request.arguments.clone(),
5143 started_at,
5144 start,
5145 false,
5146 false,
5147 format!("Tool '{}' is unavailable", request.requested_name),
5148 HashMap::new(),
5149 ToolPolicyDecisionRecord::unavailable(format!(
5150 "Tool '{}' is not registered",
5151 request.requested_name
5152 )),
5153 None,
5154 false,
5155 false,
5156 );
5157 self.finish_tool_record(&record).await;
5158 return Ok(record);
5159 };
5160
5161 let canonical_id = resolved.identity.canonical_id.clone();
5162
5163 let initial_scope_snapshot = self.get_available_tool_ids_snapshot().await?;
5164 if !initial_scope_snapshot
5165 .tool_ids
5166 .iter()
5167 .any(|id| id == &canonical_id)
5168 {
5169 let record = self.record_from_parts(
5170 &request,
5171 canonical_id.clone(),
5172 request.arguments.clone(),
5173 started_at,
5174 start,
5175 false,
5176 false,
5177 format!(
5178 "Tool '{}' is not available in the current scope",
5179 canonical_id
5180 ),
5181 HashMap::new(),
5182 ToolPolicyDecisionRecord::deny(format!(
5183 "Tool '{}' is not granted by the current top-level and state tool scope",
5184 canonical_id
5185 )),
5186 None,
5187 false,
5188 false,
5189 );
5190 self.finish_tool_record(&record).await;
5191 return Ok(record);
5192 }
5193
5194 let approval_control_snapshot = self.runtime_safety_snapshot();
5195 let security_engine = approval_control_snapshot.tool_security.clone();
5196 if let Some(reason) = fallback_state.rejection_reason(&canonical_id) {
5197 let mut metadata = HashMap::new();
5198 metadata.insert(
5199 "fallback_chain".to_string(),
5200 serde_json::to_value(&fallback_state.visited_canonical_ids).unwrap_or(Value::Null),
5201 );
5202 let record = self.record_from_parts(
5203 &request,
5204 canonical_id,
5205 request.arguments.clone(),
5206 started_at,
5207 start,
5208 false,
5209 false,
5210 format!("Denied: {reason}"),
5211 metadata,
5212 ToolPolicyDecisionRecord::deny(reason),
5213 None,
5214 false,
5215 false,
5216 );
5217 self.finish_tool_record(&record).await;
5218 return Ok(record);
5219 }
5220 let admitted_canonical_id = canonical_id.clone();
5221 let fallback_state = fallback_state.with_current(canonical_id.clone());
5222 let bindings = resolved.tool.policy_bindings();
5223 let mut executed_arguments = security_engine.prepare_tool_arguments_with_bindings(
5224 &canonical_id,
5225 &request.arguments,
5226 &bindings,
5227 );
5228 let mut metadata = HashMap::new();
5229 let safety = resolved.tool.safety_metadata();
5230 let classification = resolved.tool.classify_call(&executed_arguments);
5231 let initial_recovery_timeout_ms = self.recovery_manager.get_tool_timeout(&canonical_id);
5232 let (limits, _) = Self::effective_tool_limits(
5233 &security_engine,
5234 &canonical_id,
5235 &safety,
5236 &classification,
5237 initial_recovery_timeout_ms,
5238 )?;
5239 self.hooks
5240 .on_tool_start(&canonical_id, &executed_arguments)
5241 .await;
5242 metadata.insert(
5243 "classification".to_string(),
5244 serde_json::to_value(&classification).unwrap_or(Value::Null),
5245 );
5246 metadata.insert(
5247 "effective_limits".to_string(),
5248 serde_json::to_value(&limits).unwrap_or(Value::Null),
5249 );
5250 let policy_snapshot = security_engine.policy_snapshot(&canonical_id);
5251 if !policy_snapshot.is_null() {
5252 metadata.insert("policy_snapshot".to_string(), policy_snapshot.clone());
5253 }
5254
5255 let mut approval_record = Some(ToolApprovalRecord {
5256 status: ToolApprovalStatus::NotRequired,
5257 reason: None,
5258 modified_arguments: None,
5259 });
5260
5261 let mut security_result = security_engine
5262 .validate_tool_execution_with_bindings(&canonical_id, &executed_arguments, &bindings)
5263 .await?;
5264 if (security_result.is_allowed()
5269 || matches!(
5270 &security_result,
5271 SecurityCheckResult::RequireConfirmation { .. }
5272 ))
5273 && let Some((output, reason)) = self.host_tool_unavailability(&canonical_id)
5274 {
5275 let record = self.record_from_parts(
5276 &request,
5277 canonical_id,
5278 executed_arguments,
5279 started_at,
5280 start,
5281 false,
5282 false,
5283 output.to_string(),
5284 metadata,
5285 ToolPolicyDecisionRecord::unavailable(reason),
5286 Some(ToolApprovalRecord {
5287 status: ToolApprovalStatus::Unavailable,
5288 reason: Some(reason.to_string()),
5289 modified_arguments: None,
5290 }),
5291 false,
5292 false,
5293 );
5294 self.finish_tool_record(&record).await;
5295 return Ok(record);
5296 }
5297 match &security_result {
5298 SecurityCheckResult::Allow => {}
5299 SecurityCheckResult::Warn { message } => {
5300 warn!(tool = %canonical_id, message = %message, "Tool security warning");
5301 }
5302 SecurityCheckResult::Block { reason } => {
5303 let record = self.record_from_parts(
5304 &request,
5305 canonical_id,
5306 executed_arguments,
5307 started_at,
5308 start,
5309 false,
5310 false,
5311 format!("Denied: {}", reason),
5312 metadata,
5313 ToolPolicyDecisionRecord::deny(reason.clone()),
5314 approval_record,
5315 false,
5316 false,
5317 );
5318 self.finish_tool_record(&record).await;
5319 return Ok(record);
5320 }
5321 SecurityCheckResult::Unavailable { reason } => {
5322 let record = self.record_from_parts(
5323 &request,
5324 canonical_id,
5325 executed_arguments,
5326 started_at,
5327 start,
5328 false,
5329 false,
5330 format!("Unavailable: {}", reason),
5331 metadata,
5332 ToolPolicyDecisionRecord::unavailable(reason.clone()),
5333 approval_record,
5334 false,
5335 false,
5336 );
5337 self.finish_tool_record(&record).await;
5338 return Ok(record);
5339 }
5340 SecurityCheckResult::RequireConfirmation { message } => {
5341 if self.hitl_engine.is_none() {
5342 approval_record = Some(ToolApprovalRecord {
5343 status: ToolApprovalStatus::Unavailable,
5344 reason: Some("No HITL engine configured".to_string()),
5345 modified_arguments: None,
5346 });
5347 let record = self.record_from_parts(
5348 &request,
5349 canonical_id,
5350 executed_arguments,
5351 started_at,
5352 start,
5353 false,
5354 false,
5355 format!("Approval unavailable: {}", message),
5356 metadata,
5357 ToolPolicyDecisionRecord::approval(message.clone()),
5358 approval_record,
5359 false,
5360 false,
5361 );
5362 self.finish_tool_record(&record).await;
5363 return Ok(record);
5364 }
5365
5366 let check_result = HITLCheckResult::required(
5367 ApprovalTrigger::tool(&canonical_id, executed_arguments.clone()),
5368 HashMap::new(),
5369 message.clone(),
5370 None,
5371 );
5372 match self.request_hitl_approval(check_result).await? {
5373 ApprovalResult::Approved => {
5374 merge_approved_record(&mut approval_record);
5375 }
5376 ApprovalResult::Modified { changes } => {
5377 if let Some(obj) = executed_arguments.as_object_mut() {
5378 for (key, value) in changes {
5379 obj.insert(key, value);
5380 }
5381 }
5382 security_result = security_engine
5383 .validate_tool_execution_with_bindings(
5384 &canonical_id,
5385 &executed_arguments,
5386 &bindings,
5387 )
5388 .await?;
5389 if !matches!(
5390 security_result,
5391 SecurityCheckResult::Allow
5392 | SecurityCheckResult::Warn { .. }
5393 | SecurityCheckResult::RequireConfirmation { .. }
5394 ) {
5395 let reason = security_result
5396 .reason()
5397 .unwrap_or("modified arguments failed policy")
5398 .to_string();
5399 let record = self.record_from_parts(
5400 &request,
5401 canonical_id,
5402 executed_arguments.clone(),
5403 started_at,
5404 start,
5405 false,
5406 false,
5407 reason.clone(),
5408 metadata,
5409 ToolPolicyDecisionRecord::deny(reason),
5410 Some(ToolApprovalRecord {
5411 status: ToolApprovalStatus::Modified,
5412 reason: None,
5413 modified_arguments: Some(executed_arguments),
5414 }),
5415 false,
5416 false,
5417 );
5418 self.finish_tool_record(&record).await;
5419 return Ok(record);
5420 }
5421 approval_record = Some(ToolApprovalRecord {
5422 status: ToolApprovalStatus::Modified,
5423 reason: None,
5424 modified_arguments: Some(executed_arguments.clone()),
5425 });
5426 }
5427 ApprovalResult::Rejected { reason } => {
5428 let reason = reason.unwrap_or_else(|| "rejected".to_string());
5429 approval_record = Some(ToolApprovalRecord {
5430 status: ToolApprovalStatus::Rejected,
5431 reason: Some(reason.clone()),
5432 modified_arguments: None,
5433 });
5434 let record = self.record_from_parts(
5435 &request,
5436 canonical_id,
5437 executed_arguments,
5438 started_at,
5439 start,
5440 false,
5441 false,
5442 format!("Approval rejected: {}", reason),
5443 metadata,
5444 ToolPolicyDecisionRecord::approval(reason),
5445 approval_record,
5446 false,
5447 false,
5448 );
5449 self.finish_tool_record(&record).await;
5450 return Ok(record);
5451 }
5452 ApprovalResult::Timeout => {
5453 approval_record = Some(ToolApprovalRecord {
5454 status: ToolApprovalStatus::Timeout,
5455 reason: Some("approval timeout".to_string()),
5456 modified_arguments: None,
5457 });
5458 let record = self.record_from_parts(
5459 &request,
5460 canonical_id,
5461 executed_arguments,
5462 started_at,
5463 start,
5464 false,
5465 false,
5466 "Approval timed out".to_string(),
5467 metadata,
5468 ToolPolicyDecisionRecord::approval("approval timeout"),
5469 approval_record,
5470 false,
5471 false,
5472 );
5473 self.finish_tool_record(&record).await;
5474 return Ok(record);
5475 }
5476 }
5477 }
5478 }
5479
5480 if approval_record
5481 .as_ref()
5482 .is_some_and(|record| matches!(record.status, ToolApprovalStatus::NotRequired))
5483 && let Some(message) =
5484 security_engine.classification_approval_message(&canonical_id, &classification)
5485 {
5486 if self.hitl_engine.is_none() {
5487 approval_record = Some(ToolApprovalRecord {
5488 status: ToolApprovalStatus::Unavailable,
5489 reason: Some("No HITL engine configured".to_string()),
5490 modified_arguments: None,
5491 });
5492 let record = self.record_from_parts(
5493 &request,
5494 canonical_id,
5495 executed_arguments,
5496 started_at,
5497 start,
5498 false,
5499 false,
5500 format!("Approval unavailable: {}", message),
5501 metadata,
5502 ToolPolicyDecisionRecord::approval(message),
5503 approval_record,
5504 false,
5505 false,
5506 );
5507 self.finish_tool_record(&record).await;
5508 return Ok(record);
5509 }
5510 let check_result = HITLCheckResult::required(
5511 ApprovalTrigger::tool(&canonical_id, executed_arguments.clone()),
5512 HashMap::new(),
5513 message.clone(),
5514 None,
5515 );
5516 match self.request_hitl_approval(check_result).await? {
5517 ApprovalResult::Approved => {
5518 merge_approved_record(&mut approval_record);
5519 }
5520 ApprovalResult::Modified { changes } => {
5521 if let Some(obj) = executed_arguments.as_object_mut() {
5522 for (key, value) in changes {
5523 obj.insert(key, value);
5524 }
5525 }
5526 let modified_security = security_engine
5527 .validate_tool_execution_with_bindings(
5528 &canonical_id,
5529 &executed_arguments,
5530 &bindings,
5531 )
5532 .await?;
5533 if !matches!(
5534 modified_security,
5535 SecurityCheckResult::Allow | SecurityCheckResult::Warn { .. }
5536 ) {
5537 let reason = modified_security
5538 .reason()
5539 .unwrap_or("modified arguments failed policy")
5540 .to_string();
5541 let record = self.record_from_parts(
5542 &request,
5543 canonical_id,
5544 executed_arguments.clone(),
5545 started_at,
5546 start,
5547 false,
5548 false,
5549 reason.clone(),
5550 metadata,
5551 ToolPolicyDecisionRecord::deny(reason),
5552 Some(ToolApprovalRecord {
5553 status: ToolApprovalStatus::Modified,
5554 reason: None,
5555 modified_arguments: Some(executed_arguments),
5556 }),
5557 false,
5558 false,
5559 );
5560 self.finish_tool_record(&record).await;
5561 return Ok(record);
5562 }
5563 approval_record = Some(ToolApprovalRecord {
5564 status: ToolApprovalStatus::Modified,
5565 reason: None,
5566 modified_arguments: Some(executed_arguments.clone()),
5567 });
5568 }
5569 ApprovalResult::Rejected { reason } => {
5570 let reason = reason.unwrap_or_else(|| "rejected".to_string());
5571 let record = self.record_from_parts(
5572 &request,
5573 canonical_id,
5574 executed_arguments,
5575 started_at,
5576 start,
5577 false,
5578 false,
5579 format!("Approval rejected: {}", reason),
5580 metadata,
5581 ToolPolicyDecisionRecord::approval(reason.clone()),
5582 Some(ToolApprovalRecord {
5583 status: ToolApprovalStatus::Rejected,
5584 reason: Some(reason),
5585 modified_arguments: None,
5586 }),
5587 false,
5588 false,
5589 );
5590 self.finish_tool_record(&record).await;
5591 return Ok(record);
5592 }
5593 ApprovalResult::Timeout => {
5594 let record = self.record_from_parts(
5595 &request,
5596 canonical_id,
5597 executed_arguments,
5598 started_at,
5599 start,
5600 false,
5601 false,
5602 "Approval timed out".to_string(),
5603 metadata,
5604 ToolPolicyDecisionRecord::approval("approval timeout"),
5605 Some(ToolApprovalRecord {
5606 status: ToolApprovalStatus::Timeout,
5607 reason: Some("approval timeout".to_string()),
5608 modified_arguments: None,
5609 }),
5610 false,
5611 false,
5612 );
5613 self.finish_tool_record(&record).await;
5614 return Ok(record);
5615 }
5616 }
5617 }
5618
5619 let hitl_lang_ctx = self.build_hitl_language_context();
5620 if let Some(ref hitl_engine) = self.hitl_engine {
5621 let check_result = self
5622 .observe_purpose(
5623 ObservationPurpose::HitlLocalization,
5624 hitl_engine.check_tool_with_localization(
5625 &canonical_id,
5626 &executed_arguments,
5627 &hitl_lang_ctx,
5628 self.approval_handler.as_ref(),
5629 Some(&self.llm_registry),
5630 ),
5631 )
5632 .await?;
5633 if check_result.is_required() {
5634 match self.request_hitl_approval(check_result).await? {
5635 ApprovalResult::Approved => {
5636 merge_approved_record(&mut approval_record);
5637 }
5638 ApprovalResult::Modified { changes } => {
5639 if let Some(obj) = executed_arguments.as_object_mut() {
5640 for (key, value) in changes {
5641 obj.insert(key, value);
5642 }
5643 }
5644 let modified_security = security_engine
5645 .validate_tool_execution_with_bindings(
5646 &canonical_id,
5647 &executed_arguments,
5648 &bindings,
5649 )
5650 .await?;
5651 if !matches!(
5652 modified_security,
5653 SecurityCheckResult::Allow | SecurityCheckResult::Warn { .. }
5654 ) {
5655 let reason = modified_security
5656 .reason()
5657 .unwrap_or("modified arguments failed policy")
5658 .to_string();
5659 let record = self.record_from_parts(
5660 &request,
5661 canonical_id,
5662 executed_arguments.clone(),
5663 started_at,
5664 start,
5665 false,
5666 false,
5667 reason.clone(),
5668 metadata,
5669 ToolPolicyDecisionRecord::deny(reason),
5670 Some(ToolApprovalRecord {
5671 status: ToolApprovalStatus::Modified,
5672 reason: None,
5673 modified_arguments: Some(executed_arguments),
5674 }),
5675 false,
5676 false,
5677 );
5678 self.finish_tool_record(&record).await;
5679 return Ok(record);
5680 }
5681 approval_record = Some(ToolApprovalRecord {
5682 status: ToolApprovalStatus::Modified,
5683 reason: None,
5684 modified_arguments: Some(executed_arguments.clone()),
5685 });
5686 }
5687 ApprovalResult::Rejected { reason } => {
5688 let reason = reason.unwrap_or_else(|| "rejected".to_string());
5689 let record = self.record_from_parts(
5690 &request,
5691 canonical_id,
5692 executed_arguments,
5693 started_at,
5694 start,
5695 false,
5696 false,
5697 format!("Approval rejected: {}", reason),
5698 metadata,
5699 ToolPolicyDecisionRecord::approval(reason.clone()),
5700 Some(ToolApprovalRecord {
5701 status: ToolApprovalStatus::Rejected,
5702 reason: Some(reason),
5703 modified_arguments: None,
5704 }),
5705 false,
5706 false,
5707 );
5708 self.finish_tool_record(&record).await;
5709 return Ok(record);
5710 }
5711 ApprovalResult::Timeout => {
5712 let record = self.record_from_parts(
5713 &request,
5714 canonical_id,
5715 executed_arguments,
5716 started_at,
5717 start,
5718 false,
5719 false,
5720 "Approval timed out".to_string(),
5721 metadata,
5722 ToolPolicyDecisionRecord::approval("approval timeout"),
5723 Some(ToolApprovalRecord {
5724 status: ToolApprovalStatus::Timeout,
5725 reason: Some("approval timeout".to_string()),
5726 modified_arguments: None,
5727 }),
5728 false,
5729 false,
5730 );
5731 self.finish_tool_record(&record).await;
5732 return Ok(record);
5733 }
5734 }
5735 }
5736
5737 let condition_check = self
5738 .observe_purpose(
5739 ObservationPurpose::HitlLocalization,
5740 hitl_engine.check_conditions_with_localization(
5741 &executed_arguments,
5742 &hitl_lang_ctx,
5743 self.approval_handler.as_ref(),
5744 Some(&self.llm_registry),
5745 ),
5746 )
5747 .await?;
5748 if condition_check.is_required() {
5749 match self.request_hitl_approval(condition_check).await? {
5750 ApprovalResult::Approved => {
5751 merge_approved_record(&mut approval_record);
5752 }
5753 ApprovalResult::Modified { changes } => {
5754 if let Some(obj) = executed_arguments.as_object_mut() {
5755 for (key, value) in changes {
5756 obj.insert(key, value);
5757 }
5758 }
5759 let modified_security = security_engine
5760 .validate_tool_execution_with_bindings(
5761 &canonical_id,
5762 &executed_arguments,
5763 &bindings,
5764 )
5765 .await?;
5766 if !matches!(
5767 modified_security,
5768 SecurityCheckResult::Allow | SecurityCheckResult::Warn { .. }
5769 ) {
5770 let reason = modified_security
5771 .reason()
5772 .unwrap_or("modified arguments failed policy")
5773 .to_string();
5774 let record = self.record_from_parts(
5775 &request,
5776 canonical_id,
5777 executed_arguments,
5778 started_at,
5779 start,
5780 false,
5781 false,
5782 reason.clone(),
5783 metadata,
5784 ToolPolicyDecisionRecord::deny(reason),
5785 approval_record,
5786 false,
5787 false,
5788 );
5789 self.finish_tool_record(&record).await;
5790 return Ok(record);
5791 }
5792 approval_record = Some(ToolApprovalRecord {
5793 status: ToolApprovalStatus::Modified,
5794 reason: None,
5795 modified_arguments: Some(executed_arguments.clone()),
5796 });
5797 }
5798 ApprovalResult::Rejected { reason } => {
5799 let reason = reason.unwrap_or_else(|| "rejected".to_string());
5800 let record = self.record_from_parts(
5801 &request,
5802 canonical_id,
5803 executed_arguments,
5804 started_at,
5805 start,
5806 false,
5807 false,
5808 format!("Approval rejected: {}", reason),
5809 metadata,
5810 ToolPolicyDecisionRecord::approval(reason.clone()),
5811 Some(ToolApprovalRecord {
5812 status: ToolApprovalStatus::Rejected,
5813 reason: Some(reason),
5814 modified_arguments: None,
5815 }),
5816 false,
5817 false,
5818 );
5819 self.finish_tool_record(&record).await;
5820 return Ok(record);
5821 }
5822 ApprovalResult::Timeout => {
5823 let record = self.record_from_parts(
5824 &request,
5825 canonical_id,
5826 executed_arguments,
5827 started_at,
5828 start,
5829 false,
5830 false,
5831 "Approval timed out".to_string(),
5832 metadata,
5833 ToolPolicyDecisionRecord::approval("approval timeout"),
5834 Some(ToolApprovalRecord {
5835 status: ToolApprovalStatus::Timeout,
5836 reason: Some("approval timeout".to_string()),
5837 modified_arguments: None,
5838 }),
5839 false,
5840 false,
5841 );
5842 self.finish_tool_record(&record).await;
5843 return Ok(record);
5844 }
5845 }
5846 }
5847 }
5848
5849 executed_arguments = security_engine.prepare_tool_arguments_with_bindings(
5854 &canonical_id,
5855 &executed_arguments,
5856 &bindings,
5857 );
5858 if let Some(record) = approval_record.as_mut()
5859 && matches!(record.status, ToolApprovalStatus::Modified)
5860 {
5861 record.modified_arguments = Some(executed_arguments.clone());
5862 }
5863 let binding_security_result = security_engine
5864 .validate_tool_execution_with_bindings(&canonical_id, &executed_arguments, &bindings)
5865 .await?;
5866 let approval_confirmation_required = matches!(
5867 binding_security_result,
5868 SecurityCheckResult::RequireConfirmation { .. }
5869 ) || security_engine
5870 .classification_approval_message(
5871 &canonical_id,
5872 &resolved.tool.classify_call(&executed_arguments),
5873 )
5874 .is_some();
5875 let approval_binding = approval_record.as_ref().and_then(|record| {
5876 matches!(
5877 record.status,
5878 ToolApprovalStatus::Approved | ToolApprovalStatus::Modified
5879 )
5880 .then(|| ToolApprovalBinding {
5881 canonical_id: canonical_id.clone(),
5882 arguments: executed_arguments.clone(),
5883 confirmation_required: approval_confirmation_required,
5884 policy_version: security_engine.policy_version(),
5885 runtime_control_version: approval_control_snapshot.version,
5886 state_generation: initial_scope_snapshot.state_generation,
5887 reviewed_tool: Arc::clone(&resolved.tool),
5888 })
5889 });
5890
5891 let control_snapshot = self.runtime_safety_snapshot();
5896 let resolved = self.tools.resolve(&request.requested_name);
5897 let registry_version = self.tools.version();
5898 let mut versions = ToolDecisionVersions {
5899 policy: control_snapshot.tool_security.policy_version(),
5900 registry: registry_version,
5901 runtime_control: control_snapshot.version,
5902 state: None,
5903 };
5904 metadata.insert(
5905 "runtime_scope_snapshot".to_string(),
5906 serde_json::to_value(&control_snapshot.tool_scope_override).unwrap_or(Value::Null),
5907 );
5908 let resolved = match resolved {
5909 Some(resolved) => resolved,
5910 None => {
5911 let reason = format!(
5912 "Tool '{}' became unavailable after approval",
5913 request.requested_name
5914 );
5915 let record = self.record_from_parts_at(
5916 &request,
5917 request.requested_name.clone(),
5918 executed_arguments,
5919 started_at,
5920 start,
5921 false,
5922 false,
5923 reason.clone(),
5924 metadata,
5925 ToolPolicyDecisionRecord::unavailable(reason),
5926 approval_record,
5927 false,
5928 false,
5929 versions,
5930 );
5931 self.finish_tool_record(&record).await;
5932 return Ok(record);
5933 }
5934 };
5935
5936 let canonical_id = resolved.identity.canonical_id.clone();
5937 if let Some(reason) =
5938 fallback_state.final_rejection_reason(&admitted_canonical_id, &canonical_id)
5939 {
5940 metadata.insert(
5944 "fallback_chain".to_string(),
5945 serde_json::to_value(&fallback_state.visited_canonical_ids).unwrap_or(Value::Null),
5946 );
5947 metadata.insert(
5948 "final_resolved_canonical_id".to_string(),
5949 Value::String(canonical_id),
5950 );
5951 let record = self.record_from_parts_at(
5952 &request,
5953 admitted_canonical_id,
5954 executed_arguments,
5955 started_at,
5956 start,
5957 false,
5958 false,
5959 format!("Denied: {reason}"),
5960 metadata,
5961 ToolPolicyDecisionRecord::deny(reason),
5962 approval_record,
5963 false,
5964 false,
5965 versions,
5966 );
5967 self.finish_tool_record(&record).await;
5968 return Ok(record);
5969 }
5970 let bindings = resolved.tool.policy_bindings();
5971 let final_arguments = control_snapshot
5972 .tool_security
5973 .prepare_tool_arguments_with_bindings(&canonical_id, &executed_arguments, &bindings);
5974 if let Some(record) = approval_record.as_mut()
5975 && matches!(record.status, ToolApprovalStatus::Modified)
5976 {
5977 record.modified_arguments = Some(final_arguments.clone());
5978 }
5979 let classification = resolved.tool.classify_call(&final_arguments);
5980 let safety = resolved.tool.safety_metadata();
5981 let security_engine = control_snapshot.tool_security;
5982 let tool_config = self.recovery_manager.get_tool_config(&canonical_id).clone();
5983 let recovery_timeout_ms = self.recovery_manager.get_tool_timeout(&canonical_id);
5984 metadata.insert(
5985 "classification".to_string(),
5986 serde_json::to_value(&classification).unwrap_or(Value::Null),
5987 );
5988 let (limits, timeout) = match Self::effective_tool_limits(
5992 &security_engine,
5993 &canonical_id,
5994 &safety,
5995 &classification,
5996 recovery_timeout_ms,
5997 ) {
5998 Ok(effective) => effective,
5999 Err(error) => {
6000 let reason = error.to_string();
6001 metadata.insert(
6002 "configuration_error".to_string(),
6003 Value::String(reason.clone()),
6004 );
6005 let record = self.record_from_parts_at(
6006 &request,
6007 canonical_id,
6008 final_arguments,
6009 started_at,
6010 start,
6011 false,
6012 false,
6013 format!("Denied: {reason}"),
6014 metadata,
6015 ToolPolicyDecisionRecord::deny(reason),
6016 approval_record,
6017 false,
6018 false,
6019 versions,
6020 );
6021 self.finish_tool_record(&record).await;
6022 return Ok(record);
6023 }
6024 };
6025 let policy_snapshot = security_engine.policy_snapshot(&canonical_id);
6026 let resource_lock_keys =
6027 tool_resource_lock_keys(&canonical_id, &final_arguments, &bindings, &classification);
6028 metadata.insert(
6029 "effective_limits".to_string(),
6030 serde_json::to_value(&limits).unwrap_or(Value::Null),
6031 );
6032 metadata.insert(
6033 "resource_lock_keys".to_string(),
6034 serde_json::to_value(&resource_lock_keys).unwrap_or(Value::Null),
6035 );
6036 if policy_snapshot.is_null() {
6037 metadata.remove("policy_snapshot");
6038 } else {
6039 metadata.insert("policy_snapshot".to_string(), policy_snapshot.clone());
6040 }
6041
6042 let final_denial = |canonical_id: String,
6043 output: String,
6044 policy: ToolPolicyDecisionRecord,
6045 metadata: HashMap<String, Value>,
6046 decision_versions: ToolDecisionVersions| {
6047 self.record_from_parts_at(
6048 &request,
6049 canonical_id,
6050 final_arguments.clone(),
6051 started_at,
6052 start,
6053 false,
6054 false,
6055 output,
6056 metadata,
6057 policy,
6058 approval_record.clone(),
6059 false,
6060 false,
6061 decision_versions,
6062 )
6063 };
6064
6065 if control_snapshot.emergency_deny {
6066 let reason = "Tool execution is disabled by runtime control".to_string();
6067 let record = final_denial(
6068 canonical_id,
6069 reason.clone(),
6070 ToolPolicyDecisionRecord::deny(reason),
6071 metadata,
6072 versions,
6073 );
6074 self.finish_tool_record(&record).await;
6075 return Ok(record);
6076 }
6077
6078 let available_snapshot = self
6083 .get_available_tool_ids_snapshot_for_scope(
6084 control_snapshot.tool_scope_override.as_deref(),
6085 )
6086 .await?;
6087 versions.state = available_snapshot.state_generation;
6088 metadata.insert(
6089 "available_tool_ids_snapshot".to_string(),
6090 serde_json::to_value(&available_snapshot.tool_ids).unwrap_or(Value::Null),
6091 );
6092 metadata.insert(
6093 "state_generation_snapshot".to_string(),
6094 serde_json::to_value(available_snapshot.state_generation).unwrap_or(Value::Null),
6095 );
6096 if !available_snapshot
6097 .tool_ids
6098 .iter()
6099 .any(|tool_id| tool_id == &canonical_id)
6100 {
6101 let reason = format!(
6102 "Tool '{}' is not available in the final runtime scope",
6103 canonical_id
6104 );
6105 let record = final_denial(
6106 canonical_id,
6107 reason.clone(),
6108 ToolPolicyDecisionRecord::deny(reason),
6109 metadata,
6110 versions,
6111 );
6112 self.finish_tool_record(&record).await;
6113 return Ok(record);
6114 }
6115
6116 let final_security_result = security_engine
6121 .validate_tool_execution_with_bindings(&canonical_id, &final_arguments, &bindings)
6122 .await?;
6123 match &final_security_result {
6124 SecurityCheckResult::Block { reason } => {
6125 let record = final_denial(
6126 canonical_id,
6127 format!("Denied: {}", reason),
6128 ToolPolicyDecisionRecord::deny(reason.clone()),
6129 metadata,
6130 versions,
6131 );
6132 self.finish_tool_record(&record).await;
6133 return Ok(record);
6134 }
6135 SecurityCheckResult::Unavailable { reason } => {
6136 let record = final_denial(
6137 canonical_id,
6138 format!("Unavailable: {}", reason),
6139 ToolPolicyDecisionRecord::unavailable(reason.clone()),
6140 metadata,
6141 versions,
6142 );
6143 self.finish_tool_record(&record).await;
6144 return Ok(record);
6145 }
6146 SecurityCheckResult::Warn { message } => {
6147 warn!(tool = %canonical_id, message = %message, "Tool security warning after approval");
6148 }
6149 SecurityCheckResult::Allow | SecurityCheckResult::RequireConfirmation { .. } => {}
6150 }
6151 let final_confirmation_required = matches!(
6152 final_security_result,
6153 SecurityCheckResult::RequireConfirmation { .. }
6154 ) || security_engine
6155 .classification_approval_message(&canonical_id, &classification)
6156 .is_some();
6157 let stale_approval = approval_binding.as_ref().is_some_and(|binding| {
6158 binding.is_stale(
6159 &canonical_id,
6160 &final_arguments,
6161 final_confirmation_required,
6162 versions,
6163 &resolved.tool,
6164 )
6165 });
6166 if stale_approval {
6167 let reason = "Approval became stale before final admission".to_string();
6168 let record = final_denial(
6169 canonical_id,
6170 reason.clone(),
6171 ToolPolicyDecisionRecord::deny(reason),
6172 metadata,
6173 versions,
6174 );
6175 self.finish_tool_record(&record).await;
6176 return Ok(record);
6177 }
6178 if final_confirmation_required && approval_binding.is_none() {
6179 let reason = "Final policy requires fresh approval".to_string();
6180 let record = final_denial(
6181 canonical_id,
6182 reason.clone(),
6183 ToolPolicyDecisionRecord::approval(reason),
6184 metadata,
6185 versions,
6186 );
6187 self.finish_tool_record(&record).await;
6188 return Ok(record);
6189 }
6190
6191 if let Some((_, reason)) = self.host_tool_unavailability(&canonical_id) {
6192 let record = final_denial(
6193 canonical_id,
6194 reason.to_string(),
6195 ToolPolicyDecisionRecord::unavailable(reason),
6196 metadata,
6197 versions,
6198 );
6199 self.finish_tool_record(&record).await;
6200 return Ok(record);
6201 }
6202
6203 let Some(resource_guards) = self.acquire_tool_resource_locks(&resource_lock_keys).await
6208 else {
6209 let reason = "Tool execution cancelled while waiting for resource locks".to_string();
6213 let mut record = final_denial(
6214 canonical_id,
6215 reason.clone(),
6216 ToolPolicyDecisionRecord::deny(reason),
6217 metadata,
6218 versions,
6219 );
6220 record.cancelled = true;
6221 record.cancellation_reason = Some("runtime control cancellation".to_string());
6222 self.finish_tool_record(&record).await;
6223 return Ok(record);
6224 };
6225
6226 let admission = self.admit_tool_execution(
6231 versions.runtime_control,
6232 versions.policy,
6233 versions.state,
6234 &canonical_id,
6235 );
6236 if !matches!(admission, SecurityCheckResult::Allow) {
6237 let latest_control = self.runtime_safety_snapshot();
6238 let reason = admission
6239 .reason()
6240 .unwrap_or("tool admission was denied")
6241 .to_string();
6242 let policy = if admission.is_unavailable() {
6243 ToolPolicyDecisionRecord::unavailable(reason.clone())
6244 } else {
6245 ToolPolicyDecisionRecord::deny(reason.clone())
6246 };
6247 let record = self.record_from_parts_at(
6248 &request,
6249 canonical_id,
6250 final_arguments,
6251 started_at,
6252 start,
6253 false,
6254 false,
6255 reason,
6256 metadata,
6257 policy,
6258 approval_record,
6259 false,
6260 false,
6261 ToolDecisionVersions {
6262 policy: latest_control.tool_security.policy_version(),
6263 registry: versions.registry,
6264 runtime_control: latest_control.version,
6265 state: self
6266 .state_machine
6267 .as_ref()
6268 .map(|state_machine| state_machine.generation()),
6269 },
6270 );
6271 self.finish_tool_record_after_resource_guards(resource_guards, &record)
6272 .await;
6273 return Ok(record);
6274 }
6275 let executed_arguments = final_arguments;
6276
6277 let turn_actor = current_turn_actor_context();
6278 let actor = ToolActorContext {
6279 actor_id: turn_actor
6280 .as_ref()
6281 .and_then(|context| context.effective_actor_id().map(str::to_string))
6282 .or_else(|| self.actor_id()),
6283 origin_actor_id: turn_actor
6284 .as_ref()
6285 .and_then(|context| context.origin_actor_id.clone()),
6286 sender_agent_id: turn_actor
6287 .as_ref()
6288 .and_then(|context| context.sender_agent_id.clone()),
6289 };
6290 let tool_context = ToolExecutionContext {
6291 requested_name: request.requested_name.clone(),
6292 canonical_id: canonical_id.clone(),
6293 display_name: resolved.identity.display_name.clone(),
6294 provider_id: resolved.identity.provider_id.clone(),
6295 registry_version: versions.registry,
6296 policy_version: versions.policy,
6297 runtime_control_version: versions.runtime_control,
6298 call_id: request.call_id.clone(),
6299 source: request.source.clone(),
6300 actor,
6301 cancellation: ToolCancellationToken::new(
6302 Arc::clone(&self.runtime_control.emergency_deny),
6303 Some("runtime control cancellation".to_string()),
6304 ),
6305 started_at,
6306 deadline: None,
6307 permission: ToolPolicyDecisionRecord::allow(),
6308 approval: approval_record.clone(),
6309 classification: classification.clone(),
6310 safety,
6311 limits: limits.clone(),
6312 policy_snapshot,
6313 custom_config: security_engine.custom_config(&canonical_id),
6314 };
6315 let (mut result, timed_out, cancelled, invoked) = self
6316 .run_tool_with_retries(
6317 &canonical_id,
6318 resolved.tool.clone(),
6319 executed_arguments.clone(),
6320 tool_context,
6321 timeout,
6322 tool_config.max_retries,
6323 )
6324 .await?;
6325
6326 let fallback_tool = if !result.success && !cancelled {
6330 match &tool_config.on_failure {
6331 ToolFailureAction::Skip => {
6332 result = ToolResult::ok(format!(
6333 "{{\"skipped\": true, \"reason\": \"Tool '{}' was skipped after failure\"}}",
6334 canonical_id
6335 ));
6336 None
6337 }
6338 ToolFailureAction::Fallback { fallback_tool } => Some(fallback_tool.clone()),
6339 ToolFailureAction::ReportError => None,
6340 }
6341 } else {
6342 None
6343 };
6344
6345 let output_cap = limits.max_output_chars;
6346 let (output, output_truncated) =
6347 Self::truncate_tool_output(result.output.clone(), output_cap);
6348 if let Some(result_metadata) = result.metadata {
6349 metadata.extend(result_metadata);
6350 }
6351 let mut record = self.record_from_parts_at(
6352 &request,
6353 canonical_id,
6354 executed_arguments,
6355 started_at,
6356 start,
6357 invoked,
6358 result.success,
6359 output,
6360 metadata,
6361 ToolPolicyDecisionRecord::allow(),
6362 approval_record,
6363 timed_out,
6364 output_truncated,
6365 versions,
6366 );
6367 record.cancelled = cancelled;
6368 if cancelled {
6369 record.cancellation_reason = Some("runtime control cancellation".to_string());
6370 }
6371 if let Some(fallback_tool) = fallback_tool {
6372 let fallback_arguments = record.executed_arguments.clone();
6373 let original_tool = record.canonical_id.clone();
6374 self.finish_tool_record_after_resource_guards(resource_guards, &record)
6378 .await;
6379 let fallback_request = ToolExecutionRequest::new(
6380 request.call_id.clone(),
6381 fallback_tool,
6382 fallback_arguments,
6383 ToolCallSource::Fallback { original_tool },
6384 );
6385 return Box::pin(self.execute_tool_record_inner(fallback_request, fallback_state))
6386 .await;
6387 }
6388 self.finish_tool_record_after_resource_guards(resource_guards, &record)
6389 .await;
6390 Ok(record)
6391 }
6392
6393 #[instrument(skip(self, tool_call), fields(tool = %tool_call.name))]
6394 async fn execute_tool_smart(&self, tool_call: &ToolCall) -> Result<String> {
6395 let record = self
6396 .execute_tool_record(ToolExecutionRequest::new(
6397 tool_call.id.clone(),
6398 tool_call.name.clone(),
6399 tool_call.arguments.clone(),
6400 ToolCallSource::Model,
6401 ))
6402 .await?;
6403 if record.success {
6404 Ok(record.model_output_string())
6405 } else if matches!(record.policy.outcome, PermissionOutcome::RequiresApproval) {
6406 Err(AgentError::HITLRejected(record.model_output_string()))
6407 } else {
6408 Err(AgentError::Tool(record.model_output_string()))
6409 }
6410 }
6411
6412 async fn select_skill_candidate(&self, input: &str) -> Result<Option<SkillCandidate>> {
6418 let Some(ref router) = self.skill_router else {
6419 return Ok(None);
6420 };
6421 let available_skills = self.get_available_skills();
6422 if available_skills.is_empty() {
6423 return Ok(None);
6424 }
6425 let skill_ids: Vec<&str> = available_skills.iter().map(|s| s.id.as_str()).collect();
6426 let Some(skill_id) = self
6427 .observe_purpose(
6428 ObservationPurpose::SkillRouting,
6429 router.select_skill_filtered(input, &skill_ids),
6430 )
6431 .await?
6432 else {
6433 return Ok(None);
6434 };
6435 let skill = router
6436 .get_skill(&skill_id)
6437 .cloned()
6438 .ok_or_else(|| AgentError::Skill(format!("Skill not found: {}", skill_id)))?;
6439 info!(skill_id = %skill_id, "Skill selected");
6440 Ok(Some(SkillCandidate::new(skill_id, skill)))
6441 }
6442
6443 async fn commit_skill_candidate_route_result(
6448 &self,
6449 candidate: SkillCandidate,
6450 input: &str,
6451 ) -> Result<SkillRouteResult> {
6452 let skill_id = candidate.skill_id;
6453 let skill = candidate.skill;
6454 let expected_state_generation = self
6455 .state_machine
6456 .as_ref()
6457 .map(|state_machine| state_machine.generation());
6458 let expected_disambiguation_epoch = self.disambiguation_epoch.load(Ordering::SeqCst);
6459 if let Some(ref skill_disambig) = skill.disambiguation
6460 && skill_disambig.enabled.unwrap_or(false)
6461 && let Some(ref disambiguator) = self.disambiguation_manager
6462 {
6463 let context = self.build_disambiguation_context().await?;
6464 let state_override = self
6465 .state_machine
6466 .as_ref()
6467 .and_then(|sm| sm.current_definition())
6468 .and_then(|def| def.disambiguation.clone());
6469
6470 let disambiguation_result = self
6471 .observe_purpose(
6472 ObservationPurpose::DisambiguationDetection,
6473 disambiguator.process_input_with_override(
6474 input,
6475 &context,
6476 state_override.as_ref(),
6477 Some(skill_disambig),
6478 ),
6479 )
6480 .await?;
6481 let current_state_generation = self
6482 .state_machine
6483 .as_ref()
6484 .map(|state_machine| state_machine.generation());
6485 if current_state_generation != expected_state_generation
6486 || self.disambiguation_epoch.load(Ordering::SeqCst) != expected_disambiguation_epoch
6487 {
6488 disambiguator.clear_pending().await;
6489 *self.pending_skill_id.write() = None;
6490 return Err(AgentError::Other(
6491 "State or reset ownership changed during skill disambiguation".to_string(),
6492 ));
6493 }
6494 match disambiguation_result {
6495 DisambiguationResult::Clear => {
6496 debug!(skill_id = %skill_id, "Skill disambiguation: clear");
6497 }
6498 DisambiguationResult::NeedsClarification {
6499 question,
6500 detection,
6501 } => {
6502 let admission = self
6503 .admit_disambiguation_redispatch(
6504 expected_disambiguation_epoch,
6505 expected_state_generation,
6506 )
6507 .await?;
6508 let awaiting_confirmation = disambiguator.has_pending_confirmation().await;
6509 info!(
6510 skill_id = %skill_id,
6511 ambiguity_type = ?detection.ambiguity_type,
6512 confidence = detection.confidence,
6513 "Skill requires clarification before execution"
6514 );
6515 *self.pending_skill_id.write() = Some(skill_id.clone());
6516 let response = AgentResponse::new(&question.question).with_metadata(
6517 "disambiguation",
6518 serde_json::json!({
6519 "status": if awaiting_confirmation { "awaiting_confirmation" } else { "awaiting_clarification" },
6520 "skill_id": skill_id,
6521 "options": question.options,
6522 "clarifying": question.clarifying,
6523 "detection": {
6524 "type": detection.ambiguity_type,
6525 "confidence": detection.confidence,
6526 "what_is_unclear": detection.what_is_unclear,
6527 }
6528 }),
6529 );
6530 drop(admission);
6531 return Ok(SkillRouteResult::NeedsClarification {
6532 response,
6533 ownership: Some(DisambiguationOwnership {
6534 epoch: expected_disambiguation_epoch,
6535 state_generation: expected_state_generation,
6536 }),
6537 });
6538 }
6539 DisambiguationResult::Clarified { enriched_input, .. } => {
6540 info!(skill_id = %skill_id, enriched = %enriched_input, "Skill disambiguation clarified");
6541 let admission = self
6542 .admit_disambiguation_redispatch(
6543 expected_disambiguation_epoch,
6544 expected_state_generation,
6545 )
6546 .await?;
6547 drop(admission);
6548 let content = self.execute_skill(&skill, &enriched_input).await?;
6549 return Ok(SkillRouteResult::Response { skill_id, content });
6550 }
6551 DisambiguationResult::ProceedWithBestGuess { enriched_input } => {
6552 info!(skill_id = %skill_id, "Skill disambiguation best guess");
6553 let admission = self
6554 .admit_disambiguation_redispatch(
6555 expected_disambiguation_epoch,
6556 expected_state_generation,
6557 )
6558 .await?;
6559 drop(admission);
6560 let content = self.execute_skill(&skill, &enriched_input).await?;
6561 return Ok(SkillRouteResult::Response { skill_id, content });
6562 }
6563 DisambiguationResult::GiveUp { reason } => {
6564 warn!(skill_id = %skill_id, reason = %reason, "Skill disambiguation gave up");
6565 let apology = self
6566 .generate_localized_apology(
6567 "Generate a brief, polite apology saying you couldn't understand the request. Be concise.",
6568 &reason,
6569 )
6570 .await
6571 .unwrap_or_else(|_| {
6572 format!("I'm sorry, I couldn't understand your request: {}", reason)
6573 });
6574 return Ok(SkillRouteResult::NeedsClarification {
6575 response: AgentResponse::new(&apology),
6576 ownership: None,
6577 });
6578 }
6579 DisambiguationResult::Escalate { reason } => {
6580 info!(skill_id = %skill_id, reason = %reason, "Skill disambiguation escalating");
6581 let apology = self
6582 .generate_localized_apology(
6583 "Explain briefly that you're transferring the user to a human agent for help.",
6584 &reason,
6585 )
6586 .await
6587 .unwrap_or_else(|_| {
6588 format!("I need human assistance to help with your request: {}", reason)
6589 });
6590 return Ok(SkillRouteResult::NeedsClarification {
6591 response: AgentResponse::new(&apology),
6592 ownership: None,
6593 });
6594 }
6595 DisambiguationResult::Abandoned { .. } => {
6596 debug!(skill_id = %skill_id, "Skill disambiguation abandoned");
6597 return Ok(SkillRouteResult::NoMatch);
6598 }
6599 }
6600 }
6601 let admission = self
6602 .admit_disambiguation_redispatch(
6603 expected_disambiguation_epoch,
6604 expected_state_generation,
6605 )
6606 .await?;
6607 drop(admission);
6608 let content = self.execute_skill(&skill, input).await?;
6609 Ok(SkillRouteResult::Response { skill_id, content })
6610 }
6611
6612 async fn try_skill_route(&self, input: &str) -> Result<SkillRouteResult> {
6614 if let Some(candidate) = self.select_skill_candidate(input).await? {
6615 self.commit_skill_candidate_route_result(candidate, input)
6616 .await
6617 } else {
6618 Ok(SkillRouteResult::NoMatch)
6619 }
6620 }
6621
6622 fn skill_clarification_needs_memory_record(response: &AgentResponse) -> bool {
6625 response
6626 .metadata
6627 .as_ref()
6628 .and_then(|m| m.get("disambiguation"))
6629 .and_then(|d| d.get("status"))
6630 .and_then(|s| s.as_str())
6631 == Some("awaiting_clarification")
6632 }
6633
6634 async fn commit_winning_skill_candidate(
6641 &self,
6642 candidate: SkillCandidate,
6643 processed_input: &str,
6644 input_context: &HashMap<String, Value>,
6645 ) -> Result<Option<AgentResponse>> {
6646 self.commit_root_user_message(processed_input).await?;
6647 match self
6648 .commit_skill_candidate_route_result(candidate, processed_input)
6649 .await?
6650 {
6651 SkillRouteResult::Response { skill_id, content } => self
6652 .handle_skill_response(processed_input, &skill_id, content, input_context)
6653 .await
6654 .map(Some),
6655 SkillRouteResult::NeedsClarification {
6656 response,
6657 ownership,
6658 } => {
6659 let admission = self
6660 .admit_optional_disambiguation_ownership(ownership)
6661 .await?;
6662 if Self::skill_clarification_needs_memory_record(&response) {
6663 self.memory
6664 .add_message(ChatMessage::assistant(&response.content))
6665 .await?;
6666 }
6667 drop(admission);
6668 self.finish_turn_if_root(&response).await?;
6669 Ok(Some(response))
6670 }
6671 SkillRouteResult::NoMatch => Ok(None),
6672 }
6673 }
6674
6675 async fn execute_skill(&self, skill: &SkillDefinition, input: &str) -> Result<String> {
6677 if let Some(ref executor) = self.skill_executor {
6678 let skill_reasoning = self.get_skill_reasoning_config(skill);
6679 let skill_reflection = self.get_skill_reflection_config(skill);
6680
6681 debug!(
6682 skill_id = %skill.id,
6683 reasoning_mode = ?skill_reasoning.mode,
6684 reflection_enabled = ?skill_reflection.enabled,
6685 "Skill reasoning/reflection config"
6686 );
6687
6688 let response = self
6689 .observe_purpose(
6690 ObservationPurpose::SkillPrompt,
6691 executor.execute_with_invoker(skill, input, serde_json::json!({}), self),
6692 )
6693 .await?;
6694
6695 if skill_reflection.requires_evaluation() && skill_reflection.is_enabled() {
6696 let should_reflect = self
6697 .should_reflect_with_config(input, &response, &skill_reflection)
6698 .await?;
6699 if should_reflect {
6700 let evaluated = self
6701 .evaluate_and_retry_with_config(input, response, &skill_reflection)
6702 .await?;
6703 return Ok(evaluated);
6704 }
6705 }
6706
6707 return Ok(response);
6708 }
6709 Err(AgentError::Skill(
6710 "No skill executor configured".to_string(),
6711 ))
6712 }
6713
6714 async fn execute_skill_by_id(&self, skill_id: &str, input: &str) -> Result<String> {
6717 let skill = self
6718 .skill_router
6719 .as_ref()
6720 .and_then(|r| r.get_skill(skill_id).cloned())
6721 .ok_or_else(|| AgentError::Skill(format!("Skill not found: {}", skill_id)))?;
6722 self.execute_skill(&skill, input).await
6723 }
6724
6725 async fn should_reflect_with_config(
6726 &self,
6727 input: &str,
6728 response: &str,
6729 config: &ReflectionConfig,
6730 ) -> Result<bool> {
6731 if !config.requires_evaluation() {
6732 return Ok(false);
6733 }
6734
6735 if config.is_enabled() {
6736 return Ok(true);
6737 }
6738
6739 let evaluator_llm = config
6740 .evaluator_llm
6741 .as_ref()
6742 .and_then(|alias| self.llm_registry.get(alias).ok())
6743 .or_else(|| self.llm_registry.router().ok())
6744 .or_else(|| self.llm_registry.default().ok());
6745
6746 let Some(llm) = evaluator_llm else {
6747 return Ok(false);
6748 };
6749
6750 let response_preview: String = response.chars().take(500).collect();
6751 let prompt = format!(
6752 r#"Should this response be evaluated for quality? Consider if it's a complex or important response.
6753
6754User query: "{}"
6755Response: "{}"
6756
6757Answer YES or NO only."#,
6758 input, response_preview
6759 );
6760
6761 let messages = vec![ChatMessage::user(&prompt)];
6762 let result = self
6763 .observe_purpose(
6764 ObservationPurpose::ReflectionDecision,
6765 llm.complete(&messages, None),
6766 )
6767 .await;
6768
6769 match result {
6770 Ok(resp) => Ok(resp.content.trim().to_uppercase().contains("YES")),
6771 Err(_) => Ok(false),
6772 }
6773 }
6774
6775 async fn evaluate_and_retry_with_config(
6776 &self,
6777 input: &str,
6778 mut response: String,
6779 config: &ReflectionConfig,
6780 ) -> Result<String> {
6781 let llm = self.get_state_llm()?;
6782 let mut attempts = 0u32;
6783 let max_retries = config.max_retries;
6784
6785 loop {
6786 let evaluation = self
6787 .evaluate_response_with_config(input, &response, config)
6788 .await?;
6789
6790 if evaluation.passed || attempts >= max_retries {
6791 info!(
6792 passed = evaluation.passed,
6793 confidence = evaluation.confidence,
6794 attempts = attempts + 1,
6795 "Skill reflection evaluation complete"
6796 );
6797 return Ok(response);
6798 }
6799
6800 debug!(
6801 attempt = attempts + 1,
6802 failed_criteria = evaluation.failed_criteria().count(),
6803 "Skill response did not meet criteria, retrying"
6804 );
6805
6806 let feedback: Vec<String> = evaluation
6807 .failed_criteria()
6808 .map(|c| format!("- {}", c.criterion))
6809 .collect();
6810
6811 let retry_prompt = format!(
6812 "Your previous response did not meet these criteria:\n{}\n\nPlease provide an improved response to: {}",
6813 feedback.join("\n"),
6814 input
6815 );
6816
6817 let messages = vec![ChatMessage::user(&retry_prompt)];
6818 let retry_response = self
6819 .observe_purpose(
6820 ObservationPurpose::ReflectionEvaluation,
6821 llm.complete(&messages, None),
6822 )
6823 .await
6824 .map_err(|e| AgentError::LLM(e.to_string()))?;
6825
6826 response = retry_response.content.trim().to_string();
6827 attempts += 1;
6828 }
6829 }
6830
6831 async fn evaluate_response_with_config(
6832 &self,
6833 input: &str,
6834 response: &str,
6835 config: &ReflectionConfig,
6836 ) -> Result<EvaluationResult> {
6837 let evaluator_llm = config
6838 .evaluator_llm
6839 .as_ref()
6840 .and_then(|alias| self.llm_registry.get(alias).ok())
6841 .or_else(|| self.llm_registry.router().ok())
6842 .or_else(|| self.llm_registry.default().ok())
6843 .ok_or_else(|| AgentError::Config("No LLM available for evaluation".into()))?;
6844
6845 let criteria = &config.criteria;
6846 let criteria_list = criteria
6847 .iter()
6848 .enumerate()
6849 .map(|(i, c)| format!("{}. {}", i + 1, c))
6850 .collect::<Vec<_>>()
6851 .join("\n");
6852
6853 let prompt = format!(
6854 r#"Evaluate this response against the criteria.
6855
6856User query: "{}"
6857
6858Response to evaluate: "{}"
6859
6860Criteria:
6861{}
6862
6863For each criterion, respond with:
6864- criterion number
6865- PASS or FAIL
6866- brief reason
6867
6868Then provide overall confidence (0.0 to 1.0) and whether it passes overall.
6869
6870Format:
68711. PASS/FAIL - reason
68722. PASS/FAIL - reason
6873...
6874CONFIDENCE: 0.X
6875OVERALL: PASS/FAIL"#,
6876 input, response, criteria_list
6877 );
6878
6879 let messages = vec![ChatMessage::user(&prompt)];
6880 let eval_response = self
6881 .observe_purpose(
6882 ObservationPurpose::ReflectionEvaluation,
6883 evaluator_llm.complete(&messages, None),
6884 )
6885 .await
6886 .map_err(|e| AgentError::LLM(format!("Evaluation failed: {}", e)))?;
6887
6888 let content = eval_response.content.to_uppercase();
6889 let llm_pass = content.contains("OVERALL: PASS");
6890
6891 let confidence = content
6892 .lines()
6893 .find(|l| l.contains("CONFIDENCE:"))
6894 .and_then(|l| {
6895 l.split(':')
6896 .nth(1)
6897 .and_then(|v| v.trim().parse::<f32>().ok())
6898 })
6899 .unwrap_or(if llm_pass { 0.8 } else { 0.4 });
6900
6901 let overall_pass = llm_pass && confidence >= config.pass_threshold;
6904
6905 let mut criteria_results = Vec::new();
6906 for (i, criterion) in criteria.iter().enumerate() {
6907 let line_marker = format!("{}.", i + 1);
6908 let passed = eval_response
6909 .content
6910 .lines()
6911 .find(|l| l.contains(&line_marker))
6912 .map(|l| l.to_uppercase().contains("PASS"))
6913 .unwrap_or(overall_pass);
6914
6915 if passed {
6916 criteria_results.push(CriterionResult::pass(criterion));
6917 } else {
6918 criteria_results.push(CriterionResult::fail(criterion, "Did not meet criterion"));
6919 }
6920 }
6921
6922 Ok(EvaluationResult::new(overall_pass, confidence).with_criteria(criteria_results))
6923 }
6924
6925 async fn process_input(&self, input: &str) -> Result<ProcessData> {
6927 if let Some(processor) = self.get_state_process_processor() {
6928 let purpose = observation_purpose_for_process(processor.input_purpose_hint());
6929 return self
6930 .observe_purpose(purpose, processor.process_input(input))
6931 .await;
6932 }
6933 if let Some(ref processor) = self.process_processor {
6934 let purpose = observation_purpose_for_process(processor.input_purpose_hint());
6935 self.observe_purpose(purpose, processor.process_input(input))
6936 .await
6937 } else {
6938 Ok(ProcessData::new(input))
6939 }
6940 }
6941
6942 async fn process_output(
6944 &self,
6945 output: &str,
6946 input_context: &std::collections::HashMap<String, serde_json::Value>,
6947 ) -> Result<ProcessData> {
6948 if let Some(processor) = self.get_state_process_processor() {
6949 let purpose = observation_purpose_for_process(processor.output_purpose_hint());
6950 return self
6951 .observe_purpose(purpose, processor.process_output(output, input_context))
6952 .await;
6953 }
6954 if let Some(ref processor) = self.process_processor {
6955 let purpose = observation_purpose_for_process(processor.output_purpose_hint());
6956 self.observe_purpose(purpose, processor.process_output(output, input_context))
6957 .await
6958 } else {
6959 Ok(ProcessData::new(output))
6960 }
6961 }
6962
6963 fn get_state_process_processor(&self) -> Option<ProcessProcessor> {
6965 let sm = self.state_machine.as_ref()?;
6966 let def = sm.current_definition()?;
6967 let config = def.process.as_ref()?;
6968 let mut processor = ProcessProcessor::new(config.clone());
6969 if let Some(ref registry) = Some(self.llm_registry.clone()) {
6970 processor = processor.with_llm_registry(registry.clone());
6971 }
6972 processor = processor.with_stage_observer(Arc::new(ObservabilityProcessStageObserver));
6973 Some(processor)
6974 }
6975
6976 async fn check_turn_timeout(&self) -> Result<()> {
6978 let Some(ref sm) = self.state_machine else {
6979 return Ok(());
6980 };
6981 let Some(timeout_state) = sm.check_timeout() else {
6982 return Ok(());
6983 };
6984 let claim_admission = self.disambiguation_admission.write().await;
6985 if sm.check_timeout().as_deref() != Some(timeout_state.as_str()) {
6986 return Ok(());
6987 }
6988 let Some(reservation) = self.reserve_state_transition() else {
6989 return Ok(());
6990 };
6991 let from_state = sm.current();
6992 let expected_state_generation = sm.generation();
6993 let expected_disambiguation_epoch = self.disambiguation_epoch.load(Ordering::SeqCst);
6994 let history_before = sm.history();
6995 drop(claim_admission);
6996
6997 self.execute_state_exit_actions(&from_state).await;
6998
6999 let admission = self.disambiguation_admission.write().await;
7000 if sm.current() != from_state
7001 || sm.generation() != expected_state_generation
7002 || self.disambiguation_epoch.load(Ordering::SeqCst) != expected_disambiguation_epoch
7003 || sm.check_timeout().as_deref() != Some(timeout_state.as_str())
7004 {
7005 return Ok(());
7006 }
7007 sm.transition_to(&timeout_state, "max_turns exceeded")?;
7008 self.invalidate_pending_confirmation("state_timeout").await;
7009 let entered = sm.current();
7010 let is_reentry = Self::state_was_previously_entered(&entered, &from_state, &history_before);
7011 drop(admission);
7012
7013 self.execute_state_enter_actions(&entered, is_reentry).await;
7014 drop(reservation);
7015 info!(to = %entered, "Timeout transition");
7016 Ok(())
7017 }
7018
7019 fn increment_turn(&self) {
7020 if let Some(ref sm) = self.state_machine {
7021 sm.increment_turn();
7022 }
7023 }
7024
7025 fn transitions_available_for_commit(&self) -> Option<(Vec<Transition>, String)> {
7026 let sm = self.state_machine.as_ref()?;
7027 let current = sm.current();
7028 let transitions: Vec<_> = sm
7029 .auto_transitions()
7030 .into_iter()
7031 .filter(|t| match t.cooldown_turns {
7032 Some(cd) if cd > 0 => {
7033 let resolved = sm.config().resolve_full_path(¤t, &t.to);
7034 !sm.is_on_cooldown(&resolved, cd)
7035 }
7036 _ => true,
7037 })
7038 .collect();
7039 Some((transitions, current))
7040 }
7041
7042 fn transition_reason(transition: &Transition) -> String {
7043 if transition.when.is_empty() {
7044 "guard condition met".to_string()
7045 } else {
7046 transition.when.clone()
7047 }
7048 }
7049
7050 fn build_transition_context(
7052 &self,
7053 user_message: &str,
7054 response: &str,
7055 current_state: &str,
7056 staged: Option<&HashMap<String, Value>>,
7057 ) -> TransitionContext {
7058 let context_map = staged
7059 .map(|writes| self.build_context_with_staged(writes))
7060 .unwrap_or_else(|| self.build_context_with_overlays());
7061 TransitionContext::new(user_message, response, current_state).with_context(context_map)
7062 }
7063
7064 async fn select_transition_candidate(
7066 &self,
7067 user_message: &str,
7068 response: &str,
7069 ) -> Result<Option<TransitionCandidate>> {
7070 let Some((transitions, current_state)) = self.transitions_available_for_commit() else {
7071 return Ok(None);
7072 };
7073 let transitions: Vec<Transition> = transitions
7074 .into_iter()
7075 .filter(|transition| matches!(transition.timing, TransitionTiming::PostResponse))
7076 .collect();
7077 if transitions.is_empty() {
7078 return Ok(None);
7079 }
7080 let Some(evaluator) = self.transition_evaluator.as_ref() else {
7081 return Ok(None);
7082 };
7083 let context = self.build_transition_context(user_message, response, ¤t_state, None);
7084 let selected = self
7085 .observe_purpose(
7086 ObservationPurpose::StateTransitionEvaluation,
7087 evaluator.select_transition(&transitions, &context),
7088 )
7089 .await?;
7090 Ok(selected.map(|index| {
7091 let transition = transitions[index].clone();
7092 TransitionCandidate::new(
7093 current_state,
7094 transition.clone(),
7095 Self::transition_reason(&transition),
7096 )
7097 }))
7098 }
7099
7100 fn select_deterministic_transition_candidate(
7102 &self,
7103 user_message: &str,
7104 current_state: &str,
7105 transitions: &[Transition],
7106 staged: &HashMap<String, Value>,
7107 ) -> Option<TransitionCandidate> {
7108 let context = self.build_transition_context(user_message, "", current_state, Some(staged));
7109
7110 for transition in transitions {
7111 if let Some(guard) = transition.guard.as_ref()
7112 && evaluate_guard(guard, &context)
7113 {
7114 return Some(TransitionCandidate::new(
7115 current_state,
7116 transition.clone(),
7117 Self::transition_reason(transition),
7118 ));
7119 }
7120 }
7121
7122 let resolved_intent = context
7123 .context
7124 .get("resolved_intent")
7125 .and_then(Value::as_str)
7126 .filter(|value| !value.is_empty());
7127 if let Some(resolved_intent) = resolved_intent {
7128 for transition in transitions {
7129 if transition.intent.as_deref() == Some(resolved_intent) {
7130 return Some(TransitionCandidate::new(
7131 current_state,
7132 transition.clone(),
7133 Self::transition_reason(transition),
7134 ));
7135 }
7136 }
7137 }
7138
7139 None
7140 }
7141
7142 async fn commit_transition_candidate(&self, candidate: &TransitionCandidate) -> Result<bool> {
7144 self.commit_transition_target(&candidate.from_state, candidate.target(), &candidate.reason)
7145 .await
7146 }
7147
7148 async fn approve_transition_target(&self, from_state: &str, target: &str) -> Result<bool> {
7150 let approved = self.check_state_hitl(Some(from_state), target).await?;
7151 if !approved {
7152 info!(to = %target, "State transition rejected by HITL");
7153 }
7154 Ok(approved)
7155 }
7156
7157 async fn apply_transition_target(
7159 &self,
7160 from_state: &str,
7161 target: &str,
7162 reason: &str,
7163 staged: Option<&HashMap<String, Value>>,
7164 ) -> Result<bool> {
7165 let Some(ref sm) = self.state_machine else {
7166 return Ok(false);
7167 };
7168 let claim_admission = self.disambiguation_admission.write().await;
7169 if sm.current() != from_state {
7170 return Ok(false);
7171 }
7172 let Some(reservation) = self.reserve_state_transition() else {
7173 return Ok(false);
7174 };
7175 let expected_state_generation = sm.generation();
7176 let expected_disambiguation_epoch = self.disambiguation_epoch.load(Ordering::SeqCst);
7177 let history_before = sm.history();
7178 drop(claim_admission);
7179
7180 self.execute_state_exit_actions(from_state).await;
7181
7182 let admission = self.disambiguation_admission.write().await;
7183 if sm.current() != from_state
7184 || sm.generation() != expected_state_generation
7185 || self.disambiguation_epoch.load(Ordering::SeqCst) != expected_disambiguation_epoch
7186 {
7187 return Ok(false);
7188 }
7189 sm.transition_to(target, reason)?;
7190 self.invalidate_pending_confirmation("state_transition")
7191 .await;
7192 sm.reset_no_transition();
7193 if let Some(staged) = staged {
7194 self.commit_staged_context_writes(staged);
7195 }
7196 let entered = sm.current();
7197 let is_reentry = Self::state_was_previously_entered(&entered, from_state, &history_before);
7198 drop(admission);
7199
7200 self.execute_state_enter_actions(&entered, is_reentry).await;
7201 drop(reservation);
7202 self.hooks
7203 .on_state_transition(Some(from_state), &entered, reason)
7204 .await;
7205 info!(from = %from_state, to = %entered, "State transition");
7206 Ok(true)
7207 }
7208
7209 async fn commit_transition_target(
7211 &self,
7212 from_state: &str,
7213 target: &str,
7214 reason: &str,
7215 ) -> Result<bool> {
7216 if !self.approve_transition_target(from_state, target).await? {
7217 return Ok(false);
7218 }
7219 self.apply_transition_target(from_state, target, reason, None)
7220 .await
7221 }
7222
7223 async fn apply_pre_response_transition_candidate(
7225 &self,
7226 candidate: &TransitionCandidate,
7227 staged: &HashMap<String, Value>,
7228 processed_input: &str,
7229 ) -> Result<bool> {
7230 self.commit_root_user_message(processed_input).await?;
7231 self.apply_transition_target(
7232 &candidate.from_state,
7233 candidate.target(),
7234 &candidate.reason,
7235 Some(staged),
7236 )
7237 .await
7238 }
7239
7240 async fn commit_pre_response_transition_candidate(
7242 &self,
7243 candidate: &TransitionCandidate,
7244 staged: &HashMap<String, Value>,
7245 processed_input: &str,
7246 ) -> Result<bool> {
7247 if !self
7248 .approve_transition_target(&candidate.from_state, candidate.target())
7249 .await?
7250 {
7251 return Ok(false);
7252 }
7253 self.apply_pre_response_transition_candidate(candidate, staged, processed_input)
7254 .await
7255 }
7256
7257 async fn handle_transition_miss(&self, current_state: &str) -> Result<bool> {
7259 let Some(ref sm) = self.state_machine else {
7260 return Ok(false);
7261 };
7262 sm.increment_no_transition();
7263 let Some(fallback) = sm.check_fallback() else {
7264 return Ok(false);
7265 };
7266 self.commit_transition_target(current_state, &fallback, "fallback after no transitions")
7267 .await
7268 }
7269
7270 async fn evaluate_transitions(&self, user_message: &str, response: &str) -> Result<bool> {
7272 let Some((transitions, current_state)) = self.transitions_available_for_commit() else {
7273 return Ok(false);
7274 };
7275 if transitions.is_empty() {
7276 return Ok(false);
7277 }
7278 if let Some(candidate) = self
7279 .select_transition_candidate(user_message, response)
7280 .await?
7281 {
7282 return self.commit_transition_candidate(&candidate).await;
7283 }
7284 self.handle_transition_miss(¤t_state).await
7285 }
7286
7287 async fn try_pre_response_transition(
7289 &self,
7290 processed_input: &str,
7291 ) -> Result<Option<AgentResponse>> {
7292 let optimization = &self.runtime_config.optimization;
7293 if !optimization.enabled || !optimization.pre_response_deterministic_transitions {
7294 return Ok(None);
7295 }
7296 let Some((transitions, current_state)) = self.transitions_available_for_commit() else {
7297 return Ok(None);
7298 };
7299 let eligible: Vec<Transition> = transitions
7300 .into_iter()
7301 .filter(|transition| !transition.requires_response)
7302 .filter(|transition| matches!(transition.timing, TransitionTiming::PreResponse))
7303 .collect();
7304 if eligible.is_empty() {
7305 return Ok(None);
7306 }
7307
7308 let empty_staged = HashMap::new();
7309 let mut extracted_staged: Option<HashMap<String, Value>> = None;
7310 let mut selected: Option<(TransitionCandidate, HashMap<String, Value>)> = None;
7311
7312 for transition in &eligible {
7313 let use_extractors = optimization.pre_response_extractors || transition.run_extractors;
7314 let staged_for_eval = if use_extractors {
7315 if extracted_staged.is_none() {
7316 extracted_staged =
7317 Some(self.run_context_extractors_staged(processed_input).await);
7318 }
7319 extracted_staged.as_ref().unwrap_or(&empty_staged)
7320 } else {
7321 &empty_staged
7322 };
7323
7324 if let Some(candidate) = self.select_deterministic_transition_candidate(
7325 processed_input,
7326 ¤t_state,
7327 std::slice::from_ref(transition),
7328 staged_for_eval,
7329 ) {
7330 let staged_for_commit = if use_extractors {
7331 staged_for_eval.clone()
7332 } else {
7333 HashMap::new()
7334 };
7335 selected = Some((candidate, staged_for_commit));
7336 break;
7337 }
7338 }
7339
7340 let Some((candidate, staged)) = selected else {
7341 return Ok(None);
7342 };
7343
7344 if !self
7345 .commit_pre_response_transition_candidate(&candidate, &staged, processed_input)
7346 .await?
7347 {
7348 return Ok(None);
7349 }
7350 self.redispatch_current_state(processed_input)
7351 .await
7352 .map(Some)
7353 }
7354
7355 async fn try_speculative_branches(
7360 &self,
7361 processed_input: &str,
7362 input_context: &HashMap<String, Value>,
7363 ) -> Result<Option<AgentResponse>> {
7364 let optimization = &self.runtime_config.optimization;
7365 if !optimization.enabled {
7366 return Ok(None);
7367 }
7368
7369 let effective_reasoning_mode = self.get_effective_reasoning_config().mode.clone();
7370 if !matches!(
7371 effective_reasoning_mode,
7372 ReasoningMode::None | ReasoningMode::Auto
7373 ) {
7374 return Ok(None);
7375 }
7376
7377 let mut transition_enabled =
7378 optimization.speculative_state_transitions && self.has_parallel_transition_candidates();
7379 let mut skill_enabled = optimization.speculative_skill_routing
7380 && self.skill_router.is_some()
7381 && self.pending_skill_id.read().is_none();
7382 let mut reasoning_enabled = optimization.speculative_reasoning_auto
7383 && matches!(effective_reasoning_mode, ReasoningMode::Auto);
7384
7385 if matches!(effective_reasoning_mode, ReasoningMode::Auto)
7386 && (!reasoning_enabled || optimization.max_speculative_llm_calls_per_turn < 2)
7387 {
7388 return Ok(None);
7389 }
7390
7391 if !transition_enabled && !skill_enabled && !reasoning_enabled {
7392 return Ok(None);
7393 }
7394
7395 let mut optional_slots = optimization.max_parallel_runtime_tasks.saturating_sub(1);
7396 let mut speculative_call_slots = optimization
7397 .max_speculative_llm_calls_per_turn
7398 .saturating_sub(1);
7399 if reasoning_enabled {
7400 if optional_slots == 0 || speculative_call_slots == 0 {
7401 return Ok(None);
7402 }
7403 optional_slots -= 1;
7404 speculative_call_slots -= 1;
7405 }
7406 if transition_enabled {
7407 if optional_slots == 0 {
7408 transition_enabled = false;
7409 } else {
7410 optional_slots -= 1;
7411 }
7412 }
7413 if skill_enabled && (optional_slots == 0 || speculative_call_slots == 0) {
7414 skill_enabled = false;
7415 }
7416
7417 if !transition_enabled && !skill_enabled && !reasoning_enabled {
7418 return Ok(None);
7419 }
7420
7421 let main_kind = if transition_enabled {
7422 RuntimeOptimizationKind::ParallelStateTransition
7423 } else if skill_enabled {
7424 RuntimeOptimizationKind::SpeculativeSkillRouting
7425 } else {
7426 RuntimeOptimizationKind::SpeculativeReasoningAuto
7427 };
7428 if !self.reserve_active_speculative_llm_call(main_kind) {
7429 return Ok(None);
7430 }
7431
7432 let mut branch_set = ScheduledBranchSet::new(optimization.max_parallel_runtime_tasks)?;
7433 let main_branch = RuntimeBranch::new(
7434 RuntimeTaskPurpose::MainResponse,
7435 main_kind,
7436 RuntimeTaskPriority::Normal,
7437 RuntimeCommitBehavior::FinalResponse,
7438 );
7439 let transition_branch = RuntimeBranch::new(
7440 RuntimeTaskPurpose::StateTransition,
7441 RuntimeOptimizationKind::ParallelStateTransition,
7442 RuntimeTaskPriority::Critical,
7443 RuntimeCommitBehavior::TransitionDecision,
7444 );
7445 let skill_branch = RuntimeBranch::new(
7446 RuntimeTaskPurpose::SkillRouting,
7447 RuntimeOptimizationKind::SpeculativeSkillRouting,
7448 RuntimeTaskPriority::High,
7449 RuntimeCommitBehavior::SkillSelection,
7450 );
7451 let reasoning_branch = RuntimeBranch::new(
7452 RuntimeTaskPurpose::ReasoningJudge,
7453 RuntimeOptimizationKind::SpeculativeReasoningAuto,
7454 RuntimeTaskPriority::Normal,
7455 RuntimeCommitBehavior::ReasoningDecision,
7456 );
7457 let main_id = main_branch.branch_id();
7458 let transition_id = transition_branch.branch_id();
7459 let skill_id = skill_branch.branch_id();
7460 let reasoning_id = reasoning_branch.branch_id();
7461
7462 let main_id_for_future = main_id.clone();
7463 if !branch_set.schedule(
7464 main_branch,
7465 Box::pin(async move {
7466 match crate::optimization::observability::with_branch_observation(
7467 &main_id_for_future,
7468 main_kind,
7469 RuntimeCommitBehavior::FinalResponse,
7470 self.generate_main_response_draft(processed_input, &ReasoningMode::None),
7471 )
7472 .await
7473 {
7474 Ok(draft) => RuntimeBranchResult::MainDraft(draft),
7475 Err(error) => RuntimeBranchResult::Failed(error),
7476 }
7477 }),
7478 ) {
7479 return Ok(None);
7480 }
7481
7482 if transition_enabled {
7483 let transition_id_for_future = transition_id.clone();
7484 if !branch_set.schedule(
7485 transition_branch,
7486 Box::pin(async move {
7487 match crate::optimization::observability::with_branch_observation(
7488 &transition_id_for_future,
7489 RuntimeOptimizationKind::ParallelStateTransition,
7490 RuntimeCommitBehavior::TransitionDecision,
7491 self.select_parallel_transition_candidate(processed_input),
7492 )
7493 .await
7494 {
7495 Ok(ParallelTransitionSelection::Candidate(candidate)) => {
7496 RuntimeBranchResult::Transition(Some(candidate))
7497 }
7498 Ok(ParallelTransitionSelection::NoMatch) => {
7499 RuntimeBranchResult::Transition(None)
7500 }
7501 Ok(ParallelTransitionSelection::ReservationExhausted) => {
7502 RuntimeBranchResult::Cancelled
7503 }
7504 Err(error) => RuntimeBranchResult::Failed(error),
7505 }
7506 }),
7507 ) {
7508 transition_enabled = false;
7509 }
7510 }
7511
7512 if skill_enabled {
7513 let skill_id_for_future = skill_id.clone();
7514 if !branch_set.schedule(
7515 skill_branch,
7516 Box::pin(async move {
7517 if !self.reserve_active_speculative_llm_call(
7518 RuntimeOptimizationKind::SpeculativeSkillRouting,
7519 ) {
7520 return RuntimeBranchResult::Cancelled;
7521 }
7522 match crate::optimization::observability::with_branch_observation(
7523 &skill_id_for_future,
7524 RuntimeOptimizationKind::SpeculativeSkillRouting,
7525 RuntimeCommitBehavior::SkillSelection,
7526 self.select_skill_candidate(processed_input),
7527 )
7528 .await
7529 {
7530 Ok(candidate) => RuntimeBranchResult::Skill(candidate),
7531 Err(error) => RuntimeBranchResult::Failed(error),
7532 }
7533 }),
7534 ) {
7535 skill_enabled = false;
7536 }
7537 }
7538
7539 if reasoning_enabled {
7540 let reasoning_id_for_future = reasoning_id.clone();
7541 if !branch_set.schedule(
7542 reasoning_branch,
7543 Box::pin(async move {
7544 if !self.reserve_active_speculative_llm_call(
7545 RuntimeOptimizationKind::SpeculativeReasoningAuto,
7546 ) {
7547 return RuntimeBranchResult::Cancelled;
7548 }
7549 match crate::optimization::observability::with_branch_observation(
7550 &reasoning_id_for_future,
7551 RuntimeOptimizationKind::SpeculativeReasoningAuto,
7552 RuntimeCommitBehavior::ReasoningDecision,
7553 self.determine_reasoning_mode_strict(processed_input),
7554 )
7555 .await
7556 {
7557 Ok(mode) => RuntimeBranchResult::Reasoning(mode),
7558 Err(error) => RuntimeBranchResult::Failed(error),
7559 }
7560 }),
7561 ) {
7562 reasoning_enabled = false;
7563 }
7564 }
7565
7566 if matches!(effective_reasoning_mode, ReasoningMode::Auto) && !reasoning_enabled {
7567 self.finalize_pending_branches(branch_set.cancel_pending());
7568 return Ok(None);
7569 }
7570
7571 if !transition_enabled && !skill_enabled && !reasoning_enabled {
7572 self.finalize_pending_branches(branch_set.cancel_pending());
7573 return Ok(None);
7574 }
7575
7576 let mut main_pending = true;
7577 let mut skill_pending = skill_enabled;
7578 let mut reasoning_pending = reasoning_enabled;
7579 let mut transition_finalized = !transition_enabled;
7580 let mut skill_finalized = !skill_enabled && self.skill_router.is_none();
7583 let mut reasoning_finalized = !reasoning_enabled;
7584 let mut main_result: Option<Result<MainResponseDraft>> = None;
7585 let mut transition_candidate: Option<TransitionCandidate> = None;
7586 let mut skill_candidate: Option<SkillCandidate> = None;
7587 let mut reasoning_decision: Option<ReasoningMode> = None;
7588 let mut transition_fallback_required = false;
7589 let mut skill_fallback_required = false;
7590 let mut reasoning_fallback_required = false;
7591
7592 loop {
7593 if let Some(candidate) = transition_candidate.take() {
7594 if self
7595 .approve_transition_target(&candidate.from_state, candidate.target())
7596 .await?
7597 {
7598 self.finalize_pending_branches(branch_set.cancel_pending());
7600 if !main_pending {
7601 self.finalize_branch_loss(
7602 &main_id,
7603 main_kind,
7604 RuntimeCommitBehavior::FinalResponse,
7605 false,
7606 main_result.as_ref().map(|result| result.is_err()),
7607 );
7608 }
7609 if skill_enabled && !skill_pending {
7610 self.finalize_branch_loss(
7611 &skill_id,
7612 RuntimeOptimizationKind::SpeculativeSkillRouting,
7613 RuntimeCommitBehavior::SkillSelection,
7614 false,
7615 Some(false),
7616 );
7617 }
7618 if reasoning_enabled && !reasoning_pending {
7619 self.finalize_branch_loss(
7620 &reasoning_id,
7621 RuntimeOptimizationKind::SpeculativeReasoningAuto,
7622 RuntimeCommitBehavior::ReasoningDecision,
7623 false,
7624 Some(false),
7625 );
7626 }
7627 if !self
7628 .apply_pre_response_transition_candidate(
7629 &candidate,
7630 &HashMap::new(),
7631 processed_input,
7632 )
7633 .await?
7634 {
7635 self.finalize_optional_branch(
7636 &transition_id,
7637 RuntimeOptimizationKind::ParallelStateTransition,
7638 RuntimeCommitBehavior::TransitionDecision,
7639 "discarded",
7640 false,
7641 );
7642 return Ok(None);
7643 }
7644 self.finalize_optional_branch(
7645 &transition_id,
7646 RuntimeOptimizationKind::ParallelStateTransition,
7647 RuntimeCommitBehavior::TransitionDecision,
7648 "committed",
7649 true,
7650 );
7651 return self
7652 .redispatch_current_state(processed_input)
7653 .await
7654 .map(Some);
7655 }
7656 self.finalize_optional_branch(
7657 &transition_id,
7658 RuntimeOptimizationKind::ParallelStateTransition,
7659 RuntimeCommitBehavior::TransitionDecision,
7660 "discarded",
7661 false,
7662 );
7663 transition_finalized = true;
7664 }
7665
7666 if transition_finalized
7675 && !skill_finalized
7676 && !skill_enabled
7677 && self.skill_router.is_some()
7678 {
7679 match self.select_skill_candidate(processed_input).await {
7680 Ok(Some(candidate)) => skill_candidate = Some(candidate),
7681 Ok(None) => {}
7682 Err(error) => {
7683 self.finalize_pending_branches(branch_set.cancel_pending());
7685 return Err(error);
7686 }
7687 }
7688 skill_finalized = true;
7689 }
7690
7691 if transition_finalized && skill_candidate.is_some() {
7692 let candidate = skill_candidate.take().unwrap();
7693 if skill_enabled {
7695 self.finalize_optional_branch(
7696 &skill_id,
7697 RuntimeOptimizationKind::SpeculativeSkillRouting,
7698 RuntimeCommitBehavior::SkillSelection,
7699 "committed",
7700 true,
7701 );
7702 }
7703 if !main_pending {
7704 self.finalize_branch_loss(
7705 &main_id,
7706 main_kind,
7707 RuntimeCommitBehavior::FinalResponse,
7708 false,
7709 main_result.as_ref().map(|result| result.is_err()),
7710 );
7711 }
7712 if reasoning_enabled && !reasoning_pending {
7713 self.finalize_branch_loss(
7714 &reasoning_id,
7715 RuntimeOptimizationKind::SpeculativeReasoningAuto,
7716 RuntimeCommitBehavior::ReasoningDecision,
7717 false,
7718 Some(false),
7719 );
7720 }
7721 self.finalize_pending_branches(branch_set.cancel_pending());
7722 return self
7723 .commit_winning_skill_candidate(candidate, processed_input, input_context)
7724 .await;
7725 }
7726
7727 if transition_finalized
7728 && skill_finalized
7729 && let Some(reasoning_mode) = reasoning_decision.take()
7730 {
7731 if !matches!(reasoning_mode, ReasoningMode::None) {
7732 self.finalize_optional_branch(
7733 &reasoning_id,
7734 RuntimeOptimizationKind::SpeculativeReasoningAuto,
7735 RuntimeCommitBehavior::ReasoningDecision,
7736 "committed",
7737 true,
7738 );
7739 if !main_pending {
7740 self.finalize_branch_loss(
7741 &main_id,
7742 main_kind,
7743 RuntimeCommitBehavior::FinalResponse,
7744 false,
7745 main_result.as_ref().map(|result| result.is_err()),
7746 );
7747 }
7748 self.finalize_pending_branches(branch_set.cancel_pending());
7749 self.commit_root_user_message(processed_input).await?;
7750 return if matches!(reasoning_mode, ReasoningMode::PlanAndExecute) {
7751 self.handle_plan_and_execute(processed_input, input_context, true)
7752 .await
7753 .map(Some)
7754 } else {
7755 self.run_committed_response_loop_with_reasoning(
7756 processed_input,
7757 input_context,
7758 reasoning_mode,
7759 true,
7760 )
7761 .await
7762 .map(Some)
7763 };
7764 }
7765 self.finalize_optional_branch(
7766 &reasoning_id,
7767 RuntimeOptimizationKind::SpeculativeReasoningAuto,
7768 RuntimeCommitBehavior::ReasoningDecision,
7769 "committed",
7770 true,
7771 );
7772 reasoning_finalized = true;
7773 }
7774
7775 if transition_finalized && skill_finalized && reasoning_finalized {
7776 if transition_fallback_required
7777 || skill_fallback_required
7778 || reasoning_fallback_required
7779 {
7780 if !main_pending {
7781 self.finalize_branch_loss(
7782 &main_id,
7783 main_kind,
7784 RuntimeCommitBehavior::FinalResponse,
7785 false,
7786 main_result.as_ref().map(|result| result.is_err()),
7787 );
7788 }
7789 self.finalize_pending_branches(branch_set.cancel_pending());
7790 return Ok(None);
7791 }
7792
7793 if let Some(result) = main_result.take() {
7794 let draft = match result {
7795 Ok(draft) => draft,
7796 Err(error) => {
7797 self.finalize_optional_branch(
7798 &main_id,
7799 main_kind,
7800 RuntimeCommitBehavior::FinalResponse,
7801 "failed",
7802 false,
7803 );
7804 self.finalize_pending_branches(branch_set.cancel_pending());
7805 return Err(error);
7806 }
7807 };
7808 self.finalize_optional_branch(
7809 &main_id,
7810 main_kind,
7811 RuntimeCommitBehavior::FinalResponse,
7812 "committed",
7813 true,
7814 );
7815 self.finalize_pending_branches(branch_set.cancel_pending());
7816 return self
7817 .commit_main_response_draft(
7818 processed_input,
7819 input_context,
7820 draft,
7821 ReasoningMode::None,
7822 reasoning_enabled,
7823 )
7824 .await
7825 .map(Some);
7826 }
7827 }
7828
7829 if branch_set.is_empty() {
7830 return Ok(None);
7831 }
7832
7833 let Some(outcome) = branch_set.next_completed().await else {
7834 return Ok(None);
7835 };
7836 let branch_id = outcome.branch.branch_id();
7837 match outcome.result {
7838 RuntimeBranchResult::MainDraft(draft) => {
7839 main_pending = false;
7840 main_result = Some(Ok(draft));
7841 }
7842 RuntimeBranchResult::Transition(candidate) => {
7843 if let Some(candidate) = candidate {
7844 transition_candidate = Some(candidate);
7845 } else {
7846 self.finalize_optional_branch(
7847 &transition_id,
7848 RuntimeOptimizationKind::ParallelStateTransition,
7849 RuntimeCommitBehavior::TransitionDecision,
7850 "discarded",
7851 false,
7852 );
7853 transition_finalized = true;
7854 }
7855 }
7856 RuntimeBranchResult::Skill(candidate) => {
7857 skill_pending = false;
7858 if let Some(candidate) = candidate {
7859 skill_candidate = Some(candidate);
7860 } else {
7861 self.finalize_optional_branch(
7862 &skill_id,
7863 RuntimeOptimizationKind::SpeculativeSkillRouting,
7864 RuntimeCommitBehavior::SkillSelection,
7865 "discarded",
7866 false,
7867 );
7868 skill_finalized = true;
7869 }
7870 }
7871 RuntimeBranchResult::Reasoning(mode) => {
7872 reasoning_pending = false;
7873 reasoning_decision = Some(mode);
7874 }
7875 RuntimeBranchResult::Failed(error) => {
7876 if branch_id == main_id {
7877 main_pending = false;
7878 main_result = Some(Err(error));
7879 } else if branch_id == transition_id {
7880 self.finalize_optional_branch(
7881 &transition_id,
7882 RuntimeOptimizationKind::ParallelStateTransition,
7883 RuntimeCommitBehavior::TransitionDecision,
7884 "failed",
7885 false,
7886 );
7887 transition_finalized = true;
7888 } else if branch_id == skill_id {
7889 skill_pending = false;
7890 self.finalize_optional_branch(
7891 &skill_id,
7892 RuntimeOptimizationKind::SpeculativeSkillRouting,
7893 RuntimeCommitBehavior::SkillSelection,
7894 "failed",
7895 false,
7896 );
7897 skill_finalized = true;
7898 } else if branch_id == reasoning_id {
7899 reasoning_pending = false;
7900 self.finalize_optional_branch(
7901 &reasoning_id,
7902 RuntimeOptimizationKind::SpeculativeReasoningAuto,
7903 RuntimeCommitBehavior::ReasoningDecision,
7904 "failed",
7905 false,
7906 );
7907 reasoning_finalized = true;
7908 }
7909 }
7910 RuntimeBranchResult::Cancelled => {
7911 self.finalize_optional_branch(
7912 &branch_id,
7913 outcome.branch.optimization,
7914 outcome.branch.commit_behavior,
7915 "cancelled",
7916 false,
7917 );
7918 if branch_id == main_id {
7919 main_pending = false;
7920 main_result =
7921 Some(Err(AgentError::Other("main branch cancelled".to_string())));
7922 } else if branch_id == transition_id {
7923 transition_finalized = true;
7924 transition_fallback_required = true;
7925 } else if branch_id == skill_id {
7926 skill_pending = false;
7927 skill_finalized = true;
7928 skill_fallback_required = true;
7929 } else if branch_id == reasoning_id {
7930 reasoning_pending = false;
7931 reasoning_finalized = true;
7932 reasoning_fallback_required = true;
7933 }
7934 }
7935 }
7936 }
7937 }
7938
7939 fn finalize_pending_branches(&self, branches: Vec<RuntimeBranch>) {
7940 for branch in branches {
7941 self.finalize_optional_branch(
7942 &branch.branch_id(),
7943 branch.optimization,
7944 branch.commit_behavior,
7945 "cancelled",
7946 false,
7947 );
7948 }
7949 }
7950
7951 fn finalize_branch_loss(
7956 &self,
7957 branch_id: &str,
7958 optimization: RuntimeOptimizationKind,
7959 commit_behavior: RuntimeCommitBehavior,
7960 pending: bool,
7961 completed_failed: Option<bool>,
7962 ) {
7963 let status = if pending {
7964 "cancelled"
7965 } else if completed_failed.unwrap_or(false) {
7966 "failed"
7967 } else {
7968 "discarded"
7969 };
7970 self.finalize_optional_branch(branch_id, optimization, commit_behavior, status, false);
7971 }
7972
7973 fn finalize_optional_branch(
7978 &self,
7979 branch_id: &str,
7980 optimization: RuntimeOptimizationKind,
7981 commit_behavior: RuntimeCommitBehavior,
7982 status: &str,
7983 winner: bool,
7984 ) {
7985 crate::optimization::observability::finalize_branch(
7986 self.observability_manager.as_ref(),
7987 branch_id,
7988 status,
7989 winner,
7990 optimization,
7991 commit_behavior,
7992 );
7993 }
7994
7995 fn has_parallel_transition_candidates(&self) -> bool {
8000 self.transitions_available_for_commit()
8001 .map(|(transitions, _)| {
8002 transitions
8003 .iter()
8004 .any(|transition| matches!(transition.timing, TransitionTiming::Parallel))
8005 })
8006 .unwrap_or(false)
8007 }
8008
8009 async fn select_parallel_transition_candidate(
8014 &self,
8015 processed_input: &str,
8016 ) -> Result<ParallelTransitionSelection> {
8017 let Some((transitions, current_state)) = self.transitions_available_for_commit() else {
8018 return Ok(ParallelTransitionSelection::NoMatch);
8019 };
8020 let parallel: Vec<Transition> = transitions
8021 .into_iter()
8022 .filter(|transition| matches!(transition.timing, TransitionTiming::Parallel))
8023 .filter(|transition| !transition.requires_response)
8024 .collect();
8025 if parallel.is_empty() {
8026 return Ok(ParallelTransitionSelection::NoMatch);
8027 }
8028 let empty_staged = HashMap::new();
8029 if let Some(candidate) = self.select_deterministic_transition_candidate(
8030 processed_input,
8031 ¤t_state,
8032 ¶llel,
8033 &empty_staged,
8034 ) {
8035 return Ok(ParallelTransitionSelection::Candidate(candidate));
8036 }
8037 let when_transitions: Vec<(usize, &Transition)> = parallel
8038 .iter()
8039 .enumerate()
8040 .filter(|(_, transition)| !transition.when.trim().is_empty())
8041 .collect();
8042 if when_transitions.is_empty() {
8043 return Ok(ParallelTransitionSelection::NoMatch);
8044 }
8045 let llm = self
8046 .llm_registry
8047 .router()
8048 .or_else(|_| self.llm_registry.default())
8049 .map_err(|e| AgentError::Config(e.to_string()))?;
8050 let conditions = when_transitions
8051 .iter()
8052 .enumerate()
8053 .map(|(display_idx, (_, transition))| {
8054 format!("{}. {}", display_idx + 1, transition.when)
8055 })
8056 .collect::<Vec<_>>()
8057 .join("\n");
8058 if !self
8059 .reserve_active_speculative_llm_call(RuntimeOptimizationKind::ParallelStateTransition)
8060 {
8061 return Ok(ParallelTransitionSelection::ReservationExhausted);
8062 }
8063 let context_preview = self.branch_context_preview();
8064 let prompt = format!(
8065 "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-{}).",
8066 current_state,
8067 processed_input,
8068 context_preview,
8069 conditions,
8070 when_transitions.len()
8071 );
8072 let response = self
8073 .observe_purpose(
8074 ObservationPurpose::StateTransitionEvaluation,
8075 llm.complete(&[ChatMessage::user(prompt)], None),
8076 )
8077 .await
8078 .map_err(|e| AgentError::LLM(e.to_string()))?;
8079 let choice = response.content.trim().parse::<usize>().unwrap_or(0);
8080 if choice == 0 || choice > when_transitions.len() {
8081 return Ok(ParallelTransitionSelection::NoMatch);
8082 }
8083 let transition = when_transitions[choice - 1].1.clone();
8084 Ok(ParallelTransitionSelection::Candidate(
8085 TransitionCandidate::new(
8086 current_state,
8087 transition.clone(),
8088 Self::transition_reason(&transition),
8089 ),
8090 ))
8091 }
8092
8093 async fn redispatch_current_state(&self, processed_input: &str) -> Result<AgentResponse> {
8095 const MAX_REDISPATCH_DEPTH: u32 = 3;
8096 let current_depth = *self.redispatch_depth.read();
8097 if current_depth >= MAX_REDISPATCH_DEPTH {
8098 warn!(depth = current_depth, "Re-dispatch depth limit reached");
8099 let response = AgentResponse::new("");
8100 self.finish_turn_if_root(&response).await?;
8101 return Ok(response);
8102 }
8103 *self.redispatch_depth.write() += 1;
8104 if let Some(context) = self.active_turn_context.write().as_mut() {
8105 context.enter_redispatch();
8106 }
8107 let result = Box::pin(self.run_loop_internal(processed_input)).await;
8108 *self.redispatch_depth.write() -= 1;
8109 if let Some(context) = self.active_turn_context.write().as_mut() {
8110 context.exit_redispatch();
8111 }
8112 let response = result?;
8113 self.finish_turn_if_root(&response).await?;
8114 Ok(response)
8115 }
8116
8117 async fn finish_turn_if_root(&self, response: &AgentResponse) -> Result<()> {
8119 if *self.redispatch_depth.read() == 0 {
8120 self.post_turn_session_lifecycle().await?;
8121 if let Some(context) = self.active_turn_context.write().as_mut() {
8122 context.mark_post_turn_lifecycle_completed();
8123 }
8124 self.hooks.on_response(response).await;
8125 self.end_root_turn();
8126 }
8127 Ok(())
8128 }
8129
8130 async fn execute_state_exit_actions(&self, state_path: &str) {
8132 if let Some(ref sm) = self.state_machine
8133 && let Some(def) = sm.get_definition(state_path)
8134 && !def.on_exit.is_empty()
8135 {
8136 debug!(state = %state_path, count = def.on_exit.len(), "Executing on_exit actions");
8137 self.execute_state_actions(&def.on_exit).await;
8138 }
8139 }
8140
8141 fn state_was_previously_entered(
8143 state_path: &str,
8144 from_state: &str,
8145 history_before: &[StateTransitionEvent],
8146 ) -> bool {
8147 state_path == from_state
8148 || history_before
8149 .iter()
8150 .any(|event| event.from == state_path || event.to == state_path)
8151 }
8152
8153 async fn execute_state_enter_actions(&self, state_path: &str, is_reentry: bool) {
8155 if let Some(ref sm) = self.state_machine
8156 && let Some(def) = sm.get_definition(state_path)
8157 {
8158 if is_reentry && !def.on_reenter.is_empty() {
8159 debug!(state = %state_path, count = def.on_reenter.len(), "Executing on_reenter actions");
8160 self.execute_state_actions(&def.on_reenter).await;
8161 } else if !def.on_enter.is_empty() {
8162 debug!(state = %state_path, count = def.on_enter.len(), "Executing on_enter actions");
8163 self.execute_state_actions(&def.on_enter).await;
8164 }
8165 }
8166 }
8167
8168 async fn execute_state_actions(&self, actions: &[StateAction]) {
8170 for (action_index, action) in actions.iter().enumerate() {
8171 match action {
8172 StateAction::Tool { tool, args } => {
8173 let raw_args = args.clone().unwrap_or(Value::Object(Default::default()));
8174 let args_value = self.render_action_args(&raw_args);
8175 let state = self.state_machine.as_ref().map(|sm| sm.current());
8176 let request = ToolExecutionRequest::new(
8177 uuid::Uuid::new_v4().to_string(),
8178 tool.clone(),
8179 args_value,
8180 ToolCallSource::StateAction {
8181 state,
8182 action_index,
8183 },
8184 );
8185 match self.execute_tool_record(request).await {
8186 Ok(record) if record.success => {
8187 debug!(tool = %record.canonical_id, "State action: tool executed");
8188 let _ = self.context_manager.set(
8189 "last_tool_result",
8190 serde_json::Value::String(record.model_output_string()),
8191 );
8192 let _ = self.context_manager.set(
8193 "last_tool_record",
8194 serde_json::to_value(record).unwrap_or(Value::Null),
8195 );
8196 }
8197 Ok(record) => {
8198 warn!(tool = %record.canonical_id, error = %record.output, "State action: tool failed");
8199 }
8200 Err(e) => {
8201 warn!(tool = %tool, error = %e, "State action: tool failed")
8202 }
8203 }
8204 }
8205 StateAction::Skill { skill } => {
8206 if let Some(ref executor) = self.skill_executor {
8207 if let Some(def) = self.skills.iter().find(|s| s.id == *skill) {
8208 match executor
8209 .execute_with_invoker(def, "", serde_json::json!({}), self)
8210 .await
8211 {
8212 Ok(_) => debug!(skill = %skill, "State action: skill executed"),
8213 Err(e) => {
8214 warn!(skill = %skill, error = %e, "State action: skill failed")
8215 }
8216 }
8217 } else {
8218 warn!(skill = %skill, "State action: skill not found");
8219 }
8220 }
8221 }
8222 StateAction::SetContext { set_context } => {
8223 for (key, value) in set_context {
8224 if let Err(e) = self.context_manager.set(key, value.clone()) {
8225 warn!(key = %key, error = %e, "State action: set_context failed");
8226 } else {
8227 debug!(key = %key, "State action: context set");
8228 }
8229 }
8230 }
8231 StateAction::Prompt {
8232 prompt,
8233 llm,
8234 store_as,
8235 } => {
8236 let llm_result = if let Some(alias) = llm {
8237 self.llm_registry.get(alias)
8238 } else {
8239 self.llm_registry.default()
8240 };
8241 match llm_result {
8242 Ok(llm_provider) => {
8243 let context = self.build_context_with_overlays();
8245 let rendered_prompt = self
8246 .template_renderer
8247 .render(prompt, &context)
8248 .unwrap_or_else(|_| prompt.clone());
8249 let recent =
8250 self.memory.get_messages(Some(5)).await.unwrap_or_default();
8251 let mut messages: Vec<ChatMessage> = recent;
8252 messages.push(ChatMessage::user(&rendered_prompt));
8253 match self
8254 .observe_purpose(
8255 ObservationPurpose::StateAction,
8256 llm_provider.complete(&messages, None),
8257 )
8258 .await
8259 {
8260 Ok(response) => {
8261 if let Some(key) = store_as {
8262 let _ = self
8263 .context_manager
8264 .set(key, Value::String(response.content));
8265 debug!(key = %key, "State action: prompt result stored");
8266 }
8267 }
8268 Err(e) => {
8269 warn!(error = %e, "State action: prompt LLM call failed");
8270 }
8271 }
8272 }
8273 Err(e) => {
8274 warn!(error = %e, "State action: LLM not found for prompt");
8275 }
8276 }
8277 }
8278 }
8279 }
8280 }
8281
8282 async fn run_context_extractors_staged(&self, user_message: &str) -> HashMap<String, Value> {
8283 let extractors = match &self.state_machine {
8284 Some(sm) => match sm.current_definition() {
8285 Some(def) if !def.extract.is_empty() => def.extract.clone(),
8286 _ => return HashMap::new(),
8287 },
8288 None => return HashMap::new(),
8289 };
8290
8291 let mut staged = HashMap::new();
8292 for extractor in &extractors {
8293 let prompt = if let Some(ref custom) = extractor.llm_extract {
8294 format!(
8295 "User message:\n\"{}\"\n\nInstruction:\n{}",
8296 user_message, custom
8297 )
8298 } else if let Some(ref desc) = extractor.description {
8299 format!(
8300 "From the following message, extract: {}\n\n\
8301 Message: \"{}\"\n\n\
8302 If the information is present, return ONLY the extracted value.\n\
8303 If NOT present, return exactly: __NONE__",
8304 desc, user_message
8305 )
8306 } else {
8307 continue;
8308 };
8309
8310 let llm = match self
8311 .llm_registry
8312 .get(&extractor.llm)
8313 .or_else(|_| self.llm_registry.get("router"))
8314 .or_else(|_| self.llm_registry.get("default"))
8315 {
8316 Ok(llm) => llm,
8317 Err(e) => {
8318 warn!(key = %extractor.key, error = %e, "Extractor LLM not found");
8319 continue;
8320 }
8321 };
8322
8323 let messages = vec![ChatMessage::user(&prompt)];
8324 match self
8325 .observe_purpose(
8326 ObservationPurpose::ContextExtraction,
8327 llm.complete(&messages, None),
8328 )
8329 .await
8330 {
8331 Ok(response) => {
8332 let value = response.content.trim().to_string();
8333 if value != "__NONE__" && !value.is_empty() {
8334 staged.insert(
8335 extractor.key.clone(),
8336 serde_json::Value::String(value.clone()),
8337 );
8338 debug!(key = %extractor.key, value = %value, "Context extracted");
8339 } else if extractor.required {
8340 warn!(key = %extractor.key, "Required extraction returned no value");
8341 }
8342 }
8343 Err(e) => {
8344 warn!(key = %extractor.key, error = %e, "Context extraction LLM call failed");
8345 }
8346 }
8347 }
8348 staged
8349 }
8350
8351 fn commit_staged_context_writes(&self, staged: &HashMap<String, Value>) {
8352 for (key, value) in staged {
8353 if let Err(error) = self.context_manager.update(key, value.clone()) {
8354 warn!(key = %key, error = %error, "staged context write failed");
8355 }
8356 }
8357 }
8358
8359 async fn run_context_extractors(&self, user_message: &str) {
8361 let staged = self.run_context_extractors_staged(user_message).await;
8362 self.commit_staged_context_writes(&staged);
8363 }
8364
8365 async fn check_memory_compression(&self) -> Result<()> {
8366 if self.memory.needs_compression() {
8367 let result = self.memory.compress(None).await?;
8368 if let CompressResult::Compressed {
8369 messages_summarized,
8370 new_summary_length,
8371 tokens_saved,
8372 } = result
8373 {
8374 let event = MemoryCompressEvent::new(
8375 messages_summarized,
8376 tokens_saved,
8377 new_summary_length as u32,
8378 );
8379 self.hooks.on_memory_compress(&event).await;
8380 debug!(
8381 messages = messages_summarized,
8382 tokens_saved = tokens_saved,
8383 "Memory compressed"
8384 );
8385 }
8386 }
8387
8388 self.handle_memory_overflow().await?;
8390 self.check_memory_budget().await;
8391
8392 Ok(())
8393 }
8394
8395 async fn check_memory_budget(&self) {
8396 let Some(ref budget) = self.memory_token_budget else {
8397 return;
8398 };
8399
8400 let context = match self.memory.get_context().await {
8401 Ok(ctx) => ctx,
8402 Err(_) => return,
8403 };
8404
8405 let used_tokens = context.estimated_tokens();
8407 if budget.is_over_warn_threshold(used_tokens) {
8408 let event = MemoryBudgetEvent::new("memory", used_tokens, budget.total);
8409 self.hooks.on_memory_budget_warning(&event).await;
8410 debug!(
8411 used = used_tokens,
8412 total = budget.total,
8413 percent = event.usage_percent,
8414 "Memory budget warning"
8415 );
8416 }
8417
8418 if let Some(ref summary) = context.summary {
8420 let summary_tokens = ai_agents_memory::estimate_tokens(summary);
8421 let summary_budget = budget.allocation.summary;
8422 if summary_budget > 0 {
8423 let warn_threshold =
8424 (summary_budget as f64 * budget.warn_at_percent as f64 / 100.0) as u32;
8425 if summary_tokens >= warn_threshold {
8426 let event = MemoryBudgetEvent::new("summary", summary_tokens, summary_budget);
8427 self.hooks.on_memory_budget_warning(&event).await;
8428 }
8429 }
8430 }
8431
8432 let recent_tokens: u32 = context
8434 .messages
8435 .iter()
8436 .map(ai_agents_memory::estimate_message_tokens)
8437 .sum();
8438 let recent_budget = budget.allocation.recent_messages;
8439 if recent_budget > 0 {
8440 let warn_threshold =
8441 (recent_budget as f64 * budget.warn_at_percent as f64 / 100.0) as u32;
8442 if recent_tokens >= warn_threshold {
8443 let event = MemoryBudgetEvent::new("recent_messages", recent_tokens, recent_budget);
8444 self.hooks.on_memory_budget_warning(&event).await;
8445 }
8446 }
8447
8448 let relationship_budget = budget.allocation.relationships;
8449 if relationship_budget > 0 {
8450 let relationship_tokens = self
8451 .relationship_memory_text()
8452 .map(|text| ai_agents_memory::estimate_tokens(&text))
8453 .unwrap_or(0);
8454 let warn_threshold =
8455 (relationship_budget as f64 * budget.warn_at_percent as f64 / 100.0) as u32;
8456 if relationship_tokens >= warn_threshold {
8457 let event = MemoryBudgetEvent::new(
8458 "relationships",
8459 relationship_tokens,
8460 relationship_budget,
8461 );
8462 self.hooks.on_memory_budget_warning(&event).await;
8463 }
8464 }
8465 }
8466
8467 async fn handle_memory_overflow(&self) -> Result<()> {
8468 let Some(ref budget) = self.memory_token_budget else {
8469 return Ok(());
8470 };
8471
8472 let context = self.memory.get_context().await?;
8473 let used_tokens = context.estimated_tokens();
8474
8475 if used_tokens <= budget.total {
8476 return Ok(());
8477 }
8478
8479 match budget.overflow_strategy {
8480 OverflowStrategy::TruncateOldest => {
8481 let tokens_to_free = used_tokens - budget.total;
8482 let messages_to_evict = self.calculate_eviction_count(tokens_to_free);
8483 if messages_to_evict > 0 {
8484 self.evict_messages(messages_to_evict, EvictionReason::TokenBudgetExceeded)
8485 .await?;
8486 }
8487 }
8488 OverflowStrategy::SummarizeMore => {
8489 let max_attempts = context.total_messages.max(1);
8490 for _ in 0..max_attempts {
8491 match self.memory.compress(None).await? {
8492 CompressResult::Compressed {
8493 messages_summarized,
8494 ..
8495 } if messages_summarized > 0 => {
8496 let context = self.memory.get_context().await?;
8497 if context.estimated_tokens() <= budget.total {
8498 return Ok(());
8499 }
8500 }
8501 _ => break,
8502 }
8503 }
8504 let context = self.memory.get_context().await?;
8505 let used_tokens = context.estimated_tokens();
8506 if used_tokens > budget.total {
8507 return Err(AgentError::MemoryBudgetExceeded {
8508 used: used_tokens,
8509 budget: budget.total,
8510 });
8511 }
8512 }
8513 OverflowStrategy::Error => {
8514 return Err(AgentError::MemoryBudgetExceeded {
8515 used: used_tokens,
8516 budget: budget.total,
8517 });
8518 }
8519 }
8520 Ok(())
8521 }
8522
8523 fn calculate_eviction_count(&self, tokens_to_free: u32) -> usize {
8524 ((tokens_to_free as f64 / 50.0).ceil() as usize).max(1)
8526 }
8527
8528 async fn evict_messages(&self, count: usize, reason: EvictionReason) -> Result<()> {
8529 let evicted = self.memory.evict_oldest(count).await?;
8530 if !evicted.is_empty() {
8531 let event = MemoryEvictEvent {
8532 reason,
8533 messages_evicted: evicted.len(),
8534 importance_scores: vec![],
8535 };
8536 self.hooks.on_memory_evict(&event).await;
8537 debug!(count = evicted.len(), "Messages evicted from memory");
8538 }
8539 Ok(())
8540 }
8541
8542 #[instrument(skip(self, input), fields(agent = %self.info.name))]
8543 async fn determine_reasoning_mode(&self, input: &str) -> Result<ReasoningMode> {
8544 match self.determine_reasoning_mode_strict(input).await {
8545 Ok(mode) => Ok(mode),
8546 Err(_) => Ok(ReasoningMode::None),
8547 }
8548 }
8549
8550 async fn determine_reasoning_mode_strict(&self, input: &str) -> Result<ReasoningMode> {
8551 let effective_config = self.get_effective_reasoning_config();
8552
8553 if !matches!(effective_config.mode, ReasoningMode::Auto) {
8554 return Ok(effective_config.mode.clone());
8555 }
8556
8557 let judge_llm = effective_config
8558 .judge_llm
8559 .as_ref()
8560 .and_then(|alias| self.llm_registry.get(alias).ok())
8561 .or_else(|| self.llm_registry.router().ok())
8562 .or_else(|| self.llm_registry.default().ok());
8563
8564 let Some(llm) = judge_llm else {
8565 return Ok(ReasoningMode::None);
8566 };
8567
8568 let prompt = format!(
8569 r#"Analyze this user request and determine the appropriate reasoning mode.
8570
8571User request: "{}"
8572
8573Choose ONE of these modes:
8574- none: Simple queries, greetings, direct answers (fastest)
8575- cot: Complex analysis, multi-step reasoning, math problems
8576- react: Tasks requiring multiple tool calls with observation
8577- plan_and_execute: Complex multi-step tasks requiring coordination
8578
8579Respond with ONLY the mode name (none, cot, react, or plan_and_execute)."#,
8580 input
8581 );
8582
8583 let messages = vec![ChatMessage::user(&prompt)];
8584 let response = self
8585 .observe_purpose(
8586 ObservationPurpose::ReflectionDecision,
8587 llm.complete(&messages, None),
8588 )
8589 .await
8590 .map_err(|e| AgentError::LLM(e.to_string()))?;
8591
8592 let mode_str = response.content.trim().to_lowercase();
8593 Ok(match mode_str.as_str() {
8594 "cot" => ReasoningMode::CoT,
8595 "react" => ReasoningMode::React,
8596 "plan_and_execute" => ReasoningMode::PlanAndExecute,
8597 _ => ReasoningMode::None,
8598 })
8599 }
8600
8601 async fn should_reflect(&self, input: &str, response: &str) -> Result<bool> {
8602 let effective_config = self.get_effective_reflection_config();
8603
8604 if !effective_config.requires_evaluation() {
8605 return Ok(false);
8606 }
8607
8608 if effective_config.is_enabled() {
8609 return Ok(true);
8610 }
8611
8612 let evaluator_llm = effective_config
8613 .evaluator_llm
8614 .as_ref()
8615 .and_then(|alias| self.llm_registry.get(alias).ok())
8616 .or_else(|| self.llm_registry.router().ok())
8617 .or_else(|| self.llm_registry.default().ok());
8618
8619 let Some(llm) = evaluator_llm else {
8620 return Ok(false);
8621 };
8622
8623 let response_preview: String = response.chars().take(500).collect();
8624 let prompt = format!(
8625 r#"Should this response be evaluated for quality? Consider if it's a complex or important response.
8626
8627User query: "{}"
8628Response: "{}"
8629
8630Answer YES or NO only."#,
8631 input, response_preview
8632 );
8633
8634 let messages = vec![ChatMessage::user(&prompt)];
8635 let result = self
8636 .observe_purpose(
8637 ObservationPurpose::ReflectionDecision,
8638 llm.complete(&messages, None),
8639 )
8640 .await;
8641
8642 match result {
8643 Ok(resp) => Ok(resp.content.trim().to_uppercase().contains("YES")),
8644 Err(_) => Ok(false),
8645 }
8646 }
8647
8648 fn build_cot_system_prompt(&self, base_prompt: &str) -> String {
8649 format!(
8650 "{}\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>",
8651 base_prompt
8652 )
8653 }
8654
8655 fn build_react_system_prompt(&self, base_prompt: &str) -> String {
8656 format!(
8657 "{}\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>",
8658 base_prompt
8659 )
8660 }
8661
8662 async fn generate_plan(&self, input: &str) -> Result<Plan> {
8663 let effective = self.get_effective_reasoning_config();
8664 let planning_config = effective.get_planning();
8665
8666 let planner_llm = planning_config
8667 .and_then(|c| c.planner_llm.as_ref())
8668 .and_then(|alias| self.llm_registry.get(alias).ok())
8669 .or_else(|| self.llm_registry.router().ok())
8670 .or_else(|| self.llm_registry.default().ok())
8671 .ok_or_else(|| AgentError::Config("No LLM available for planning".into()))?;
8672
8673 let mut available_tool_ids: Vec<String> = self
8674 .get_available_tool_ids()
8675 .await
8676 .unwrap_or_else(|_| self.tools.list_ids());
8677 let mut available_skills: Vec<String> = self.skills.iter().map(|s| s.id.clone()).collect();
8678
8679 if let Some(config) = planning_config {
8681 if !config.available.tools.is_all() {
8682 available_tool_ids.retain(|t| config.available.tools.allows(t));
8683 }
8684 if !config.available.skills.is_all() {
8685 available_skills.retain(|s| config.available.skills.allows(s));
8686 }
8687 }
8688
8689 let tool_descriptions: Vec<String> = available_tool_ids
8692 .iter()
8693 .filter_map(|id| {
8694 self.tools.get(id).map(|tool| {
8695 let schema = tool.input_schema();
8696 let args_desc = schema
8697 .get("properties")
8698 .and_then(|p| serde_json::to_string(p).ok())
8699 .unwrap_or_else(|| "{}".to_string());
8700 format!(
8701 "- {} ({}): {}\n Arguments: {}",
8702 id,
8703 tool.name(),
8704 tool.description(),
8705 args_desc
8706 )
8707 })
8708 })
8709 .collect();
8710
8711 let tools_section = if tool_descriptions.is_empty() {
8712 "Available tools: none".to_string()
8713 } else {
8714 format!("Available tools:\n{}", tool_descriptions.join("\n"))
8715 };
8716
8717 let skills_section = if available_skills.is_empty() {
8718 "Available skills: none".to_string()
8719 } else {
8720 format!("Available skills: {}", available_skills.join(", "))
8721 };
8722
8723 let prompt = format!(
8724 r#"Create a step-by-step plan to accomplish this goal.
8725
8726Goal: "{}"
8727
8728{}
8729
8730{}
8731
8732Create a plan with clear steps. For each step, specify:
8733- description: What this step accomplishes
8734- action_type: "tool", "skill", "think", or "respond"
8735- action_target: The tool/skill id (if applicable)
8736- args: The arguments object matching the tool's schema (if action_type is "tool")
8737- dependencies: List of step IDs this depends on (empty if none)
8738
8739Respond in JSON format:
8740{{
8741 "steps": [
8742 {{"id": "step1", "description": "...", "action_type": "tool", "action_target": "tool_id", "args": {{"required_field": "value"}}, "dependencies": []}},
8743 {{"id": "step2", "description": "...", "action_type": "think", "action_target": "...", "dependencies": ["step1"]}}
8744 ]
8745}}"#,
8746 input, tools_section, skills_section,
8747 );
8748
8749 let messages = vec![ChatMessage::user(&prompt)];
8750 let response = self
8751 .observe_purpose(
8752 ObservationPurpose::PlanGeneration,
8753 planner_llm.complete(&messages, None),
8754 )
8755 .await
8756 .map_err(|e| AgentError::LLM(format!("Planning failed: {}", e)))?;
8757
8758 let mut plan = Plan::new(input);
8759
8760 if let Some(json_start) = response.content.find('{')
8761 && let Some(json_end) = response.content.rfind('}')
8762 {
8763 let json_str = &response.content[json_start..=json_end];
8764 if let Ok(parsed) = serde_json::from_str::<serde_json::Value>(json_str)
8765 && let Some(steps) = parsed.get("steps").and_then(|s| s.as_array())
8766 {
8767 for step_value in steps {
8768 let id = step_value
8769 .get("id")
8770 .and_then(|v| v.as_str())
8771 .unwrap_or("step");
8772 let desc = step_value
8773 .get("description")
8774 .and_then(|v| v.as_str())
8775 .unwrap_or("");
8776 let action_type = step_value
8777 .get("action_type")
8778 .and_then(|v| v.as_str())
8779 .unwrap_or("think");
8780 let action_target = step_value
8781 .get("action_target")
8782 .and_then(|v| v.as_str())
8783 .unwrap_or("");
8784 let args = step_value
8785 .get("args")
8786 .cloned()
8787 .unwrap_or(serde_json::json!({}));
8788 let deps: Vec<String> = step_value
8789 .get("dependencies")
8790 .and_then(|v| v.as_array())
8791 .map(|arr| {
8792 arr.iter()
8793 .filter_map(|v| v.as_str().map(String::from))
8794 .collect()
8795 })
8796 .unwrap_or_default();
8797
8798 let action = match action_type {
8799 "tool" => PlanAction::tool(action_target, args),
8800 "skill" => PlanAction::skill(action_target),
8801 "respond" => PlanAction::respond(action_target),
8802 _ => PlanAction::think(desc),
8803 };
8804
8805 let step = PlanStep::new(desc, action)
8806 .with_id(id)
8807 .with_dependencies(deps);
8808 plan.add_step(step);
8809 }
8810 }
8811 }
8812
8813 if plan.steps.is_empty() {
8814 plan.add_step(PlanStep::new(
8815 "Process the request",
8816 PlanAction::think(input),
8817 ));
8818 plan.add_step(PlanStep::new(
8819 "Provide response",
8820 PlanAction::respond("Answer based on analysis"),
8821 ));
8822 }
8823
8824 Ok(plan)
8825 }
8826
8827 async fn execute_plan(&self, plan: &mut Plan) -> Result<String> {
8828 let llm = self.get_state_llm()?;
8829 let mut results: HashMap<String, serde_json::Value> = HashMap::new();
8830 let effective = self.get_effective_reasoning_config();
8831 let max_steps = effective.get_planning().map(|c| c.max_steps).unwrap_or(10);
8832
8833 plan.status = PlanStatus::InProgress;
8834
8835 for step_idx in 0..plan.steps.len().min(max_steps as usize) {
8836 let step = &plan.steps[step_idx];
8837
8838 let deps_satisfied = step.dependencies.iter().all(|dep| {
8839 plan.steps
8840 .iter()
8841 .find(|s| &s.id == dep)
8842 .map(|s| s.status.is_completed())
8843 .unwrap_or(false)
8844 });
8845
8846 if !deps_satisfied {
8847 continue;
8848 }
8849
8850 plan.steps[step_idx].mark_running();
8851
8852 let result = match &plan.steps[step_idx].action {
8853 PlanAction::Tool { tool, args } => {
8854 let has_dep_results = plan.steps[step_idx]
8860 .dependencies
8861 .iter()
8862 .any(|dep| results.contains_key(dep));
8863
8864 let final_args = if has_dep_results {
8865 let dep_context: String = plan.steps[step_idx]
8866 .dependencies
8867 .iter()
8868 .filter_map(|dep| results.get(dep).map(|r| format!("{}: {}", dep, r)))
8869 .collect::<Vec<_>>()
8870 .join("\n");
8871
8872 let tool_schema = self
8873 .tools
8874 .get(tool)
8875 .map(|t| {
8876 let schema = t.input_schema();
8877 let props = schema
8878 .get("properties")
8879 .and_then(|p| serde_json::to_string(p).ok())
8880 .unwrap_or_else(|| "{}".to_string());
8881 format!(
8882 "{}: {}\nArguments schema: {}",
8883 t.id(),
8884 t.description(),
8885 props
8886 )
8887 })
8888 .unwrap_or_default();
8889
8890 let step_desc = &plan.steps[step_idx].description;
8891 let arg_prompt = format!(
8892 "Generate the JSON arguments for a tool call.\n\n\
8893 Tool: {}\n\n\
8894 Task: {}\n\n\
8895 Previous step results:\n{}\n\n\
8896 Planner's draft arguments: {}\n\n\
8897 Produce ONLY a valid JSON object with the correct argument values.\n\
8898 Use actual values from the previous step results, not template references.",
8899 tool_schema,
8900 step_desc,
8901 dep_context,
8902 serde_json::to_string(args).unwrap_or_default()
8903 );
8904 let messages = vec![ChatMessage::user(&arg_prompt)];
8905 match self
8906 .observe_purpose(
8907 ObservationPurpose::PlanStep,
8908 llm.complete(&messages, None),
8909 )
8910 .await
8911 {
8912 Ok(resp) => {
8913 let content = resp.content.trim();
8914 let json_start = content.find('{');
8916 let json_end = content.rfind('}');
8917 if let (Some(start), Some(end)) = (json_start, json_end) {
8918 serde_json::from_str(&content[start..=end])
8919 .unwrap_or_else(|_| args.clone())
8920 } else {
8921 args.clone()
8922 }
8923 }
8924 Err(_) => args.clone(),
8925 }
8926 } else {
8927 args.clone()
8928 };
8929
8930 let request = ToolExecutionRequest::new(
8931 uuid::Uuid::new_v4().to_string(),
8932 tool.clone(),
8933 final_args,
8934 ToolCallSource::Plan {
8935 step_index: step_idx,
8936 },
8937 );
8938 match self.execute_tool_record(request).await {
8939 Ok(record) if record.success => {
8940 serde_json::json!({ "output": record.model_output_string() })
8941 }
8942 Ok(record) => {
8943 plan.steps[step_idx].mark_failed(record.model_output_string());
8944 continue;
8945 }
8946 Err(e) => {
8947 plan.steps[step_idx].mark_failed(e.to_string());
8948 continue;
8949 }
8950 }
8951 }
8952 PlanAction::Skill { skill } => {
8953 if let Some(skill_def) = self.skills.iter().find(|s| &s.id == skill) {
8954 if let Some(ref executor) = self.skill_executor {
8955 match executor
8956 .execute_with_invoker(skill_def, "", serde_json::json!({}), self)
8957 .await
8958 {
8959 Ok(output) => serde_json::json!({ "output": output }),
8960 Err(e) => {
8961 plan.steps[step_idx].mark_failed(e.to_string());
8962 continue;
8963 }
8964 }
8965 } else {
8966 serde_json::json!({ "output": "Skill executor not available" })
8967 }
8968 } else {
8969 plan.steps[step_idx].mark_failed("Skill not found");
8970 continue;
8971 }
8972 }
8973 PlanAction::Think { prompt } => {
8974 let context: String = results
8975 .iter()
8976 .map(|(k, v)| format!("{}: {}", k, v))
8977 .collect::<Vec<_>>()
8978 .join("\n");
8979
8980 let think_prompt = format!("Context:\n{}\n\nTask: {}", context, prompt);
8981 let messages = vec![ChatMessage::user(&think_prompt)];
8982
8983 match self
8984 .observe_purpose(
8985 ObservationPurpose::PlanStep,
8986 llm.complete(&messages, None),
8987 )
8988 .await
8989 {
8990 Ok(resp) => serde_json::json!({ "output": resp.content }),
8991 Err(e) => {
8992 plan.steps[step_idx].mark_failed(e.to_string());
8993 continue;
8994 }
8995 }
8996 }
8997 PlanAction::Respond { template } => {
8998 let context: String = results
8999 .iter()
9000 .map(|(k, v)| format!("{}: {}", k, v))
9001 .collect::<Vec<_>>()
9002 .join("\n");
9003
9004 let respond_prompt = format!(
9005 "Based on this context:\n{}\n\nGenerate a response following this template/instruction: {}",
9006 context, template
9007 );
9008 let messages = vec![ChatMessage::user(&respond_prompt)];
9009
9010 match self
9011 .observe_purpose(
9012 ObservationPurpose::PlanStep,
9013 llm.complete(&messages, None),
9014 )
9015 .await
9016 {
9017 Ok(resp) => serde_json::json!({ "output": resp.content }),
9018 Err(e) => {
9019 plan.steps[step_idx].mark_failed(e.to_string());
9020 continue;
9021 }
9022 }
9023 }
9024 };
9025
9026 results.insert(plan.steps[step_idx].id.clone(), result.clone());
9027 plan.steps[step_idx].mark_completed(Some(result));
9028 }
9029
9030 let has_failures = plan.steps.iter().any(|s| s.status.is_failed());
9032 if has_failures {
9033 let failed_ids: Vec<String> = plan
9034 .steps
9035 .iter()
9036 .filter(|s| s.status.is_failed())
9037 .map(|s| s.id.clone())
9038 .collect();
9039 plan.status = PlanStatus::Failed {
9040 error: format!("Steps failed: {}", failed_ids.join(", ")),
9041 };
9042 } else {
9043 plan.status = PlanStatus::Completed;
9044 }
9045
9046 let all_outputs: Vec<String> = plan
9048 .steps
9049 .iter()
9050 .filter(|s| s.status.is_completed())
9051 .filter_map(|s| {
9052 s.result
9053 .as_ref()
9054 .and_then(|r| r.get("output"))
9055 .and_then(|o| o.as_str())
9056 .map(|o| format!("{}: {}", s.description, o))
9057 })
9058 .collect();
9059
9060 if all_outputs.is_empty() {
9061 return Ok("Plan execution completed but produced no results.".to_string());
9062 }
9063
9064 if all_outputs.len() == 1 {
9065 return Ok(all_outputs.into_iter().next().unwrap());
9066 }
9067
9068 let context = all_outputs.join("\n\n");
9070 let prompt = format!(
9071 "You completed a multi-step plan for: \"{}\"\n\nStep results:\n{}\n\nProvide a coherent final response that synthesizes these results.",
9072 plan.goal, context
9073 );
9074 let messages = vec![ChatMessage::user(&prompt)];
9075 match self
9076 .observe_purpose(ObservationPurpose::PlanStep, llm.complete(&messages, None))
9077 .await
9078 {
9079 Ok(resp) => Ok(resp.content.trim().to_string()),
9080 Err(_) => Ok(context),
9081 }
9082 }
9083
9084 async fn evaluate_response(&self, input: &str, response: &str) -> Result<EvaluationResult> {
9085 let effective_config = self.get_effective_reflection_config();
9086 self.evaluate_response_with_config(input, response, &effective_config)
9087 .await
9088 }
9089
9090 fn extract_thinking(&self, content: &str) -> (Option<String>, String) {
9091 if let Some(start) = content.find("<thinking>")
9092 && let Some(end) = content.find("</thinking>")
9093 {
9094 let thinking = content[start + 10..end].trim().to_string();
9095 let answer = content[end + 11..].trim().to_string();
9096 return (Some(thinking), answer);
9097 }
9098 (None, content.to_string())
9099 }
9100
9101 fn format_response_with_thinking(&self, thinking: Option<&str>, answer: &str) -> String {
9102 match self.get_effective_reasoning_config().output {
9103 ReasoningOutput::Hidden => answer.to_string(),
9104 ReasoningOutput::Visible => {
9105 if let Some(t) = thinking {
9106 format!("Thinking:\n{}\n\nAnswer:\n{}", t, answer)
9107 } else {
9108 answer.to_string()
9109 }
9110 }
9111 ReasoningOutput::Tagged => {
9112 if let Some(t) = thinking {
9113 format!("<thinking>{}</thinking>\n{}", t, answer)
9114 } else {
9115 answer.to_string()
9116 }
9117 }
9118 }
9119 }
9120
9121 fn disambiguation_question_response(
9124 question: &ClarificationQuestion,
9125 detection: &AmbiguityDetectionResult,
9126 awaiting_confirmation: bool,
9127 ) -> AgentResponse {
9128 let status = if awaiting_confirmation {
9129 "awaiting_confirmation"
9130 } else {
9131 "awaiting_clarification"
9132 };
9133 AgentResponse::new(&question.question).with_metadata(
9134 "disambiguation",
9135 serde_json::json!({
9136 "status": status,
9137 "options": question.options,
9138 "clarifying": question.clarifying,
9139 "detection": {
9140 "type": detection.ambiguity_type,
9141 "confidence": detection.confidence,
9142 "what_is_unclear": detection.what_is_unclear,
9143 }
9144 }),
9145 )
9146 }
9147
9148 async fn resolve_disambiguation(&self, input: &str) -> Result<DisambiguationDispatch> {
9160 let Some(ref disambiguator) = self.disambiguation_manager else {
9161 return Ok(DisambiguationDispatch::Proceed(input.to_string()));
9162 };
9163 let disambiguation_context = self.build_disambiguation_context().await?;
9164
9165 let state_override = self
9167 .state_machine
9168 .as_ref()
9169 .and_then(|sm| sm.current_definition())
9170 .and_then(|def| def.disambiguation.clone());
9171
9172 let state_generation = self
9173 .state_machine
9174 .as_ref()
9175 .map(|state_machine| state_machine.generation());
9176 let disambiguation_epoch = self.disambiguation_epoch.load(Ordering::SeqCst);
9177 let mut disambiguation_result = self
9178 .observe_purpose(
9179 ObservationPurpose::DisambiguationDetection,
9180 disambiguator.process_input_with_override(
9181 input,
9182 &disambiguation_context,
9183 state_override.as_ref(),
9184 None,
9185 ),
9186 )
9187 .await?;
9188 let current_state_generation = self
9189 .state_machine
9190 .as_ref()
9191 .map(|state_machine| state_machine.generation());
9192 if current_state_generation != state_generation
9193 || self.disambiguation_epoch.load(Ordering::SeqCst) != disambiguation_epoch
9194 {
9195 disambiguator.clear_pending().await;
9196 *self.pending_skill_id.write() = None;
9197 disambiguation_result = DisambiguationResult::Abandoned { new_input: None };
9198 info!(
9199 confirmation_event = "invalidated",
9200 invalidation_reason = "state_generation_changed",
9201 "Disambiguation result invalidated before redispatch"
9202 );
9203 }
9204 match disambiguation_result {
9205 DisambiguationResult::Clear => {
9206 debug!("Input is clear, proceeding normally");
9207 Ok(DisambiguationDispatch::Proceed(input.to_string()))
9208 }
9209 DisambiguationResult::NeedsClarification {
9210 question,
9211 detection,
9212 } => {
9213 let admission = match self
9214 .admit_disambiguation_redispatch(disambiguation_epoch, state_generation)
9215 .await
9216 {
9217 Ok(admission) => admission,
9218 Err(error) => {
9219 *self.pending_skill_id.write() = None;
9220 return Err(error);
9221 }
9222 };
9223 let awaiting_confirmation = disambiguator.has_pending_confirmation().await;
9224 info!(
9225 ambiguity_type = ?detection.ambiguity_type,
9226 confidence = detection.confidence,
9227 "Input requires clarification"
9228 );
9229
9230 self.commit_root_user_message(input).await?;
9233 self.memory
9234 .add_message(ChatMessage::assistant(&question.question))
9235 .await?;
9236
9237 let response = Self::disambiguation_question_response(
9238 &question,
9239 &detection,
9240 awaiting_confirmation,
9241 );
9242 drop(admission);
9243 self.finish_turn_if_root(&response).await?;
9244 Ok(DisambiguationDispatch::Terminal(response))
9245 }
9246 DisambiguationResult::Clarified {
9247 enriched_input,
9248 resolved,
9249 ..
9250 } => {
9251 let admission = match self
9252 .admit_disambiguation_redispatch(disambiguation_epoch, state_generation)
9253 .await
9254 {
9255 Ok(admission) => admission,
9256 Err(error) => {
9257 *self.pending_skill_id.write() = None;
9258 return Err(error);
9259 }
9260 };
9261 info!(
9262 resolved_count = resolved.len(),
9263 enriched = %enriched_input,
9264 "Input clarified, injecting resolved intent into context"
9265 );
9266
9267 for (key, value) in &resolved {
9270 let context_key = format!("disambiguation.{}", key);
9271 let _ = self.context_manager.set(&context_key, value.clone());
9272 }
9273
9274 if let Some(intent) = resolved.get("intent") {
9275 let _ = self.context_manager.set("resolved_intent", intent.clone());
9276 }
9277
9278 let _ = self
9279 .context_manager
9280 .set("disambiguation.resolved", serde_json::Value::Bool(true));
9281
9282 let skill_id = self.pending_skill_id.read().clone();
9286 drop(admission);
9287 if let Some(skill_id) = skill_id {
9288 info!(skill_id = %skill_id, "Re-checking skill disambiguation on clarified input");
9289 return Ok(DisambiguationDispatch::RecheckSkill {
9290 skill_id,
9291 enriched_input,
9292 disambiguation_epoch,
9293 state_generation,
9294 });
9295 }
9296 Ok(DisambiguationDispatch::Proceed(enriched_input))
9297 }
9298 DisambiguationResult::ProceedWithBestGuess { enriched_input } => {
9299 info!("Proceeding with best guess interpretation");
9300
9301 let skill_id = self.pending_skill_id.read().clone();
9303 if let Some(skill_id) = skill_id {
9304 info!(skill_id = %skill_id, "Re-checking skill disambiguation on best-guess input");
9305 return Ok(DisambiguationDispatch::RecheckSkill {
9306 skill_id,
9307 enriched_input,
9308 disambiguation_epoch,
9309 state_generation,
9310 });
9311 }
9312 Ok(DisambiguationDispatch::Proceed(enriched_input))
9313 }
9314 DisambiguationResult::GiveUp { reason } => {
9315 *self.pending_skill_id.write() = None;
9316 warn!(reason = %reason, "Disambiguation gave up");
9317 let apology = self
9318 .generate_localized_apology(
9319 "Generate a brief, polite apology saying you couldn't understand the request. Be concise.",
9320 &reason,
9321 )
9322 .await
9323 .unwrap_or_else(|_| {
9324 format!("I'm sorry, I couldn't understand your request: {}", reason)
9325 });
9326 let response = AgentResponse::new(&apology);
9327 self.finish_turn_if_root(&response).await?;
9328 Ok(DisambiguationDispatch::Terminal(response))
9329 }
9330 DisambiguationResult::Escalate { reason } => {
9331 *self.pending_skill_id.write() = None;
9332 info!(reason = %reason, "Escalating to human");
9333 if let Some(ref hitl) = self.hitl_engine {
9334 let trigger =
9335 ApprovalTrigger::condition("disambiguation_escalation", reason.clone());
9336 let mut context_map = HashMap::new();
9337 context_map.insert("original_input".to_string(), serde_json::json!(input));
9338 context_map.insert("reason".to_string(), serde_json::json!(&reason));
9339 let check_result = HITLCheckResult::required(
9340 trigger,
9341 context_map,
9342 format!("User request needs human assistance: {}", reason),
9343 Some(hitl.config().default_timeout_seconds),
9344 );
9345 let result = self.request_hitl_approval(check_result).await?;
9346 if matches!(
9347 result,
9348 ApprovalResult::Approved | ApprovalResult::Modified { .. }
9349 ) {
9350 return Ok(DisambiguationDispatch::Proceed(input.to_string()));
9352 }
9353 }
9354 let apology = self
9355 .generate_localized_apology(
9356 "Explain briefly that you're transferring the user to a human agent for help.",
9357 &reason,
9358 )
9359 .await
9360 .unwrap_or_else(|_| {
9361 format!("I need human assistance to help with your request: {}", reason)
9362 });
9363 let response = AgentResponse::new(&apology);
9364 self.finish_turn_if_root(&response).await?;
9365 Ok(DisambiguationDispatch::Terminal(response))
9366 }
9367 DisambiguationResult::Abandoned { new_input } => {
9368 *self.pending_skill_id.write() = None;
9369
9370 info!(
9371 has_new_input = new_input.is_some(),
9372 "Clarification abandoned by user"
9373 );
9374
9375 self.commit_root_user_message(input).await?;
9376
9377 match new_input {
9378 Some(fresh_input) => {
9379 Ok(DisambiguationDispatch::Proceed(fresh_input))
9382 }
9383 None => {
9384 let ack = self
9386 .generate_localized_apology(
9387 "The user changed their mind about their previous request. \
9388 Generate a brief, friendly acknowledgment (e.g. 'OK, no problem. What else can I help with?'). \
9389 Do NOT apologize excessively. Be concise.",
9390 "User abandoned clarification",
9391 )
9392 .await
9393 .unwrap_or_else(|_| {
9394 "OK, no problem. What else can I help with?".to_string()
9395 });
9396
9397 self.memory
9398 .add_message(ChatMessage::assistant(&ack))
9399 .await?;
9400
9401 let response = AgentResponse::new(&ack);
9402 self.finish_turn_if_root(&response).await?;
9403 Ok(DisambiguationDispatch::Terminal(response))
9404 }
9405 }
9406 }
9407 }
9408 }
9409
9410 async fn run_loop(&self, input: &str) -> Result<AgentResponse> {
9414 self.init_storage().await?;
9418 self.begin_root_turn();
9419 let _root_cleanup = RootTurnCleanup::new(self);
9420 info!(input_len = input.len(), "Starting chat");
9421
9422 self.hooks.on_message_received(input).await;
9423
9424 if !self.context_initialized.swap(true, Ordering::SeqCst) {
9428 self.context_manager.initialize().await?;
9429 debug!("Context manager initialized (defaults, env, builtins)");
9430 }
9431
9432 self.check_turn_timeout().await?;
9433 self.context_manager.refresh_per_turn().await?;
9434
9435 self.clear_disambiguation_context();
9438
9439 let input_to_run = match self.resolve_disambiguation(input).await? {
9442 DisambiguationDispatch::Terminal(response) => return Ok(response),
9443 DisambiguationDispatch::RecheckSkill {
9444 skill_id,
9445 enriched_input,
9446 disambiguation_epoch,
9447 state_generation,
9448 } => {
9449 return self
9450 .recheck_skill_disambiguation(
9451 &skill_id,
9452 &enriched_input,
9453 disambiguation_epoch,
9454 state_generation,
9455 )
9456 .await;
9457 }
9458 DisambiguationDispatch::Proceed(input) => input,
9459 };
9460
9461 self.run_loop_internal(&input_to_run).await
9462 }
9463
9464 async fn generate_localized_apology(&self, instruction: &str, reason: &str) -> Result<String> {
9466 let llm = self.llm_registry.router().map_err(|e| {
9467 AgentError::LLM(format!(
9468 "Router LLM not available for localized response: {}",
9469 e
9470 ))
9471 })?;
9472
9473 let recent: Vec<String> = self
9474 .memory
9475 .get_messages(Some(3))
9476 .await?
9477 .iter()
9478 .map(|m| m.content.clone())
9479 .collect();
9480
9481 let context_hint = if recent.is_empty() {
9482 String::new()
9483 } else {
9484 format!(
9485 "\nRecent conversation (detect the user's language from this):\n{}\n",
9486 recent.join("\n")
9487 )
9488 };
9489
9490 let prompt = format!(
9491 "{}\nReason: {}\n{}Respond in the same language as the user. Output ONLY the message, nothing else.",
9492 instruction, reason, context_hint
9493 );
9494
9495 let messages = vec![ChatMessage::user(&prompt)];
9496 let response = self
9497 .observe_purpose(
9498 ObservationPurpose::DisambiguationClarification,
9499 llm.complete(&messages, None),
9500 )
9501 .await
9502 .map_err(|e| AgentError::LLM(format!("Localized response generation failed: {}", e)))?;
9503
9504 Ok(response.content.trim().to_string())
9505 }
9506
9507 fn render_action_args(&self, args: &Value) -> Value {
9511 let context = self.build_context_with_overlays();
9512 match args {
9513 Value::Object(map) => {
9514 let mut rendered = serde_json::Map::new();
9515 for (k, v) in map {
9516 match v {
9517 Value::String(s) if s.contains("{{") => {
9518 match self.template_renderer.render(s, &context) {
9519 Ok(rendered_str) => {
9520 rendered.insert(k.clone(), Value::String(rendered_str));
9521 }
9522 Err(_) => {
9523 rendered.insert(k.clone(), v.clone());
9524 }
9525 }
9526 }
9527 _ => {
9528 rendered.insert(k.clone(), v.clone());
9529 }
9530 }
9531 }
9532 Value::Object(rendered)
9533 }
9534 _ => args.clone(),
9535 }
9536 }
9537
9538 fn clear_disambiguation_context(&self) {
9540 let _ = self
9541 .context_manager
9542 .set("resolved_intent", serde_json::Value::Null);
9543
9544 let all = self.context_manager.get_all();
9545 for key in all.keys() {
9546 if key.starts_with("disambiguation.") {
9547 let _ = self.context_manager.set(key, serde_json::Value::Null);
9548 }
9549 }
9550 }
9551
9552 async fn recheck_skill_disambiguation(
9558 &self,
9559 skill_id: &str,
9560 enriched_input: &str,
9561 expected_disambiguation_epoch: u64,
9562 expected_state_generation: Option<u64>,
9563 ) -> Result<AgentResponse> {
9564 let skill = self
9565 .skill_router
9566 .as_ref()
9567 .and_then(|r| r.get_skill(skill_id).cloned());
9568
9569 if let Some(ref skill) = skill
9571 && let Some(ref skill_disambig) = skill.disambiguation
9572 && skill_disambig.enabled.unwrap_or(false)
9573 && let Some(ref disambiguator) = self.disambiguation_manager
9574 {
9575 let context = self.build_disambiguation_context().await?;
9576 let state_override = self
9577 .state_machine
9578 .as_ref()
9579 .and_then(|sm| sm.current_definition())
9580 .and_then(|def| def.disambiguation.clone());
9581
9582 let disambiguation_result = self
9583 .observe_purpose(
9584 ObservationPurpose::DisambiguationDetection,
9585 disambiguator.process_input_with_override(
9586 enriched_input,
9587 &context,
9588 state_override.as_ref(),
9589 Some(skill_disambig),
9590 ),
9591 )
9592 .await?;
9593 let current_state_generation = self
9594 .state_machine
9595 .as_ref()
9596 .map(|state_machine| state_machine.generation());
9597 if current_state_generation != expected_state_generation
9598 || self.disambiguation_epoch.load(Ordering::SeqCst) != expected_disambiguation_epoch
9599 {
9600 disambiguator.clear_pending().await;
9601 *self.pending_skill_id.write() = None;
9602 return Err(AgentError::Other(
9603 "State or reset ownership changed during skill disambiguation recheck"
9604 .to_string(),
9605 ));
9606 }
9607 match disambiguation_result {
9608 DisambiguationResult::Clear => {
9609 debug!(skill_id = %skill_id, "Skill re-check: all fields present");
9610 }
9611 DisambiguationResult::NeedsClarification {
9612 question,
9613 detection,
9614 } => {
9615 let admission = self
9616 .admit_disambiguation_redispatch(
9617 expected_disambiguation_epoch,
9618 expected_state_generation,
9619 )
9620 .await?;
9621 let awaiting_confirmation = disambiguator.has_pending_confirmation().await;
9622 info!(
9623 skill_id = %skill_id,
9624 ambiguity_type = ?detection.ambiguity_type,
9625 what_is_unclear = ?detection.what_is_unclear,
9626 "Skill re-check: still missing fields, asking again"
9627 );
9628 self.memory
9632 .add_message(ChatMessage::user(enriched_input))
9633 .await?;
9634 self.memory
9635 .add_message(ChatMessage::assistant(&question.question))
9636 .await?;
9637
9638 let response = AgentResponse::new(&question.question).with_metadata(
9639 "disambiguation",
9640 serde_json::json!({
9641 "status": if awaiting_confirmation { "awaiting_confirmation" } else { "awaiting_clarification" },
9642 "skill_id": skill_id,
9643 "options": question.options,
9644 "clarifying": question.clarifying,
9645 "detection": {
9646 "type": detection.ambiguity_type,
9647 "confidence": detection.confidence,
9648 "what_is_unclear": detection.what_is_unclear,
9649 }
9650 }),
9651 );
9652 drop(admission);
9653 self.finish_turn_if_root(&response).await?;
9654 return Ok(response);
9655 }
9656 DisambiguationResult::Clarified {
9657 enriched_input: re_enriched,
9658 ..
9659 } => {
9660 debug!(skill_id = %skill_id, "Skill re-check: clarified immediately, executing");
9661 let admission = self
9662 .admit_disambiguation_redispatch(
9663 expected_disambiguation_epoch,
9664 expected_state_generation,
9665 )
9666 .await?;
9667 *self.pending_skill_id.write() = None;
9668 drop(admission);
9669 let skill_response = self.execute_skill_by_id(skill_id, &re_enriched).await?;
9670 self.memory
9671 .add_message(ChatMessage::user(&re_enriched))
9672 .await?;
9673 return self
9674 .handle_skill_response(
9675 &re_enriched,
9676 skill_id,
9677 skill_response,
9678 &HashMap::new(),
9679 )
9680 .await;
9681 }
9682 DisambiguationResult::ProceedWithBestGuess {
9683 enriched_input: re_enriched,
9684 } => {
9685 debug!(skill_id = %skill_id, "Skill re-check: proceeding with best guess");
9686 let admission = self
9687 .admit_disambiguation_redispatch(
9688 expected_disambiguation_epoch,
9689 expected_state_generation,
9690 )
9691 .await?;
9692 *self.pending_skill_id.write() = None;
9693 drop(admission);
9694 let skill_response = self.execute_skill_by_id(skill_id, &re_enriched).await?;
9695 self.memory
9696 .add_message(ChatMessage::user(&re_enriched))
9697 .await?;
9698 return self
9699 .handle_skill_response(
9700 &re_enriched,
9701 skill_id,
9702 skill_response,
9703 &HashMap::new(),
9704 )
9705 .await;
9706 }
9707 DisambiguationResult::GiveUp { reason } => {
9708 *self.pending_skill_id.write() = None;
9709 let apology = self
9710 .generate_localized_apology(
9711 "Generate a brief, polite apology saying you couldn't understand the request. Be concise.",
9712 &reason,
9713 )
9714 .await
9715 .unwrap_or_else(|_| {
9716 format!("I'm sorry, I couldn't understand your request: {}", reason)
9717 });
9718 let response = AgentResponse::new(&apology);
9719 self.finish_turn_if_root(&response).await?;
9720 return Ok(response);
9721 }
9722 DisambiguationResult::Escalate { reason } => {
9723 *self.pending_skill_id.write() = None;
9724 let apology = self
9725 .generate_localized_apology(
9726 "Explain briefly that you're transferring the user to a human agent for help.",
9727 &reason,
9728 )
9729 .await
9730 .unwrap_or_else(|_| {
9731 format!("I need human assistance to help with your request: {}", reason)
9732 });
9733 let response = AgentResponse::new(&apology);
9734 self.finish_turn_if_root(&response).await?;
9735 return Ok(response);
9736 }
9737 DisambiguationResult::Abandoned { new_input } => {
9738 *self.pending_skill_id.write() = None;
9741 debug!(skill_id = %skill_id, "Skill re-check: abandoned by user");
9742 if let Some(fresh) = new_input {
9743 return self.run_loop_internal(&fresh).await;
9744 }
9745 let ack = self
9746 .generate_localized_apology(
9747 "The user changed their mind about their previous request. \
9748 Generate a brief, friendly acknowledgment (e.g. 'OK, no problem. What else can I help with?'). \
9749 Do NOT apologize excessively. Be concise.",
9750 "User abandoned clarification",
9751 )
9752 .await
9753 .unwrap_or_else(|_| {
9754 "OK, no problem. What else can I help with?".to_string()
9755 });
9756 self.memory
9757 .add_message(ChatMessage::assistant(&ack))
9758 .await?;
9759 let response = AgentResponse::new(&ack);
9760 self.finish_turn_if_root(&response).await?;
9761 return Ok(response);
9762 }
9763 }
9764 }
9765
9766 let admission = self
9768 .admit_disambiguation_redispatch(
9769 expected_disambiguation_epoch,
9770 expected_state_generation,
9771 )
9772 .await?;
9773 *self.pending_skill_id.write() = None;
9774 drop(admission);
9775 let skill_response = self.execute_skill_by_id(skill_id, enriched_input).await?;
9776 self.memory
9777 .add_message(ChatMessage::user(enriched_input))
9778 .await?;
9779 self.handle_skill_response(enriched_input, skill_id, skill_response, &HashMap::new())
9780 .await
9781 }
9782
9783 async fn handle_skill_response(
9786 &self,
9787 processed_input: &str,
9788 skill_id: &str,
9789 skill_response: String,
9790 input_context: &HashMap<String, Value>,
9791 ) -> Result<AgentResponse> {
9792 let output_data = self.process_output(&skill_response, input_context).await?;
9793 let final_response = output_data.content;
9794
9795 self.memory
9796 .add_message(ChatMessage::assistant(&final_response))
9797 .await?;
9798
9799 self.check_memory_compression().await?;
9800
9801 self.increment_turn();
9802 self.evaluate_transitions(processed_input, &final_response)
9803 .await?;
9804
9805 let response = AgentResponse::new(final_response)
9806 .with_metadata("skill_id", serde_json::json!(skill_id));
9807 self.finish_turn_if_root(&response).await?;
9808 Ok(response)
9809 }
9810
9811 async fn handle_plan_and_execute(
9814 &self,
9815 processed_input: &str,
9816 input_context: &HashMap<String, Value>,
9817 auto_detected: bool,
9818 ) -> Result<AgentResponse> {
9819 let effective = self.get_effective_reasoning_config();
9820 let plan_reflection = effective
9821 .get_planning()
9822 .map(|c| c.reflection.clone())
9823 .unwrap_or_default();
9824
9825 let max_attempts = if plan_reflection.enabled {
9826 1 + plan_reflection.max_replans
9827 } else {
9828 1
9829 };
9830
9831 let mut plan = self.generate_plan(processed_input).await?;
9832 info!(
9833 plan_id = %plan.id,
9834 steps = plan.steps.len(),
9835 "Plan generated"
9836 );
9837
9838 let mut plan_result = String::new();
9839
9840 for attempt in 0..max_attempts {
9841 *self.current_plan.write() = Some(plan.clone());
9842 plan_result = self.execute_plan(&mut plan).await?;
9843
9844 info!(
9845 plan_status = ?plan.status,
9846 completed_steps = plan.completed_steps().count(),
9847 attempt = attempt + 1,
9848 "Plan execution completed"
9849 );
9850
9851 if !plan_reflection.enabled {
9852 break;
9853 }
9854
9855 let has_failures = plan.steps.iter().any(|s| s.status.is_failed());
9856 if !has_failures {
9857 break;
9858 }
9859
9860 if attempt + 1 >= max_attempts {
9861 break;
9862 }
9863
9864 match plan_reflection.on_step_failure {
9865 StepFailureAction::Replan => {
9866 info!(attempt = attempt + 1, "Plan had failures, replanning");
9867 plan = self.generate_plan(processed_input).await?;
9868 }
9869 StepFailureAction::Abort => {
9870 warn!("Plan step failed, aborting");
9871 break;
9872 }
9873 StepFailureAction::Skip | StepFailureAction::Continue => {
9874 break;
9875 }
9876 }
9877 }
9878
9879 *self.current_plan.write() = Some(plan);
9880
9881 let output_data = self.process_output(&plan_result, input_context).await?;
9882 let final_content = output_data.content;
9883
9884 self.memory
9885 .add_message(ChatMessage::assistant(&final_content))
9886 .await?;
9887
9888 self.check_memory_compression().await?;
9889 self.increment_turn();
9890 self.evaluate_transitions(processed_input, &final_content)
9891 .await?;
9892
9893 let reasoning_metadata =
9894 ReasoningMetadata::new(ReasoningMode::PlanAndExecute).with_auto_detected(auto_detected);
9895
9896 let response = AgentResponse::new(&final_content).with_metadata(
9897 "reasoning",
9898 serde_json::to_value(&reasoning_metadata).unwrap_or_default(),
9899 );
9900
9901 self.finish_turn_if_root(&response).await?;
9902 Ok(response)
9903 }
9904
9905 fn inject_reasoning_prompt(
9907 &self,
9908 messages: &mut [ChatMessage],
9909 reasoning_mode: &ReasoningMode,
9910 is_first_iteration: bool,
9911 ) {
9912 if !is_first_iteration {
9913 return;
9914 }
9915 match reasoning_mode {
9916 ReasoningMode::CoT => {
9917 if let Some(msg) = messages.first_mut()
9918 && matches!(msg.role, ai_agents_core::Role::System)
9919 {
9920 msg.content = self.build_cot_system_prompt(&msg.content);
9921 debug!("Applied Chain-of-Thought system prompt");
9922 }
9923 }
9924 ReasoningMode::React => {
9925 if let Some(msg) = messages.first_mut()
9926 && matches!(msg.role, ai_agents_core::Role::System)
9927 {
9928 msg.content = self.build_react_system_prompt(&msg.content);
9929 debug!("Applied ReAct system prompt");
9930 }
9931 }
9932 _ => {}
9933 }
9934 }
9935
9936 async fn generate_main_response_draft(
9941 &self,
9942 processed_input: &str,
9943 reasoning_mode: &ReasoningMode,
9944 ) -> Result<MainResponseDraft> {
9945 let llm = self.get_state_llm()?;
9946 let protocol = self.main_tool_protocol(llm.as_ref(), true).await?;
9947 let mut messages = self
9948 .build_messages_internal(false, Some(processed_input), protocol.choice.is_none())
9949 .await?;
9950 self.inject_reasoning_prompt(&mut messages, reasoning_mode, true);
9951 let response = self
9952 .complete_main_llm_with_recovery(llm, &messages, &protocol)
9953 .await?;
9954 let content = response.content.trim().to_string();
9955 let (thinking, answer) = self.extract_thinking(&content);
9956 if let Some(calls) = self.parse_main_tool_calls(&content, &protocol)? {
9957 return Ok(MainResponseDraft::ToolCalls {
9958 raw_content: content,
9959 calls,
9960 thinking,
9961 });
9962 }
9963 Ok(MainResponseDraft::Text {
9964 raw_content: answer,
9965 thinking,
9966 })
9967 }
9968
9969 async fn commit_main_response_draft(
9974 &self,
9975 processed_input: &str,
9976 input_context: &HashMap<String, Value>,
9977 draft: MainResponseDraft,
9978 reasoning_mode: ReasoningMode,
9979 auto_detected: bool,
9980 ) -> Result<AgentResponse> {
9981 self.commit_root_user_message(processed_input).await?;
9982 match draft {
9983 MainResponseDraft::Text {
9984 raw_content,
9985 thinking,
9986 } => {
9987 self.finish_text_response_from_model(CommittedTextResponse {
9988 processed_input,
9989 input_context,
9990 answer: raw_content,
9991 reasoning_mode,
9992 auto_detected,
9993 iterations: 1,
9994 thinking_content: thinking,
9995 all_tool_calls: Vec::new(),
9996 })
9997 .await
9998 }
9999 MainResponseDraft::ToolCalls {
10000 raw_content,
10001 calls,
10002 thinking: _,
10003 } => {
10004 let mut all_tool_calls = Vec::new();
10005 match self
10006 .handle_tool_calls(
10007 processed_input,
10008 &raw_content,
10009 calls,
10010 &mut all_tool_calls,
10011 None,
10012 )
10013 .await?
10014 {
10015 ToolCallOutcome::Rejected(response) => {
10016 self.finish_turn_if_root(&response).await?;
10017 Ok(response)
10018 }
10019 ToolCallOutcome::Continue | ToolCallOutcome::TransitionFired => {
10020 self.continue_after_committed_tool_draft(processed_input)
10021 .await
10022 }
10023 }
10024 }
10025 }
10026 }
10027
10028 async fn continue_after_committed_tool_draft(
10033 &self,
10034 processed_input: &str,
10035 ) -> Result<AgentResponse> {
10036 *self.redispatch_depth.write() += 1;
10037 if let Some(context) = self.active_turn_context.write().as_mut() {
10038 context.enter_redispatch();
10039 }
10040 let result = Box::pin(self.run_loop_internal(processed_input)).await;
10041 *self.redispatch_depth.write() -= 1;
10042 if let Some(context) = self.active_turn_context.write().as_mut() {
10043 context.exit_redispatch();
10044 }
10045 let response = result?;
10046 self.finish_turn_if_root(&response).await?;
10047 Ok(response)
10048 }
10049
10050 async fn finish_text_response_from_model(
10055 &self,
10056 response: CommittedTextResponse<'_>,
10057 ) -> Result<AgentResponse> {
10058 let CommittedTextResponse {
10059 processed_input,
10060 input_context,
10061 answer,
10062 reasoning_mode,
10063 auto_detected,
10064 iterations,
10065 thinking_content,
10066 all_tool_calls,
10067 } = response;
10068 let output_data = self.process_output(&answer, input_context).await?;
10069 let mut final_content = if output_data.metadata.rejected {
10070 output_data
10071 .metadata
10072 .rejection_reason
10073 .unwrap_or_else(|| answer.to_string())
10074 } else {
10075 output_data.content
10076 };
10077 let llm = self.get_state_llm()?;
10078 let reflection_metadata;
10079 (final_content, reflection_metadata) = self
10080 .run_reflection(&*llm, processed_input, final_content)
10081 .await?;
10082 final_content =
10083 self.format_response_with_thinking(thinking_content.as_deref(), &final_content);
10084 let final_content = {
10085 let result = self
10086 .post_loop_processing(processed_input, final_content)
10087 .await?;
10088 self.apply_post_loop_result(processed_input, result)
10089 .await?
10090 .content
10091 };
10092 let response = self.build_agent_response(AgentResponseParts {
10093 content: final_content,
10094 all_tool_calls,
10095 reasoning_mode,
10096 auto_detected,
10097 iterations,
10098 thinking: thinking_content,
10099 reflection_metadata,
10100 });
10101 self.finish_turn_if_root(&response).await?;
10102 Ok(response)
10103 }
10104
10105 async fn run_committed_response_loop_with_reasoning(
10110 &self,
10111 processed_input: &str,
10112 input_context: &HashMap<String, Value>,
10113 reasoning_mode: ReasoningMode,
10114 auto_detected: bool,
10115 ) -> Result<AgentResponse> {
10116 self.commit_root_user_message(processed_input).await?;
10117 let llm = self.get_state_llm()?;
10118 let mut iterations = 0u32;
10119 let mut all_tool_calls = Vec::new();
10120 let mut thinking_content = None;
10121 loop {
10122 let effective_max = if reasoning_mode != ReasoningMode::None {
10123 let rc = self.get_effective_reasoning_config();
10124 self.max_iterations.min(rc.max_iterations)
10125 } else {
10126 self.max_iterations
10127 };
10128 if iterations >= effective_max {
10129 return Err(AgentError::Other(format!(
10130 "Max iterations ({}) exceeded",
10131 effective_max
10132 )));
10133 }
10134 iterations += 1;
10135 *self.iteration_count.write() = iterations;
10136 let protocol = self.main_tool_protocol(llm.as_ref(), false).await?;
10137 let mut messages = self
10138 .build_messages_internal(true, None, protocol.choice.is_none())
10139 .await?;
10140 self.inject_reasoning_prompt(&mut messages, &reasoning_mode, iterations == 1);
10141 self.hooks.on_llm_start(&messages).await;
10142 let llm_start = Instant::now();
10143 let response = self
10144 .complete_main_llm_with_recovery(Arc::clone(&llm), &messages, &protocol)
10145 .await?;
10146 let llm_duration_ms = llm_start.elapsed().as_millis() as u64;
10147 self.hooks.on_llm_complete(&response, llm_duration_ms).await;
10148 let content = response.content.trim();
10149 if let Some(tool_calls) = self.parse_main_tool_calls(content, &protocol)? {
10150 match self
10151 .handle_tool_calls(
10152 processed_input,
10153 content,
10154 tool_calls,
10155 &mut all_tool_calls,
10156 None,
10157 )
10158 .await?
10159 {
10160 ToolCallOutcome::Continue | ToolCallOutcome::TransitionFired => continue,
10161 ToolCallOutcome::Rejected(resp) => {
10162 self.finish_turn_if_root(&resp).await?;
10163 return Ok(resp);
10164 }
10165 }
10166 }
10167 let (extracted_thinking, answer) = self.extract_thinking(content);
10168 if extracted_thinking.is_some() {
10169 thinking_content = extracted_thinking;
10170 }
10171 return self
10172 .finish_text_response_from_model(CommittedTextResponse {
10173 processed_input,
10174 input_context,
10175 answer,
10176 reasoning_mode,
10177 auto_detected,
10178 iterations,
10179 thinking_content,
10180 all_tool_calls,
10181 })
10182 .await;
10183 }
10184 }
10185
10186 async fn handle_tool_calls(
10192 &self,
10193 processed_input: &str,
10194 content: &str,
10195 tool_calls: Vec<ToolCall>,
10196 all_tool_calls: &mut Vec<ToolCall>,
10197 mut events: Option<&mut Vec<StreamChunk>>,
10198 ) -> Result<ToolCallOutcome> {
10199 let include_tool_events = self.streaming.include_tool_events;
10200 let transition_content = native_readable_projection(content)
10204 .map_err(|error| AgentError::LLM(error.to_string()))?;
10205 let transition_fired = self
10206 .evaluate_transitions(processed_input, &transition_content)
10207 .await?;
10208 if transition_fired {
10209 self.memory
10210 .add_message(ChatMessage::assistant(
10211 "(Transitioned to new state — tool call handled by workflow)",
10212 ))
10213 .await?;
10214 if let Some(events) = events.as_deref_mut()
10215 && self.streaming.include_state_events
10216 && let Some(state) = self.current_state()
10217 {
10218 events.push(StreamChunk::state_transition(None, state));
10219 }
10220 return Ok(ToolCallOutcome::TransitionFired);
10221 }
10222
10223 self.memory
10225 .add_message(ChatMessage::assistant(content))
10226 .await?;
10227 self.remember_committed_native_exchange(content).await?;
10228 let native_tool_call = Self::is_native_tool_call_content(content)?;
10229
10230 if let Some(events) = events.as_deref_mut()
10231 && include_tool_events
10232 {
10233 for tool_call in &tool_calls {
10234 events.push(StreamChunk::tool_start(&tool_call.id, &tool_call.name));
10235 }
10236 }
10237 let results = self.execute_tools_parallel(&tool_calls).await;
10238 let mut rejection = None;
10239
10240 for ((_id, result), tool_call) in results.into_iter().zip(tool_calls.iter()) {
10241 match result {
10242 Ok(output) => {
10243 if let Some(events) = events.as_deref_mut()
10244 && include_tool_events
10245 {
10246 events.push(StreamChunk::tool_result(
10247 &tool_call.id,
10248 &tool_call.name,
10249 &output,
10250 true,
10251 ));
10252 }
10253 self.memory
10254 .add_message(Self::tool_result_message(
10255 tool_call,
10256 &output,
10257 native_tool_call,
10258 )?)
10259 .await?;
10260 }
10261 Err(e) => {
10262 if matches!(e, AgentError::HITLRejected(_)) {
10263 if !native_tool_call {
10264 self.memory
10265 .add_message(ChatMessage::assistant(format!(
10266 "The operation was rejected by the approver: {e}"
10267 )))
10268 .await?;
10269 return Ok(ToolCallOutcome::Rejected(AgentResponse {
10270 content: format!("Operation cancelled: {e}"),
10271 metadata: None,
10272 tool_calls: Some(all_tool_calls.clone()),
10273 }));
10274 }
10275 if rejection.is_none() {
10276 rejection = Some(e.to_string());
10277 }
10278 }
10279 if let Some(events) = events.as_deref_mut()
10280 && include_tool_events
10281 {
10282 events.push(StreamChunk::tool_result(
10283 &tool_call.id,
10284 &tool_call.name,
10285 e.to_string(),
10286 false,
10287 ));
10288 }
10289 self.memory
10290 .add_message(Self::tool_result_message(
10291 tool_call,
10292 &format!("Error: {}", e),
10293 native_tool_call,
10294 )?)
10295 .await?;
10296 }
10297 }
10298 all_tool_calls.push(tool_call.clone());
10299 if let Some(events) = events.as_deref_mut()
10300 && include_tool_events
10301 {
10302 events.push(StreamChunk::tool_end(&tool_call.id));
10303 }
10304 }
10305 if let Some(rejection) = rejection {
10306 self.memory
10307 .add_message(ChatMessage::assistant(format!(
10308 "The operation was rejected by the approver: {rejection}"
10309 )))
10310 .await?;
10311 return Ok(ToolCallOutcome::Rejected(AgentResponse {
10312 content: format!("Operation cancelled: {rejection}"),
10313 metadata: None,
10314 tool_calls: Some(all_tool_calls.clone()),
10315 }));
10316 }
10317 Ok(ToolCallOutcome::Continue)
10318 }
10319
10320 async fn run_reflection(
10322 &self,
10323 llm: &dyn LLMProvider,
10324 processed_input: &str,
10325 mut content: String,
10326 ) -> Result<(String, Option<ReflectionMetadata>)> {
10327 let should_reflect = self.should_reflect(processed_input, &content).await?;
10328 if !should_reflect {
10329 return Ok((content, None));
10330 }
10331
10332 info!("Starting response reflection evaluation");
10333 let mut attempts = 0u32;
10334 let max_retries = self.reflection_config.max_retries;
10335 let mut history: Vec<ReflectionAttempt> = Vec::new();
10336
10337 loop {
10338 let evaluation = self.evaluate_response(processed_input, &content).await?;
10339
10340 if evaluation.passed || attempts >= max_retries {
10341 info!(
10342 passed = evaluation.passed,
10343 confidence = evaluation.confidence,
10344 attempts = attempts + 1,
10345 "Reflection evaluation complete"
10346 );
10347 let reflection_metadata = Some(
10348 ReflectionMetadata::new(evaluation)
10349 .with_attempts(attempts + 1)
10350 .with_history(history),
10351 );
10352 return Ok((content, reflection_metadata));
10353 }
10354
10355 debug!(
10356 attempt = attempts + 1,
10357 failed_criteria = evaluation.failed_criteria().count(),
10358 "Response did not meet criteria, retrying"
10359 );
10360
10361 history.push(
10362 ReflectionAttempt::new(&content, evaluation.clone())
10363 .with_feedback("Response did not meet quality criteria"),
10364 );
10365
10366 let feedback: Vec<String> = evaluation
10367 .failed_criteria()
10368 .map(|c| format!("- {}", c.criterion))
10369 .collect();
10370
10371 let retry_prompt = format!(
10372 "Your previous response did not meet these criteria:\n{}\n\nPlease provide an improved response.",
10373 feedback.join("\n")
10374 );
10375
10376 self.memory
10377 .add_message(ChatMessage::user(&retry_prompt))
10378 .await?;
10379
10380 let retry_messages = self.build_messages().await?;
10381 let retry_response = self
10382 .observe_purpose(
10383 ObservationPurpose::ReflectionEvaluation,
10384 llm.complete(&retry_messages, None),
10385 )
10386 .await
10387 .map_err(|e| AgentError::LLM(e.to_string()))?;
10388
10389 content = retry_response.content.trim().to_string();
10390 attempts += 1;
10391 }
10392 }
10393
10394 async fn post_loop_processing(
10397 &self,
10398 processed_input: &str,
10399 content: String,
10400 ) -> Result<PostLoopResult> {
10401 self.increment_turn();
10406
10407 self.run_context_extractors(processed_input).await;
10409
10410 let transitioned = self.evaluate_transitions(processed_input, &content).await?;
10411
10412 if !transitioned {
10413 self.memory
10414 .add_message(ChatMessage::assistant(&content))
10415 .await?;
10416 self.check_memory_compression().await?;
10417 return Ok(PostLoopResult::NoTransition(content));
10418 }
10419
10420 if !self.should_regenerate_after_transition() {
10422 self.memory
10423 .add_message(ChatMessage::assistant(&content))
10424 .await?;
10425 self.check_memory_compression().await?;
10426 return Ok(PostLoopResult::Transitioned {
10427 content,
10428 regenerated: false,
10429 });
10430 }
10431
10432 if self.needs_redispatch_for_new_state() {
10436 info!("Post-transition NeedsRedispatch: new state requires full dispatch");
10437 return Ok(PostLoopResult::NeedsRedispatch);
10440 }
10441
10442 self.memory
10445 .add_message(ChatMessage::assistant(&content))
10446 .await?;
10447 self.check_memory_compression().await?;
10448
10449 let new_llm = self.get_state_llm()?;
10455 let mut final_content;
10456
10457 for post_iter in 0..self.max_iterations {
10458 let protocol = self.main_tool_protocol(new_llm.as_ref(), false).await?;
10459 let new_messages = self
10460 .build_messages_internal(true, None, protocol.choice.is_none())
10461 .await?;
10462 if post_iter == 0
10463 && let Some(system_msg) = new_messages.first()
10464 && system_msg.role == ai_agents_core::Role::System
10465 {
10466 debug!(
10467 prompt_preview =
10468 &system_msg.content[system_msg.content.len().saturating_sub(200)..],
10469 "Post-transition system prompt (last 200 chars)"
10470 );
10471 }
10472
10473 let new_response = self
10474 .complete_main_llm_with_recovery(Arc::clone(&new_llm), &new_messages, &protocol)
10475 .await?;
10476 final_content = new_response.content.trim().to_string();
10477
10478 if let Some(tool_calls) = self.parse_main_tool_calls(&final_content, &protocol)? {
10481 let native_tool_call = Self::is_native_tool_call_content(&final_content)?;
10482 debug!(
10483 post_iter = post_iter,
10484 tools = tool_calls.len(),
10485 "Post-transition tool call detected, executing"
10486 );
10487
10488 self.memory
10489 .add_message(ChatMessage::assistant(&final_content))
10490 .await?;
10491 self.remember_committed_native_exchange(&final_content)
10492 .await?;
10493
10494 let results = self.execute_tools_parallel(&tool_calls).await;
10495 let mut rejection = None;
10496 for ((_id, result), tool_call) in results.into_iter().zip(tool_calls.iter()) {
10497 match result {
10498 Ok(output) => {
10499 self.memory
10500 .add_message(Self::tool_result_message(
10501 tool_call,
10502 &output,
10503 native_tool_call,
10504 )?)
10505 .await?;
10506 }
10507 Err(e) => {
10508 if native_tool_call
10509 && rejection.is_none()
10510 && matches!(e, AgentError::HITLRejected(_))
10511 {
10512 rejection = Some(e.to_string());
10513 }
10514 self.memory
10515 .add_message(Self::tool_result_message(
10516 tool_call,
10517 &format!("Error: {}", e),
10518 native_tool_call,
10519 )?)
10520 .await?;
10521 }
10522 }
10523 }
10524 if let Some(rejection) = rejection {
10525 self.memory
10526 .add_message(ChatMessage::assistant(format!(
10527 "The operation was rejected by the approver: {rejection}"
10528 )))
10529 .await?;
10530 return Err(AgentError::HITLRejected(rejection));
10531 }
10532 continue;
10534 }
10535
10536 self.memory
10538 .add_message(ChatMessage::assistant(&final_content))
10539 .await?;
10540 return Ok(PostLoopResult::Transitioned {
10541 content: final_content,
10542 regenerated: true,
10543 });
10544 }
10545
10546 final_content = "Post-transition processing completed.".to_string();
10548 self.memory
10549 .add_message(ChatMessage::assistant(&final_content))
10550 .await?;
10551
10552 Ok(PostLoopResult::Transitioned {
10553 content: final_content,
10554 regenerated: true,
10555 })
10556 }
10557
10558 fn should_regenerate_after_transition(&self) -> bool {
10561 if let Some(ref sm) = self.state_machine {
10562 if !sm.config().regenerate_on_transition {
10564 return false;
10565 }
10566 if let Some(def) = sm.current_definition()
10568 && let Some(regen) = def.regenerate_on_enter
10569 {
10570 return regen;
10571 }
10572 }
10573 true
10574 }
10575
10576 fn needs_redispatch_for_new_state(&self) -> bool {
10579 if let Some(ref sm) = self.state_machine
10580 && let Some(def) = sm.current_definition()
10581 {
10582 if def.concurrent.is_some()
10583 || def.group_chat.is_some()
10584 || def.pipeline.is_some()
10585 || def.handoff.is_some()
10586 || def.delegate.is_some()
10587 {
10588 return true;
10589 }
10590 let effective = self.get_effective_reasoning_config();
10592 if !matches!(effective.mode, ReasoningMode::None) {
10593 return true;
10594 }
10595 }
10596 false
10597 }
10598
10599 async fn apply_post_loop_result(
10605 &self,
10606 processed_input: &str,
10607 result: PostLoopResult,
10608 ) -> Result<AppliedPostLoop> {
10609 match result {
10610 PostLoopResult::NoTransition(content) => Ok(AppliedPostLoop {
10611 content,
10612 transitioned: false,
10613 regenerated: false,
10614 }),
10615 PostLoopResult::Transitioned {
10616 content,
10617 regenerated,
10618 } => Ok(AppliedPostLoop {
10619 content,
10620 transitioned: true,
10621 regenerated,
10622 }),
10623 PostLoopResult::NeedsRedispatch => {
10624 const MAX_REDISPATCH_DEPTH: u32 = 3;
10625 let current_depth = *self.redispatch_depth.read();
10626 if current_depth >= MAX_REDISPATCH_DEPTH {
10627 warn!(
10628 depth = current_depth,
10629 "Post-transition re-dispatch depth limit reached, returning empty response"
10630 );
10631 let content = String::new();
10632 self.memory
10633 .add_message(ChatMessage::assistant(&content))
10634 .await?;
10635 return Ok(AppliedPostLoop {
10637 content,
10638 transitioned: true,
10639 regenerated: false,
10640 });
10641 }
10642 *self.redispatch_depth.write() += 1;
10643 if let Some(context) = self.active_turn_context.write().as_mut() {
10644 context.enter_redispatch();
10645 }
10646 info!(
10647 depth = current_depth + 1,
10648 "Re-dispatching for new state after transition"
10649 );
10650 let resp = Box::pin(self.run_loop_internal(processed_input)).await;
10651 *self.redispatch_depth.write() -= 1;
10652 if let Some(context) = self.active_turn_context.write().as_mut() {
10653 context.exit_redispatch();
10654 }
10655 resp.map(|r| AppliedPostLoop {
10656 content: r.content,
10657 transitioned: true,
10658 regenerated: true,
10659 })
10660 }
10661 }
10662 }
10663
10664 fn build_agent_response(&self, parts: AgentResponseParts) -> AgentResponse {
10666 let AgentResponseParts {
10667 content,
10668 all_tool_calls,
10669 reasoning_mode,
10670 auto_detected,
10671 iterations,
10672 thinking,
10673 reflection_metadata,
10674 } = parts;
10675 let reasoning_metadata = ReasoningMetadata::new(reasoning_mode.clone())
10676 .with_thinking(thinking.clone().unwrap_or_default())
10677 .with_iterations(iterations)
10678 .with_auto_detected(auto_detected);
10679
10680 let mut response = AgentResponse::new(&content);
10681 if !all_tool_calls.is_empty() {
10682 response = response.with_tool_calls(all_tool_calls);
10683 }
10684
10685 if let Some(state) = self.current_state() {
10686 response = response.with_metadata("current_state", serde_json::json!(state));
10687 }
10688
10689 response = response.with_metadata(
10690 "reasoning",
10691 serde_json::to_value(&reasoning_metadata).unwrap_or_default(),
10692 );
10693
10694 if let Some(ref refl_meta) = reflection_metadata {
10695 response = response.with_metadata(
10696 "reflection",
10697 serde_json::to_value(refl_meta).unwrap_or_default(),
10698 );
10699 }
10700
10701 response
10702 }
10703
10704 async fn handle_delegated_state(
10706 &self,
10707 input: &str,
10708 delegate_id: &str,
10709 state_def: &ai_agents_state::StateDefinition,
10710 ) -> Result<AgentResponse> {
10711 use std::time::Instant;
10712
10713 let registry = self.spawner_registry.as_ref().ok_or_else(|| {
10714 AgentError::Config(format!(
10715 "State delegates to '{}' but no agent registry is configured. \
10716 Add a spawner section with auto_spawn to your YAML.",
10717 delegate_id
10718 ))
10719 })?;
10720
10721 let state_name = self
10722 .state_machine
10723 .as_ref()
10724 .map(|sm| sm.current())
10725 .unwrap_or_else(|| "unknown".to_string());
10726
10727 self.hooks.on_delegate_start(delegate_id, &state_name).await;
10728 let start = Instant::now();
10729
10730 let delegate = registry.get(delegate_id).ok_or_else(|| {
10731 AgentError::Other(format!(
10732 "State '{}' delegates to '{}' but no agent with that ID exists in the registry.",
10733 state_name, delegate_id
10734 ))
10735 })?;
10736
10737 let context_mode = state_def.delegate_context.clone().unwrap_or_default();
10739 let effective_input = self
10740 .observe_purpose(
10741 ObservationPurpose::OrchestrationRouting,
10742 crate::orchestration::context::prepare_delegate_input(
10743 input,
10744 &context_mode,
10745 &*self.memory,
10746 self.llm_registry.get("router").ok().as_deref(),
10747 ),
10748 )
10749 .await?;
10750
10751 let response = delegate
10752 .chat_with_actor_context(&effective_input, self.outbound_actor_context())
10753 .await?;
10754
10755 let duration_ms = start.elapsed().as_millis() as u64;
10756 self.hooks
10757 .on_delegate_complete(delegate_id, &state_name, duration_ms)
10758 .await;
10759
10760 let ctx_key = format!("delegation.{}.last_response", delegate_id);
10762 let _ = self.context_manager.set(
10763 &ctx_key,
10764 serde_json::Value::String(response.content.clone()),
10765 );
10766
10767 let _ = self.context_manager.set(
10769 "orchestration",
10770 serde_json::json!({
10771 "type": "delegate",
10772 "agent": delegate_id,
10773 "state": state_name,
10774 "response": response.content,
10775 "duration_ms": duration_ms,
10776 }),
10777 );
10778
10779 self.commit_root_user_message(input).await?;
10780
10781 let post_result = self
10784 .post_loop_processing(
10785 input,
10786 format!("[Delegated to {}]: {}", delegate_id, response.content),
10787 )
10788 .await?;
10789 let final_content = self
10790 .apply_post_loop_result(input, post_result)
10791 .await?
10792 .content;
10793
10794 let mut result = AgentResponse::new(final_content);
10795
10796 let metadata = serde_json::json!({
10797 "orchestration": {
10798 "type": "delegate",
10799 "agent": delegate_id,
10800 "state": state_name,
10801 "response": response.content,
10802 "duration_ms": duration_ms,
10803 }
10804 });
10805 result.metadata = Some(
10806 serde_json::from_value::<std::collections::HashMap<String, serde_json::Value>>(
10807 metadata,
10808 )
10809 .unwrap_or_default(),
10810 );
10811
10812 self.finish_turn_if_root(&result).await?;
10813 Ok(result)
10814 }
10815
10816 async fn handle_concurrent_state(
10818 &self,
10819 input: &str,
10820 config: &ai_agents_state::ConcurrentStateConfig,
10821 ) -> Result<AgentResponse> {
10822 use std::time::Instant;
10823
10824 let registry = self.spawner_registry.as_ref().ok_or_else(|| {
10825 AgentError::Config(
10826 "Concurrent state requires an agent registry. Add a spawner section.".into(),
10827 )
10828 })?;
10829
10830 let context_mode = config.context_mode.clone().unwrap_or_default();
10835 let context_input = self
10836 .observe_purpose(
10837 ObservationPurpose::OrchestrationRouting,
10838 crate::orchestration::context::prepare_delegate_input(
10839 input,
10840 &context_mode,
10841 &*self.memory,
10842 self.llm_registry.get("router").ok().as_deref(),
10843 ),
10844 )
10845 .await?;
10846
10847 let effective_input = if let Some(ref tmpl) = config.input {
10848 render_concurrent_template(tmpl, &context_input, &self.build_context_with_overlays())
10849 .unwrap_or_else(|_| context_input.clone())
10850 } else {
10851 context_input
10852 };
10853
10854 let start = Instant::now();
10855
10856 let llm_name = config
10857 .aggregation
10858 .synthesizer_llm
10859 .as_deref()
10860 .unwrap_or("router");
10861 let llm_provider = self.llm_registry.get(llm_name).ok();
10862
10863 let vote_parallelism = if self.runtime_config.optimization.enabled
10864 && self
10865 .runtime_config
10866 .optimization
10867 .parallel_orchestration_vote_extraction
10868 {
10869 Some(self.runtime_config.optimization.max_parallel_runtime_tasks)
10870 } else {
10871 None
10872 };
10873
10874 let result = self
10875 .observe_purpose(
10876 ObservationPurpose::OrchestrationAggregation,
10877 scope_actor_context(
10878 self.outbound_actor_context(),
10879 crate::orchestration::concurrent(
10880 registry,
10881 &effective_input,
10882 &config.agents,
10883 &config.aggregation,
10884 llm_provider.as_deref(),
10885 config.min_required,
10886 config.timeout_ms,
10887 config.on_partial_failure.clone(),
10888 vote_parallelism,
10889 ),
10890 ),
10891 )
10892 .await?;
10893
10894 let duration_ms = start.elapsed().as_millis() as u64;
10895 let agent_ids: Vec<String> = config.agents.iter().map(|a| a.id().to_string()).collect();
10896 let strategy = format!("{:?}", config.aggregation.strategy);
10897 self.hooks
10898 .on_concurrent_complete(&agent_ids, &strategy, duration_ms)
10899 .await;
10900
10901 let _ = self.context_manager.set(
10903 "concurrent.result",
10904 serde_json::Value::String(result.response.content.clone()),
10905 );
10906
10907 let agents_json: Vec<serde_json::Value> = result
10909 .agent_results
10910 .iter()
10911 .map(|ar| {
10912 serde_json::json!({
10913 "id": ar.agent_id,
10914 "response": ar.response.as_ref().map(|r| r.content.as_str()),
10915 "success": ar.success,
10916 "error": ar.error,
10917 "duration_ms": ar.duration_ms,
10918 })
10919 })
10920 .collect();
10921
10922 let _ = self.context_manager.set(
10924 "orchestration",
10925 serde_json::json!({
10926 "type": "concurrent",
10927 "result": result.response.content,
10928 "strategy": strategy,
10929 "agents": agents_json,
10930 "duration_ms": duration_ms,
10931 }),
10932 );
10933
10934 self.commit_root_user_message(input).await?;
10935
10936 let post_result = self
10937 .post_loop_processing(input, result.response.content.clone())
10938 .await?;
10939 let final_content = self
10940 .apply_post_loop_result(input, post_result)
10941 .await?
10942 .content;
10943
10944 let mut response = AgentResponse::new(final_content);
10945 let metadata = serde_json::json!({
10946 "orchestration": {
10947 "type": "concurrent",
10948 "result": result.response.content,
10949 "strategy": strategy,
10950 "agents": agents_json,
10951 "duration_ms": duration_ms,
10952 }
10953 });
10954 response.metadata = Some(
10955 serde_json::from_value::<std::collections::HashMap<String, serde_json::Value>>(
10956 metadata,
10957 )
10958 .unwrap_or_default(),
10959 );
10960
10961 self.finish_turn_if_root(&response).await?;
10962 Ok(response)
10963 }
10964
10965 async fn handle_group_chat_state(
10967 &self,
10968 input: &str,
10969 config: &ai_agents_state::GroupChatStateConfig,
10970 ) -> Result<AgentResponse> {
10971 use std::time::Instant;
10972
10973 let registry = self.spawner_registry.as_ref().ok_or_else(|| {
10974 AgentError::Config(
10975 "Group chat state requires an agent registry. Add a spawner section.".into(),
10976 )
10977 })?;
10978
10979 let start = Instant::now();
10980
10981 let llm_provider = self.llm_registry.get("router").ok();
10982
10983 let context_mode = config.context_mode.clone().unwrap_or_default();
10985 let context_input = self
10986 .observe_purpose(
10987 ObservationPurpose::OrchestrationRouting,
10988 crate::orchestration::context::prepare_delegate_input(
10989 input,
10990 &context_mode,
10991 &*self.memory,
10992 self.llm_registry.get("router").ok().as_deref(),
10993 ),
10994 )
10995 .await?;
10996
10997 let effective_topic = if let Some(ref tmpl) = config.input {
10999 render_concurrent_template(tmpl, &context_input, &self.build_context_with_overlays())
11000 .unwrap_or_else(|_| context_input.clone())
11001 } else {
11002 context_input
11003 };
11004
11005 let result = self
11006 .observe_purpose(
11007 ObservationPurpose::OrchestrationConversation,
11008 scope_actor_context(
11009 self.outbound_actor_context(),
11010 crate::orchestration::group_chat(
11011 registry,
11012 &effective_topic,
11013 config,
11014 llm_provider.as_deref(),
11015 Some(&*self.hooks),
11016 ),
11017 ),
11018 )
11019 .await?;
11020
11021 let duration_ms = start.elapsed().as_millis() as u64;
11022
11023 let _ = self.context_manager.set(
11025 "group_chat.conclusion",
11026 serde_json::Value::String(result.response.content.clone()),
11027 );
11028
11029 let transcript_json: Vec<serde_json::Value> = result
11031 .transcript
11032 .iter()
11033 .map(|t| {
11034 serde_json::json!({
11035 "speaker": t.speaker,
11036 "round": t.round,
11037 "content": t.content,
11038 })
11039 })
11040 .collect();
11041
11042 let _ = self.context_manager.set(
11044 "orchestration",
11045 serde_json::json!({
11046 "type": "group_chat",
11047 "conclusion": result.response.content,
11048 "transcript": transcript_json,
11049 "rounds": result.rounds_completed,
11050 "termination": result.termination_reason,
11051 "duration_ms": duration_ms,
11052 }),
11053 );
11054
11055 self.commit_root_user_message(input).await?;
11056
11057 let post_result = self
11058 .post_loop_processing(input, result.response.content.clone())
11059 .await?;
11060 let final_content = self
11061 .apply_post_loop_result(input, post_result)
11062 .await?
11063 .content;
11064
11065 let mut response = AgentResponse::new(final_content);
11066 let metadata = serde_json::json!({
11067 "orchestration": {
11068 "type": "group_chat",
11069 "conclusion": result.response.content,
11070 "transcript": transcript_json,
11071 "rounds": result.rounds_completed,
11072 "termination": result.termination_reason,
11073 "duration_ms": duration_ms,
11074 }
11075 });
11076 response.metadata = Some(
11077 serde_json::from_value::<std::collections::HashMap<String, serde_json::Value>>(
11078 metadata,
11079 )
11080 .unwrap_or_default(),
11081 );
11082
11083 self.finish_turn_if_root(&response).await?;
11084 Ok(response)
11085 }
11086
11087 async fn handle_pipeline_state(
11089 &self,
11090 input: &str,
11091 config: &ai_agents_state::PipelineStateConfig,
11092 ) -> Result<AgentResponse> {
11093 use std::time::Instant;
11094
11095 let registry = self.spawner_registry.as_ref().ok_or_else(|| {
11096 AgentError::Config(
11097 "Pipeline state requires an agent registry. Add a spawner section.".into(),
11098 )
11099 })?;
11100
11101 let start = Instant::now();
11102
11103 let stages: Vec<crate::orchestration::PipelineStage> = config
11104 .stages
11105 .iter()
11106 .map(|entry| {
11107 let mut stage = crate::orchestration::PipelineStage::id(entry.id());
11108 if let Some(tmpl) = entry.input() {
11109 stage = stage.with_input(tmpl);
11110 }
11111 stage
11112 })
11113 .collect();
11114
11115 let context_mode = config.context_mode.clone().unwrap_or_default();
11117 let context_input = self
11118 .observe_purpose(
11119 ObservationPurpose::OrchestrationRouting,
11120 crate::orchestration::context::prepare_delegate_input(
11121 input,
11122 &context_mode,
11123 &*self.memory,
11124 self.llm_registry.get("router").ok().as_deref(),
11125 ),
11126 )
11127 .await?;
11128
11129 let context_values = self.build_context_with_overlays();
11130 let result = self
11131 .observe_purpose(
11132 ObservationPurpose::OrchestrationRouting,
11133 scope_actor_context(
11134 self.outbound_actor_context(),
11135 crate::orchestration::pipeline(
11136 registry,
11137 &context_input,
11138 &stages,
11139 config.timeout_ms,
11140 Some(&*self.hooks),
11141 Some(&context_values),
11142 ),
11143 ),
11144 )
11145 .await?;
11146
11147 let duration_ms = start.elapsed().as_millis() as u64;
11148
11149 let _ = self.context_manager.set(
11151 "pipeline.result",
11152 serde_json::Value::String(result.response.content.clone()),
11153 );
11154
11155 let stages_json: Vec<serde_json::Value> = result
11157 .stage_outputs
11158 .iter()
11159 .map(|s| {
11160 serde_json::json!({
11161 "agent_id": s.agent_id,
11162 "output": s.output,
11163 "duration_ms": s.duration_ms,
11164 "skipped": s.skipped,
11165 })
11166 })
11167 .collect();
11168
11169 let _ = self.context_manager.set(
11171 "orchestration",
11172 serde_json::json!({
11173 "type": "pipeline",
11174 "result": result.response.content,
11175 "stages": stages_json,
11176 "duration_ms": duration_ms,
11177 }),
11178 );
11179
11180 self.commit_root_user_message(input).await?;
11181
11182 let post_result = self
11183 .post_loop_processing(input, result.response.content.clone())
11184 .await?;
11185 let final_content = self
11186 .apply_post_loop_result(input, post_result)
11187 .await?
11188 .content;
11189
11190 let mut response = AgentResponse::new(final_content);
11191 let metadata = serde_json::json!({
11192 "orchestration": {
11193 "type": "pipeline",
11194 "result": result.response.content,
11195 "stages": stages_json,
11196 "duration_ms": duration_ms,
11197 }
11198 });
11199 response.metadata = Some(
11200 serde_json::from_value::<std::collections::HashMap<String, serde_json::Value>>(
11201 metadata,
11202 )
11203 .unwrap_or_default(),
11204 );
11205
11206 self.finish_turn_if_root(&response).await?;
11207 Ok(response)
11208 }
11209
11210 async fn handle_handoff_state(
11212 &self,
11213 input: &str,
11214 config: &ai_agents_state::HandoffStateConfig,
11215 ) -> Result<AgentResponse> {
11216 use std::time::Instant;
11217
11218 let registry = self.spawner_registry.as_ref().ok_or_else(|| {
11219 AgentError::Config(
11220 "Handoff state requires an agent registry. Add a spawner section.".into(),
11221 )
11222 })?;
11223
11224 let llm = self
11225 .llm_registry
11226 .get("router")
11227 .map_err(|_| AgentError::Config("Handoff state requires a router LLM.".into()))?;
11228
11229 let start = Instant::now();
11230
11231 let context_mode = config.context_mode.clone().unwrap_or_default();
11233 let context_input = self
11234 .observe_purpose(
11235 ObservationPurpose::OrchestrationRouting,
11236 crate::orchestration::context::prepare_delegate_input(
11237 input,
11238 &context_mode,
11239 &*self.memory,
11240 self.llm_registry.get("router").ok().as_deref(),
11241 ),
11242 )
11243 .await?;
11244
11245 let effective_input = if let Some(ref tmpl) = config.input {
11247 render_concurrent_template(tmpl, &context_input, &self.build_context_with_overlays())
11248 .unwrap_or_else(|_| context_input.clone())
11249 } else {
11250 context_input
11251 };
11252
11253 let result = self
11254 .observe_purpose(
11255 ObservationPurpose::OrchestrationRouting,
11256 scope_actor_context(
11257 self.outbound_actor_context(),
11258 crate::orchestration::handoff(
11259 registry,
11260 &effective_input,
11261 &config.initial_agent,
11262 &config.available_agents,
11263 config.max_handoffs,
11264 llm.as_ref(),
11265 Some(&*self.hooks),
11266 ),
11267 ),
11268 )
11269 .await?;
11270
11271 let duration_ms = start.elapsed().as_millis() as u64;
11272
11273 let _ = self.context_manager.set(
11275 "handoff.result",
11276 serde_json::Value::String(result.response.content.clone()),
11277 );
11278
11279 let chain_json: Vec<serde_json::Value> = result
11281 .handoff_chain
11282 .iter()
11283 .map(|h| {
11284 serde_json::json!({
11285 "from": h.from_agent,
11286 "to": h.to_agent,
11287 "reason": h.reason,
11288 })
11289 })
11290 .collect();
11291
11292 let _ = self.context_manager.set(
11294 "orchestration",
11295 serde_json::json!({
11296 "type": "handoff",
11297 "result": result.response.content,
11298 "final_agent": result.final_agent,
11299 "handoff_chain": chain_json,
11300 "duration_ms": duration_ms,
11301 }),
11302 );
11303
11304 self.commit_root_user_message(input).await?;
11305
11306 let post_result = self
11307 .post_loop_processing(input, result.response.content.clone())
11308 .await?;
11309 let final_content = self
11310 .apply_post_loop_result(input, post_result)
11311 .await?
11312 .content;
11313
11314 let mut response = AgentResponse::new(final_content);
11315 let metadata = serde_json::json!({
11316 "orchestration": {
11317 "type": "handoff",
11318 "result": result.response.content,
11319 "final_agent": result.final_agent,
11320 "handoff_chain": chain_json,
11321 "duration_ms": duration_ms,
11322 }
11323 });
11324 response.metadata = Some(
11325 serde_json::from_value::<std::collections::HashMap<String, serde_json::Value>>(
11326 metadata,
11327 )
11328 .unwrap_or_default(),
11329 );
11330
11331 self.finish_turn_if_root(&response).await?;
11332 Ok(response)
11333 }
11334
11335 async fn run_loop_internal(&self, input: &str) -> Result<AgentResponse> {
11337 self.begin_root_turn();
11338 self.pre_turn_session_lifecycle().await;
11340
11341 let input_data = self.process_input(input).await?;
11342 self.update_active_turn_context(&input_data.content, input_data.context.clone());
11343
11344 for (key, value) in &input_data.context {
11347 let _ = self.context_manager.set(key, value.clone());
11348 }
11349
11350 if input_data.metadata.rejected {
11351 let reason = input_data
11352 .metadata
11353 .rejection_reason
11354 .unwrap_or_else(|| "Input rejected".to_string());
11355 warn!(reason = %reason, "Input rejected");
11356 let response = AgentResponse::new(reason);
11357 self.finish_turn_if_root(&response).await?;
11358 return Ok(response);
11359 }
11360
11361 let processed_input = &input_data.content;
11362
11363 if let Some(response) = self.try_pre_response_transition(processed_input).await? {
11364 return Ok(response);
11365 }
11366
11367 if let Some(ref sm) = self.state_machine
11369 && let Some(def) = sm.current_definition()
11370 {
11371 if let Some(ref delegate_id) = def.delegate {
11372 return self
11373 .handle_delegated_state(processed_input, delegate_id, &def)
11374 .await;
11375 }
11376 if let Some(ref concurrent_config) = def.concurrent {
11377 return self
11378 .handle_concurrent_state(processed_input, concurrent_config)
11379 .await;
11380 }
11381 if let Some(ref group_chat_config) = def.group_chat {
11382 return self
11383 .handle_group_chat_state(processed_input, group_chat_config)
11384 .await;
11385 }
11386 if let Some(ref pipeline_config) = def.pipeline {
11387 return self
11388 .handle_pipeline_state(processed_input, pipeline_config)
11389 .await;
11390 }
11391 if let Some(ref handoff_config) = def.handoff {
11392 return self
11393 .handle_handoff_state(processed_input, handoff_config)
11394 .await;
11395 }
11396 }
11397
11398 if let Some(response) =
11403 Box::pin(self.try_speculative_branches(processed_input, &input_data.context)).await?
11404 {
11405 return Ok(response);
11406 }
11407
11408 match self.try_skill_route(processed_input).await? {
11409 SkillRouteResult::Response { skill_id, content } => {
11410 self.commit_root_user_message(processed_input).await?;
11411 return self
11412 .handle_skill_response(processed_input, &skill_id, content, &input_data.context)
11413 .await;
11414 }
11415 SkillRouteResult::NeedsClarification {
11416 response,
11417 ownership,
11418 } => {
11419 let admission = self
11420 .admit_optional_disambiguation_ownership(ownership)
11421 .await?;
11422 self.commit_root_user_message(processed_input).await?;
11423 if Self::skill_clarification_needs_memory_record(&response) {
11424 self.memory
11427 .add_message(ChatMessage::assistant(&response.content))
11428 .await?;
11429 }
11430 drop(admission);
11431 self.finish_turn_if_root(&response).await?;
11432 return Ok(response);
11433 }
11434 SkillRouteResult::NoMatch => {} }
11436
11437 let effective_reasoning = self.get_effective_reasoning_config();
11438 let reasoning_mode = self.determine_reasoning_mode(processed_input).await?;
11439 let auto_detected = matches!(effective_reasoning.mode, ReasoningMode::Auto);
11440
11441 info!(
11442 reasoning_mode = ?reasoning_mode,
11443 auto_detected = auto_detected,
11444 reflection_enabled = ?self.reflection_config.enabled,
11445 "Reasoning mode determined"
11446 );
11447
11448 if matches!(reasoning_mode, ReasoningMode::PlanAndExecute) {
11449 self.commit_root_user_message(processed_input).await?;
11450 return self
11451 .handle_plan_and_execute(processed_input, &input_data.context, auto_detected)
11452 .await;
11453 }
11454
11455 self.commit_root_user_message(processed_input).await?;
11456
11457 let mut iterations = 0u32;
11458 let mut all_tool_calls: Vec<ToolCall> = Vec::new();
11459 let mut thinking_content: Option<String> = None;
11460
11461 let llm = self.get_state_llm()?;
11462
11463 loop {
11464 let effective_max = if reasoning_mode != ReasoningMode::None {
11466 let rc = self.get_effective_reasoning_config();
11467 self.max_iterations.min(rc.max_iterations)
11468 } else {
11469 self.max_iterations
11470 };
11471
11472 if iterations >= effective_max {
11473 let err = AgentError::Other(format!("Max iterations ({}) exceeded", effective_max));
11474 self.hooks.on_error(&err).await;
11475 error!(iterations = iterations, "Max iterations exceeded");
11476 return Err(err);
11477 }
11478 iterations += 1;
11479 *self.iteration_count.write() = iterations;
11480
11481 debug!(iteration = iterations, max = effective_max, "LLM call");
11482
11483 let protocol = self.main_tool_protocol(llm.as_ref(), false).await?;
11484 let mut messages = self
11485 .build_messages_internal(true, None, protocol.choice.is_none())
11486 .await?;
11487 self.inject_reasoning_prompt(&mut messages, &reasoning_mode, iterations == 1);
11488
11489 self.hooks.on_llm_start(&messages).await;
11490 let llm_start = Instant::now();
11491 let response = self
11492 .complete_main_llm_with_recovery(Arc::clone(&llm), &messages, &protocol)
11493 .await?;
11494
11495 let llm_duration_ms = llm_start.elapsed().as_millis() as u64;
11496 self.hooks.on_llm_complete(&response, llm_duration_ms).await;
11497
11498 let content = response.content.trim();
11499
11500 if let Some(tool_calls) = self.parse_main_tool_calls(content, &protocol)? {
11501 match self
11502 .handle_tool_calls(
11503 processed_input,
11504 content,
11505 tool_calls,
11506 &mut all_tool_calls,
11507 None,
11508 )
11509 .await?
11510 {
11511 ToolCallOutcome::Continue | ToolCallOutcome::TransitionFired => continue,
11512 ToolCallOutcome::Rejected(resp) => {
11513 self.finish_turn_if_root(&resp).await?;
11514 return Ok(resp);
11515 }
11516 }
11517 }
11518
11519 let (extracted_thinking, answer) = self.extract_thinking(content);
11520 if extracted_thinking.is_some() {
11521 thinking_content = extracted_thinking;
11522 }
11523
11524 let output_data = self.process_output(&answer, &input_data.context).await?;
11525
11526 let mut final_content = if output_data.metadata.rejected {
11527 output_data
11528 .metadata
11529 .rejection_reason
11530 .unwrap_or_else(|| answer.to_string())
11531 } else {
11532 output_data.content
11533 };
11534
11535 let reflection_metadata;
11537 (final_content, reflection_metadata) = self
11538 .run_reflection(&*llm, processed_input, final_content)
11539 .await?;
11540
11541 final_content =
11542 self.format_response_with_thinking(thinking_content.as_deref(), &final_content);
11543
11544 let final_content = {
11548 let result = self
11549 .post_loop_processing(processed_input, final_content)
11550 .await?;
11551 self.apply_post_loop_result(processed_input, result)
11552 .await?
11553 .content
11554 };
11555
11556 let reflected = reflection_metadata.is_some();
11557 let reasoning_mode_debug = format!("{:?}", reasoning_mode);
11558
11559 let response = self.build_agent_response(AgentResponseParts {
11560 content: final_content,
11561 all_tool_calls,
11562 reasoning_mode,
11563 auto_detected,
11564 iterations,
11565 thinking: thinking_content,
11566 reflection_metadata,
11567 });
11568
11569 self.finish_turn_if_root(&response).await?;
11570
11571 let tool_call_count = response.tool_calls.as_ref().map(|tc| tc.len()).unwrap_or(0);
11572 info!(
11573 tool_calls = tool_call_count,
11574 response_len = response.content.len(),
11575 reasoning_mode = %reasoning_mode_debug,
11576 reflected = reflected,
11577 "Chat completed"
11578 );
11579 return Ok(response);
11580 }
11581 }
11582
11583 async fn generate_buffered_streaming_draft(
11584 &self,
11585 processed_input: &str,
11586 routing_resolved: Arc<AtomicBool>,
11587 ) -> Result<StreamingDraftResult> {
11588 let llm = self.get_state_llm()?;
11589 if llm.configured_tool_choice().is_some() {
11590 let draft = self
11591 .generate_main_response_draft(processed_input, &ReasoningMode::None)
11592 .await?;
11593 return Ok(StreamingDraftResult::new(draft, Vec::new()));
11594 }
11595 let protocol = self.main_tool_protocol(llm.as_ref(), true).await?;
11597 let messages = self.build_messages_for_draft(processed_input).await?;
11598 let source = self
11599 .open_main_stream_with_recovery(Arc::clone(&llm), &messages, &protocol)
11600 .await?;
11601 let mut buffer = crate::optimization::StreamBranchBuffer::new(self.streaming.buffer_size)?;
11602 let mut chunks = Vec::new();
11603 let mut accumulated = String::new();
11604 match source {
11605 MainStreamSource::StaticResponse(text) => {
11606 accumulated.push_str(&text);
11607 let stream_chunk = StreamChunk::content(text);
11608 if routing_resolved.load(Ordering::SeqCst) {
11609 chunks.push(stream_chunk);
11610 } else {
11611 buffer.push(stream_chunk)?;
11612 }
11613 }
11614 MainStreamSource::Stream(mut stream) => {
11615 while let Some(chunk_result) = stream.next().await {
11616 let chunk = chunk_result.map_err(|e| AgentError::LLM(e.to_string()))?;
11617 accumulated.push_str(&chunk.delta);
11618 let stream_chunk = StreamChunk::content(chunk.delta);
11619 if routing_resolved.load(Ordering::SeqCst) {
11620 chunks.push(stream_chunk);
11621 } else {
11622 buffer.push(stream_chunk)?;
11623 }
11624 }
11625 }
11626 }
11627 chunks.splice(0..0, buffer.drain());
11628 let content = accumulated.trim().to_string();
11629 let draft = if let Some(calls) = self.parse_tool_calls(&content)? {
11630 MainResponseDraft::ToolCalls {
11631 raw_content: content,
11632 calls,
11633 thinking: None,
11634 }
11635 } else {
11636 MainResponseDraft::Text {
11637 raw_content: content,
11638 thinking: None,
11639 }
11640 };
11641 Ok(StreamingDraftResult::new(draft, chunks))
11642 }
11643
11644 async fn try_buffered_streaming_branches(
11645 &self,
11646 processed_input: &str,
11647 input_context: &HashMap<String, Value>,
11648 ) -> Result<Option<(AgentResponse, Vec<StreamChunk>)>> {
11649 let optimization = &self.runtime_config.optimization;
11650 if !optimization.enabled {
11651 return Ok(None);
11652 }
11653 if !matches!(
11659 self.get_effective_reasoning_config().mode,
11660 ReasoningMode::None
11661 ) {
11662 return Ok(None);
11663 }
11664 let transition_enabled =
11665 optimization.speculative_state_transitions && self.has_parallel_transition_candidates();
11666 if !transition_enabled {
11667 return Ok(None);
11668 }
11669 let mut branch_scheduler =
11670 TurnBranchScheduler::new(optimization.max_parallel_runtime_tasks)?;
11671 if !branch_scheduler.reserve_task() {
11672 return Ok(None);
11673 }
11674 if !self
11675 .reserve_active_speculative_llm_call(RuntimeOptimizationKind::BufferedStreamingRouting)
11676 {
11677 branch_scheduler.release_task();
11678 return Ok(None);
11679 }
11680 if !branch_scheduler.reserve_task() {
11681 branch_scheduler.release_task();
11682 return Ok(None);
11683 }
11684 let mut main_branch = RuntimeBranch::new(
11685 RuntimeTaskPurpose::MainResponse,
11686 RuntimeOptimizationKind::BufferedStreamingRouting,
11687 RuntimeTaskPriority::Normal,
11688 RuntimeCommitBehavior::FinalResponse,
11689 );
11690 let mut transition_branch = RuntimeBranch::new(
11691 RuntimeTaskPurpose::StateTransition,
11692 RuntimeOptimizationKind::ParallelStateTransition,
11693 RuntimeTaskPriority::Critical,
11694 RuntimeCommitBehavior::TransitionDecision,
11695 );
11696 let main_id = main_branch.branch_id();
11697 let transition_id = transition_branch.branch_id();
11698 let routing_resolved = Arc::new(AtomicBool::new(false));
11699 let mut main_future =
11700 Box::pin(crate::optimization::observability::with_branch_observation(
11701 &main_id,
11702 RuntimeOptimizationKind::BufferedStreamingRouting,
11703 RuntimeCommitBehavior::FinalResponse,
11704 self.generate_buffered_streaming_draft(
11705 processed_input,
11706 Arc::clone(&routing_resolved),
11707 ),
11708 ));
11709 let mut transition_future =
11710 Box::pin(crate::optimization::observability::with_branch_observation(
11711 &transition_id,
11712 RuntimeOptimizationKind::ParallelStateTransition,
11713 RuntimeCommitBehavior::TransitionDecision,
11714 self.select_parallel_transition_candidate(processed_input),
11715 ));
11716 let mut main_pending = true;
11717 let mut transition_pending = true;
11718 let mut main_result: Option<Result<StreamingDraftResult>> = None;
11719 let mut transition_finalized = false;
11720 let mut transition_candidate: Option<TransitionCandidate> = None;
11721 loop {
11722 if let Some(candidate) = transition_candidate.take() {
11723 if self
11724 .approve_transition_target(&candidate.from_state, candidate.target())
11725 .await?
11726 {
11727 drop(main_future);
11729 drop(transition_future);
11730 self.finalize_branch_loss(
11731 &main_id,
11732 RuntimeOptimizationKind::BufferedStreamingRouting,
11733 RuntimeCommitBehavior::FinalResponse,
11734 main_pending,
11735 main_result.as_ref().map(|result| result.is_err()),
11736 );
11737 if !self
11738 .apply_pre_response_transition_candidate(
11739 &candidate,
11740 &HashMap::new(),
11741 processed_input,
11742 )
11743 .await?
11744 {
11745 self.finalize_optional_branch(
11746 &transition_id,
11747 RuntimeOptimizationKind::ParallelStateTransition,
11748 RuntimeCommitBehavior::TransitionDecision,
11749 "discarded",
11750 false,
11751 );
11752 return Ok(None);
11753 }
11754 self.finalize_optional_branch(
11755 &transition_id,
11756 RuntimeOptimizationKind::ParallelStateTransition,
11757 RuntimeCommitBehavior::TransitionDecision,
11758 "committed",
11759 true,
11760 );
11761 let response = self.redispatch_current_state(processed_input).await?;
11762 return Ok(Some((
11763 response.clone(),
11764 vec![StreamChunk::content(response.content)],
11765 )));
11766 }
11767 self.finalize_optional_branch(
11768 &transition_id,
11769 RuntimeOptimizationKind::ParallelStateTransition,
11770 RuntimeCommitBehavior::TransitionDecision,
11771 "discarded",
11772 false,
11773 );
11774 transition_finalized = true;
11775 }
11776 if transition_finalized && !routing_resolved.load(Ordering::SeqCst) {
11782 match self
11783 .resolve_buffered_skill_after_transition(processed_input, &routing_resolved)
11784 .await
11785 {
11786 Ok(Some(candidate)) => {
11787 drop(main_future);
11789 drop(transition_future);
11790 self.finalize_branch_loss(
11791 &main_id,
11792 RuntimeOptimizationKind::BufferedStreamingRouting,
11793 RuntimeCommitBehavior::FinalResponse,
11794 main_pending,
11795 main_result.as_ref().map(|result| result.is_err()),
11796 );
11797 return match self
11798 .commit_winning_skill_candidate(
11799 candidate,
11800 processed_input,
11801 input_context,
11802 )
11803 .await?
11804 {
11805 Some(response) => Ok(Some((
11806 response.clone(),
11807 vec![StreamChunk::content(response.content)],
11808 ))),
11809 None => Ok(None),
11810 };
11811 }
11812 Ok(None) => {}
11813 Err(error) => {
11814 drop(main_future);
11815 drop(transition_future);
11816 self.finalize_branch_loss(
11817 &main_id,
11818 RuntimeOptimizationKind::BufferedStreamingRouting,
11819 RuntimeCommitBehavior::FinalResponse,
11820 main_pending,
11821 main_result.as_ref().map(|result| result.is_err()),
11822 );
11823 return Err(error);
11824 }
11825 }
11826 }
11827 if transition_finalized
11828 && routing_resolved.load(Ordering::SeqCst)
11829 && let Some(result) = main_result.take()
11830 {
11831 let stream_draft = match result {
11832 Ok(stream_draft) => stream_draft,
11833 Err(error) => {
11834 self.finalize_optional_branch(
11835 &main_id,
11836 RuntimeOptimizationKind::BufferedStreamingRouting,
11837 RuntimeCommitBehavior::FinalResponse,
11838 "failed",
11839 false,
11840 );
11841 return Err(error);
11842 }
11843 };
11844 let raw_draft_content = stream_draft.draft.raw_content().to_string();
11845 let buffered_chunks = stream_draft.chunks;
11846 self.finalize_optional_branch(
11847 &main_id,
11848 RuntimeOptimizationKind::BufferedStreamingRouting,
11849 RuntimeCommitBehavior::FinalResponse,
11850 "committed",
11851 true,
11852 );
11853 let response = self
11854 .commit_main_response_draft(
11855 processed_input,
11856 input_context,
11857 stream_draft.draft,
11858 ReasoningMode::None,
11859 false,
11860 )
11861 .await?;
11862 let chunks = if response.content == raw_draft_content {
11863 buffered_chunks
11864 } else {
11865 vec![StreamChunk::content(response.content.clone())]
11866 };
11867 return Ok(Some((response, chunks)));
11868 }
11869 tokio::select! {
11870 result = &mut main_future, if main_pending => {
11871 main_pending = false;
11872 main_branch.transition_to(RuntimeBranchStatus::Completed)?;
11873 main_result = Some(result);
11874 }
11875 result = &mut transition_future, if transition_pending => {
11876 transition_pending = false;
11877 transition_branch.transition_to(RuntimeBranchStatus::Completed)?;
11878 match result {
11879 Ok(ParallelTransitionSelection::Candidate(candidate)) => {
11880 transition_candidate = Some(candidate)
11881 }
11882 Ok(ParallelTransitionSelection::NoMatch) => {
11883 self.finalize_optional_branch(
11884 &transition_id,
11885 RuntimeOptimizationKind::ParallelStateTransition,
11886 RuntimeCommitBehavior::TransitionDecision,
11887 "discarded",
11888 false,
11889 );
11890 transition_finalized = true;
11891 }
11892 Ok(ParallelTransitionSelection::ReservationExhausted) => {
11893 self.finalize_optional_branch(
11894 &transition_id,
11895 RuntimeOptimizationKind::ParallelStateTransition,
11896 RuntimeCommitBehavior::TransitionDecision,
11897 "cancelled",
11898 false,
11899 );
11900 routing_resolved.store(true, Ordering::SeqCst);
11901 self.finalize_branch_loss(
11902 &main_id,
11903 RuntimeOptimizationKind::BufferedStreamingRouting,
11904 RuntimeCommitBehavior::FinalResponse,
11905 main_pending,
11906 main_result.as_ref().map(|result| result.is_err()),
11907 );
11908 return Ok(None);
11909 }
11910 Err(_) => {
11911 self.finalize_optional_branch(
11912 &transition_id,
11913 RuntimeOptimizationKind::ParallelStateTransition,
11914 RuntimeCommitBehavior::TransitionDecision,
11915 "failed",
11916 false,
11917 );
11918 transition_finalized = true;
11919 }
11920 }
11921 }
11922 }
11923 }
11924 }
11925
11926 async fn resolve_buffered_skill_after_transition(
11932 &self,
11933 processed_input: &str,
11934 routing_resolved: &AtomicBool,
11935 ) -> Result<Option<SkillCandidate>> {
11936 let candidate = if self.skill_router.is_some() {
11937 self.select_skill_candidate(processed_input).await?
11938 } else {
11939 None
11940 };
11941 if candidate.is_none() {
11942 routing_resolved.store(true, Ordering::SeqCst);
11943 }
11944 Ok(candidate)
11945 }
11946
11947 fn run_loop_internal_stream<'a>(
11951 &'a self,
11952 input: &'a str,
11953 terminal: RuntimeStreamTerminalSlot,
11954 ) -> Pin<Box<dyn Stream<Item = StreamChunk> + Send + 'a>> {
11955 let include_state_events = self.streaming.include_state_events;
11956
11957 Box::pin(async_stream::stream! {
11958 self.begin_root_turn();
11959 self.pre_turn_session_lifecycle().await;
11961
11962 let input_data = match self.process_input(input).await {
11963 Ok(data) => data,
11964 Err(e) => {
11965 yield StreamChunk::error(e.to_string());
11966 return;
11967 }
11968 };
11969 self.update_active_turn_context(&input_data.content, input_data.context.clone());
11970
11971 for (key, value) in &input_data.context {
11973 let _ = self.context_manager.set(key, value.clone());
11974 }
11975
11976 if input_data.metadata.rejected {
11977 let reason = input_data
11978 .metadata
11979 .rejection_reason
11980 .unwrap_or_else(|| "Input rejected".to_string());
11981 warn!(reason = %reason, "Input rejected (stream)");
11982 let response = AgentResponse::new(&reason);
11985 if let Err(e) = self.finish_turn_if_root(&response).await {
11986 yield StreamChunk::error(e.to_string());
11987 return;
11988 }
11989 yield StreamChunk::content(&reason);
11990 record_runtime_stream_final(&terminal, response);
11991 yield StreamChunk::Done {};
11992 return;
11993 }
11994
11995 let processed_input = &input_data.content;
11996
11997 let streaming_policy = self.runtime_config.optimization.streaming_policy;
11998
11999 if self.runtime_config.optimization.enabled
12006 && !matches!(
12007 streaming_policy,
12008 crate::optimization::StreamingOptimizationPolicy::Disabled
12009 )
12010 {
12011 match self.try_pre_response_transition(processed_input).await {
12012 Ok(Some(response)) => {
12013 yield StreamChunk::content(&response.content);
12014 record_runtime_stream_final(&terminal, response);
12015 yield StreamChunk::Done {};
12016 return;
12017 }
12018 Ok(None) => {}
12019 Err(e) => {
12020 yield StreamChunk::error(e.to_string());
12021 return;
12022 }
12023 }
12024 }
12025
12026 if self.runtime_config.optimization.enabled
12027 && matches!(
12028 streaming_policy,
12029 crate::optimization::StreamingOptimizationPolicy::BufferUntilRoutingDone
12030 )
12031 {
12032 match Box::pin(self.try_buffered_streaming_branches(processed_input, &input_data.context)).await {
12037 Ok(Some((response, chunks))) => {
12038 for chunk in chunks {
12039 yield chunk;
12040 }
12041 record_runtime_stream_final(&terminal, response);
12042 yield StreamChunk::Done {};
12043 return;
12044 }
12045 Ok(None) => {}
12046 Err(e) => {
12047 yield StreamChunk::error(e.to_string());
12048 return;
12049 }
12050 }
12051 }
12052
12053 if let Some(ref sm) = self.state_machine
12055 && let Some(def) = sm.current_definition()
12056 {
12057 let orchestration_result = if let Some(ref delegate_id) = def.delegate {
12058 Some(self.handle_delegated_state(processed_input, delegate_id, &def).await)
12059 } else if let Some(ref concurrent_config) = def.concurrent {
12060 Some(self.handle_concurrent_state(processed_input, concurrent_config).await)
12061 } else if let Some(ref group_chat_config) = def.group_chat {
12062 Some(self.handle_group_chat_state(processed_input, group_chat_config).await)
12063 } else if let Some(ref pipeline_config) = def.pipeline {
12064 Some(self.handle_pipeline_state(processed_input, pipeline_config).await)
12065 } else if let Some(ref handoff_config) = def.handoff {
12066 Some(self.handle_handoff_state(processed_input, handoff_config).await)
12067 } else {
12068 None
12069 };
12070
12071 if let Some(result) = orchestration_result {
12072 match result {
12073 Ok(response) => {
12074 yield StreamChunk::content(&response.content);
12075 record_runtime_stream_final(&terminal, response);
12076 yield StreamChunk::Done {};
12077 }
12078 Err(e) => {
12079 yield StreamChunk::error(e.to_string());
12080 }
12081 }
12082 return;
12083 }
12084 }
12085
12086 match self.try_skill_route(processed_input).await {
12088 Ok(SkillRouteResult::Response { skill_id, content }) => {
12089 if let Err(e) = self.commit_root_user_message(processed_input).await {
12090 yield StreamChunk::error(e.to_string());
12091 return;
12092 }
12093 match self.handle_skill_response(processed_input, &skill_id, content, &input_data.context).await {
12094 Ok(resp) => {
12095 yield StreamChunk::content(&resp.content);
12096 record_runtime_stream_final(&terminal, resp);
12097 yield StreamChunk::Done {};
12098 return;
12099 }
12100 Err(e) => {
12101 yield StreamChunk::error(e.to_string());
12102 return;
12103 }
12104 }
12105 }
12106 Ok(SkillRouteResult::NeedsClarification {
12107 response,
12108 ownership,
12109 }) => {
12110 let admission = match self
12111 .admit_optional_disambiguation_ownership(ownership)
12112 .await
12113 {
12114 Ok(admission) => admission,
12115 Err(e) => {
12116 yield StreamChunk::error(e.to_string());
12117 return;
12118 }
12119 };
12120 if let Err(e) = self.commit_root_user_message(processed_input).await {
12121 yield StreamChunk::error(e.to_string());
12122 return;
12123 }
12124 if Self::skill_clarification_needs_memory_record(&response)
12126 && let Err(e) = self.memory.add_message(ChatMessage::assistant(&response.content)).await
12127 {
12128 yield StreamChunk::error(e.to_string());
12129 return;
12130 }
12131 drop(admission);
12132 if let Err(e) = self.finish_turn_if_root(&response).await {
12133 yield StreamChunk::error(e.to_string());
12134 return;
12135 }
12136 yield StreamChunk::content(&response.content);
12137 record_runtime_stream_final(&terminal, response);
12138 yield StreamChunk::Done {};
12139 return;
12140 }
12141 Ok(SkillRouteResult::NoMatch) => {} Err(e) => {
12143 yield StreamChunk::error(e.to_string());
12144 return;
12145 }
12146 }
12147
12148 let effective_reasoning = self.get_effective_reasoning_config();
12150 let reasoning_mode = match self.determine_reasoning_mode(processed_input).await {
12151 Ok(mode) => mode,
12152 Err(e) => {
12153 yield StreamChunk::error(e.to_string());
12154 return;
12155 }
12156 };
12157 let auto_detected = matches!(effective_reasoning.mode, ReasoningMode::Auto);
12158
12159 info!(
12160 reasoning_mode = ?reasoning_mode,
12161 auto_detected = auto_detected,
12162 "Reasoning mode determined (stream)"
12163 );
12164
12165 if matches!(reasoning_mode, ReasoningMode::PlanAndExecute) {
12167 if let Err(e) = self.commit_root_user_message(processed_input).await {
12168 yield StreamChunk::error(e.to_string());
12169 return;
12170 }
12171 match self.handle_plan_and_execute(processed_input, &input_data.context, auto_detected).await {
12172 Ok(resp) => {
12173 yield StreamChunk::content(&resp.content);
12174 record_runtime_stream_final(&terminal, resp);
12175 yield StreamChunk::Done {};
12176 return;
12177 }
12178 Err(e) => {
12179 yield StreamChunk::error(e.to_string());
12180 return;
12181 }
12182 }
12183 }
12184
12185 if let Err(e) = self.commit_root_user_message(processed_input).await {
12186 yield StreamChunk::error(e.to_string());
12187 return;
12188 }
12189
12190 let llm = match self.get_state_llm() {
12191 Ok(llm) => llm,
12192 Err(e) => {
12193 yield StreamChunk::error(e.to_string());
12194 return;
12195 }
12196 };
12197
12198 let mut iterations = 0u32;
12199 let mut all_tool_calls: Vec<ToolCall> = Vec::new();
12200 let mut thinking_content: Option<String> = None;
12201
12202 loop {
12203 let effective_max = if reasoning_mode != ReasoningMode::None {
12205 let rc = self.get_effective_reasoning_config();
12206 self.max_iterations.min(rc.max_iterations)
12207 } else {
12208 self.max_iterations
12209 };
12210
12211 if iterations >= effective_max {
12212 let err_msg = format!("Max iterations ({}) exceeded", effective_max);
12213 let err = AgentError::Other(err_msg.clone());
12214 self.hooks.on_error(&err).await;
12215 error!(iterations = iterations, "Max iterations exceeded (stream)");
12216 yield StreamChunk::error(err_msg);
12217 return;
12218 }
12219 iterations += 1;
12220 *self.iteration_count.write() = iterations;
12221
12222 debug!(iteration = iterations, max = effective_max, "LLM call (stream)");
12223
12224 let protocol = match self.main_tool_protocol(llm.as_ref(), false).await {
12225 Ok(protocol) => protocol,
12226 Err(e) => {
12227 yield StreamChunk::error(e.to_string());
12228 return;
12229 }
12230 };
12231 let mut messages = match self
12232 .build_messages_internal(true, None, protocol.choice.is_none())
12233 .await
12234 {
12235 Ok(m) => m,
12236 Err(e) => {
12237 yield StreamChunk::error(e.to_string());
12238 return;
12239 }
12240 };
12241 self.inject_reasoning_prompt(&mut messages, &reasoning_mode, iterations == 1);
12242
12243 self.hooks.on_llm_start(&messages).await;
12244 let llm_start = Instant::now();
12245
12246 let reflection_active = self
12249 .should_reflect(processed_input, "")
12250 .await
12251 .unwrap_or_default();
12252
12253 let buffered_decision = reflection_active || protocol.choice.is_some();
12254 let content = if buffered_decision {
12255 let response = match self
12259 .complete_main_llm_with_recovery(
12260 Arc::clone(&llm),
12261 &messages,
12262 &protocol,
12263 )
12264 .await
12265 {
12266 Ok(r) => r,
12267 Err(e) => {
12268 yield StreamChunk::error(e.to_string());
12269 return;
12270 }
12271 };
12272 let llm_duration_ms = llm_start.elapsed().as_millis() as u64;
12273 self.hooks.on_llm_complete(&response, llm_duration_ms).await;
12274 response.content.trim().to_string()
12275 } else {
12276 let source = match self
12278 .open_main_stream_with_recovery(Arc::clone(&llm), &messages, &protocol)
12279 .await
12280 {
12281 Ok(source) => source,
12282 Err(e) => {
12283 yield StreamChunk::error(e.to_string());
12284 return;
12285 }
12286 };
12287 let mut accumulated = String::new();
12288 match source {
12289 MainStreamSource::StaticResponse(text) => {
12290 accumulated.push_str(&text);
12291 yield StreamChunk::content(text);
12292 }
12293 MainStreamSource::Stream(mut stream_inner) => {
12294 while let Some(chunk_result) = stream_inner.next().await {
12295 match chunk_result {
12296 Ok(chunk) => {
12297 accumulated.push_str(&chunk.delta);
12298 yield StreamChunk::content(chunk.delta);
12299 }
12300 Err(e) => {
12301 yield StreamChunk::error(e.to_string());
12303 return;
12304 }
12305 }
12306 }
12307 }
12308 }
12309 let llm_duration_ms = llm_start.elapsed().as_millis() as u64;
12310 let llm_response = ai_agents_core::LLMResponse::new(
12312 accumulated.trim(),
12313 ai_agents_core::FinishReason::Stop,
12314 );
12315 self.hooks.on_llm_complete(&llm_response, llm_duration_ms).await;
12316 accumulated.trim().to_string()
12317 };
12318
12319 let parsed_tool_calls = match self.parse_main_tool_calls(&content, &protocol) {
12321 Ok(calls) => calls,
12322 Err(error) => {
12323 yield StreamChunk::error(error.to_string());
12324 return;
12325 }
12326 };
12327 if let Some(tool_calls) = parsed_tool_calls {
12328 let mut events = Vec::new();
12331 let outcome = self
12332 .handle_tool_calls(
12333 processed_input,
12334 &content,
12335 tool_calls,
12336 &mut all_tool_calls,
12337 Some(&mut events),
12338 )
12339 .await;
12340 for chunk in events.drain(..) {
12341 yield chunk;
12342 }
12343 match outcome {
12344 Ok(ToolCallOutcome::Continue) | Ok(ToolCallOutcome::TransitionFired) => continue,
12345 Ok(ToolCallOutcome::Rejected(response)) => {
12346 if let Err(finalize_error) = self.finish_turn_if_root(&response).await {
12347 yield StreamChunk::error(finalize_error.to_string());
12348 return;
12349 }
12350 let legacy_error = response.content.clone();
12351 record_runtime_stream_final(&terminal, response);
12352 yield StreamChunk::error(legacy_error);
12353 yield StreamChunk::Done {};
12354 return;
12355 }
12356 Err(e) => {
12357 yield StreamChunk::error(e.to_string());
12358 return;
12359 }
12360 }
12361 }
12362
12363 let (extracted_thinking, answer) = self.extract_thinking(&content);
12365 if extracted_thinking.is_some() {
12366 thinking_content = extracted_thinking;
12367 }
12368
12369 let output_data = match self.process_output(&answer, &input_data.context).await {
12370 Ok(d) => d,
12371 Err(e) => {
12372 yield StreamChunk::error(e.to_string());
12373 return;
12374 }
12375 };
12376
12377 let final_content = if output_data.metadata.rejected {
12378 output_data
12379 .metadata
12380 .rejection_reason
12381 .unwrap_or_else(|| answer.to_string())
12382 } else {
12383 output_data.content
12384 };
12385
12386 let (final_content, reflection_metadata) = match self
12388 .run_reflection(&*llm, processed_input, final_content)
12389 .await
12390 {
12391 Ok(r) => r,
12392 Err(e) => {
12393 yield StreamChunk::error(e.to_string());
12394 return;
12395 }
12396 };
12397
12398 let final_content = self.format_response_with_thinking(
12399 thinking_content.as_deref(),
12400 &final_content,
12401 );
12402
12403 if buffered_decision {
12405 yield StreamChunk::content(&final_content);
12406 }
12407
12408 let post_result = match self
12412 .post_loop_processing(processed_input, final_content)
12413 .await
12414 {
12415 Ok(r) => r,
12416 Err(e) => {
12417 yield StreamChunk::error(e.to_string());
12418 return;
12419 }
12420 };
12421
12422 let applied = match self.apply_post_loop_result(processed_input, post_result).await {
12423 Ok(applied) => applied,
12424 Err(e) => {
12425 yield StreamChunk::error(e.to_string());
12426 return;
12427 }
12428 };
12429
12430 if applied.transitioned {
12431 if include_state_events
12432 && let Some(state) = self.current_state()
12433 {
12434 yield StreamChunk::state_transition(None, state);
12435 }
12436 if applied.regenerated {
12442 yield StreamChunk::content(&applied.content);
12443 }
12444 }
12445 let final_content = applied.content;
12446
12447 let final_response = self.build_agent_response(AgentResponseParts {
12449 content: final_content,
12450 all_tool_calls,
12451 reasoning_mode,
12452 auto_detected,
12453 iterations,
12454 thinking: thinking_content,
12455 reflection_metadata,
12456 });
12457 if let Err(e) = self.finish_turn_if_root(&final_response).await {
12458 yield StreamChunk::error(e.to_string());
12459 return;
12460 }
12461
12462 record_runtime_stream_final(&terminal, final_response);
12463 yield StreamChunk::Done {};
12464 return;
12465 }
12466 })
12467 }
12468
12469 fn run_loop_stream<'a>(
12472 &'a self,
12473 input: &'a str,
12474 terminal: RuntimeStreamTerminalSlot,
12475 ) -> Pin<Box<dyn Stream<Item = StreamChunk> + Send + 'a>> {
12476 Box::pin(async_stream::stream! {
12477 self.begin_root_turn();
12478 let _root_cleanup = RootTurnCleanup::new(self);
12479 self.hooks.on_message_received(input).await;
12480
12481 if !self.context_initialized.swap(true, Ordering::SeqCst) {
12483 if let Err(e) = self.context_manager.initialize().await {
12484 yield StreamChunk::error(e.to_string());
12485 return;
12486 }
12487 debug!("Context manager initialized (defaults, env, builtins)");
12488 }
12489
12490 if let Err(e) = self.check_turn_timeout().await {
12491 yield StreamChunk::error(e.to_string());
12492 return;
12493 }
12494 if let Err(e) = self.context_manager.refresh_per_turn().await {
12495 yield StreamChunk::error(e.to_string());
12496 return;
12497 }
12498
12499 self.clear_disambiguation_context();
12501
12502 let input_to_run = match self.resolve_disambiguation(input).await {
12506 Err(e) => {
12507 yield StreamChunk::error(e.to_string());
12508 return;
12509 }
12510 Ok(DisambiguationDispatch::Terminal(response)) => {
12511 yield StreamChunk::content(&response.content);
12512 record_runtime_stream_final(&terminal, response);
12513 yield StreamChunk::Done {};
12514 return;
12515 }
12516 Ok(DisambiguationDispatch::RecheckSkill {
12517 skill_id,
12518 enriched_input,
12519 disambiguation_epoch,
12520 state_generation,
12521 }) => {
12522 match self
12523 .recheck_skill_disambiguation(
12524 &skill_id,
12525 &enriched_input,
12526 disambiguation_epoch,
12527 state_generation,
12528 )
12529 .await
12530 {
12531 Ok(resp) => {
12532 yield StreamChunk::content(&resp.content);
12533 record_runtime_stream_final(&terminal, resp);
12534 yield StreamChunk::Done {};
12535 return;
12536 }
12537 Err(e) => {
12538 yield StreamChunk::error(e.to_string());
12539 return;
12540 }
12541 }
12542 }
12543 Ok(DisambiguationDispatch::Proceed(input)) => input,
12544 };
12545
12546 let mut inner = self.run_loop_internal_stream(&input_to_run, Arc::clone(&terminal));
12547 while let Some(chunk) = inner.next().await {
12548 yield chunk;
12549 }
12550 })
12551 }
12552
12553 pub fn info(&self) -> AgentInfo {
12554 self.info.clone()
12555 }
12556
12557 pub fn skills(&self) -> &[SkillDefinition] {
12558 &self.skills
12559 }
12560
12561 async fn reset_runtime_state(&self) -> Result<()> {
12563 let _admission = self.disambiguation_admission.write().await;
12564 if self.state_transition_reserved.load(Ordering::SeqCst) {
12565 return Err(AgentError::Other(
12566 "Cannot reset while a state transition is in progress".to_string(),
12567 ));
12568 }
12569 self.disambiguation_epoch.fetch_add(1, Ordering::SeqCst);
12570 *self.pending_skill_id.write() = None;
12571 if let Some(disambiguator) = self.disambiguation_manager.as_ref() {
12572 disambiguator.clear_pending().await;
12573 }
12574 self.memory.clear().await?;
12575 self.active_native_exchanges.write().clear();
12576 *self.iteration_count.write() = 0;
12577 self.tool_call_history.write().clear();
12578 if let Some(ref sm) = self.state_machine {
12579 sm.reset();
12580 }
12581 Ok(())
12582 }
12583
12584 pub async fn reset(&self) -> Result<()> {
12586 self.reset_runtime_state().await
12587 }
12588
12589 pub fn max_context_tokens(&self) -> u32 {
12590 self.max_context_tokens
12591 }
12592
12593 pub fn llm_registry(&self) -> &Arc<LLMRegistry> {
12594 &self.llm_registry
12595 }
12596
12597 pub fn state_machine(&self) -> Option<&Arc<StateMachine>> {
12598 self.state_machine.as_ref()
12599 }
12600
12601 pub fn context_manager(&self) -> &Arc<ContextManager> {
12602 &self.context_manager
12603 }
12604
12605 pub fn tool_call_history(&self) -> Vec<ToolCallRecord> {
12606 self.tool_call_history.read().clone()
12607 }
12608
12609 pub fn memory_token_budget(&self) -> Option<&MemoryTokenBudget> {
12610 self.memory_token_budget.as_ref()
12611 }
12612
12613 pub fn parallel_tools_config(&self) -> &ParallelToolsConfig {
12614 &self.parallel_tools
12615 }
12616
12617 pub fn streaming_config(&self) -> &StreamingConfig {
12618 &self.streaming
12619 }
12620
12621 pub fn hooks(&self) -> &Arc<dyn AgentHooks> {
12622 &self.hooks
12623 }
12624
12625 pub fn hitl_engine(&self) -> Option<&HITLEngine> {
12626 self.hitl_engine.as_ref()
12627 }
12628
12629 pub fn approval_handler(&self) -> &Arc<dyn ApprovalHandler> {
12630 &self.approval_handler
12631 }
12632
12633 fn build_hitl_language_context(&self) -> HashMap<String, Value> {
12635 let mut ctx = HashMap::new();
12636 for key in &["user.language", "input.detected.language", "language"] {
12637 if let Some(val) = self.context_manager.get(key) {
12638 ctx.insert(key.to_string(), val);
12639 }
12640 }
12641 ctx
12642 }
12643
12644 async fn request_hitl_approval(&self, check_result: HITLCheckResult) -> Result<ApprovalResult> {
12646 let Some(request) = check_result.into_request() else {
12647 return Ok(ApprovalResult::Approved);
12648 };
12649
12650 self.hooks.on_approval_requested(&request).await;
12651
12652 let timeout = request.timeout;
12653
12654 let raw_result = if let Some(duration) = timeout {
12655 match tokio::time::timeout(
12656 duration,
12657 self.approval_handler.request_approval(request.clone()),
12658 )
12659 .await
12660 {
12661 Ok(result) => result,
12662 Err(_) => ApprovalResult::timeout(),
12663 }
12664 } else {
12665 self.approval_handler
12666 .request_approval(request.clone())
12667 .await
12668 };
12669
12670 self.hooks
12671 .on_approval_result(&request.id, &raw_result)
12672 .await;
12673
12674 let (outcome, effective_result): (ApprovalResolvedOutcome, Result<ApprovalResult>) =
12675 match &raw_result {
12676 ApprovalResult::Approved => (
12677 ApprovalResolvedOutcome::Approved,
12678 Ok(ApprovalResult::Approved),
12679 ),
12680 ApprovalResult::Rejected { reason } => (
12681 ApprovalResolvedOutcome::Rejected {
12682 reason: reason.clone(),
12683 },
12684 Ok(ApprovalResult::Rejected {
12685 reason: reason.clone(),
12686 }),
12687 ),
12688 ApprovalResult::Modified { changes } => (
12689 ApprovalResolvedOutcome::Modified {
12690 changes: changes.clone(),
12691 },
12692 Ok(ApprovalResult::Modified {
12693 changes: changes.clone(),
12694 }),
12695 ),
12696 ApprovalResult::Timeout => {
12697 if let Some(ref engine) = self.hitl_engine {
12698 match engine.config().on_timeout {
12699 TimeoutAction::Approve => (
12700 ApprovalResolvedOutcome::Approved,
12701 Ok(ApprovalResult::Approved),
12702 ),
12703 TimeoutAction::Reject => {
12704 let reason = Some("Timeout".to_string());
12705 (
12706 ApprovalResolvedOutcome::Rejected {
12707 reason: reason.clone(),
12708 },
12709 Ok(ApprovalResult::Rejected { reason }),
12710 )
12711 }
12712 TimeoutAction::Error => {
12713 let message = "HITL approval timeout".to_string();
12714 (
12715 ApprovalResolvedOutcome::Error {
12716 message: message.clone(),
12717 },
12718 Err(AgentError::Other(message)),
12719 )
12720 }
12721 }
12722 } else {
12723 let reason = Some("Timeout (no engine)".to_string());
12724 (
12725 ApprovalResolvedOutcome::Rejected {
12726 reason: reason.clone(),
12727 },
12728 Ok(ApprovalResult::Rejected { reason }),
12729 )
12730 }
12731 }
12732 };
12733
12734 self.hooks
12735 .on_approval_resolved(&request, &raw_result, &outcome)
12736 .await;
12737
12738 effective_result
12739 }
12740
12741 pub async fn check_state_hitl(&self, from: Option<&str>, to: &str) -> Result<bool> {
12742 if let Some(ref hitl_engine) = self.hitl_engine {
12743 let hitl_lang_ctx = self.build_hitl_language_context();
12744 let check_result = self
12745 .observe_purpose(
12746 ObservationPurpose::HitlLocalization,
12747 hitl_engine.check_state_transition_with_localization(
12748 from,
12749 to,
12750 &hitl_lang_ctx,
12751 self.approval_handler.as_ref(),
12752 Some(&self.llm_registry),
12753 ),
12754 )
12755 .await?;
12756 if check_result.is_required() {
12757 let result = self.request_hitl_approval(check_result).await?;
12758 return Ok(matches!(
12759 result,
12760 ApprovalResult::Approved | ApprovalResult::Modified { .. }
12761 ));
12762 }
12763 }
12764 Ok(true)
12765 }
12766
12767 async fn execute_tools_parallel(
12769 &self,
12770 tool_calls: &[ToolCall],
12771 ) -> Vec<(String, Result<String>)> {
12772 let can_run_parallel = tool_calls.iter().all(|tc| {
12773 self.tools
12774 .resolve(&tc.name)
12775 .map(|resolved| resolved.tool.classify_call(&tc.arguments).concurrency_safe)
12776 .unwrap_or(false)
12777 });
12778
12779 if !self.parallel_tools.enabled || tool_calls.len() <= 1 || !can_run_parallel {
12780 let mut results = Vec::new();
12781 for tc in tool_calls {
12782 let result = self
12783 .observe_purpose(
12784 current_observation_context()
12785 .map(|context| context.purpose)
12786 .unwrap_or_default(),
12787 self.execute_tool_smart(tc),
12788 )
12789 .await;
12790 results.push((tc.id.clone(), result));
12791 }
12792 return results;
12793 }
12794
12795 let chunks: Vec<_> = tool_calls
12796 .chunks(self.parallel_tools.max_parallel)
12797 .collect();
12798
12799 let mut all_results = Vec::new();
12800
12801 for chunk in chunks {
12802 let futures: Vec<_> = chunk
12803 .iter()
12804 .map(|tc| {
12805 let tc = tc.clone();
12806 async move {
12807 let result = self.execute_tool_smart(&tc).await;
12808 (tc.id.clone(), result)
12809 }
12810 })
12811 .collect();
12812
12813 let results = futures::future::join_all(futures).await;
12814 all_results.extend(results);
12815 }
12816
12817 all_results
12818 }
12819
12820 pub async fn chat_stream<'a>(
12824 &'a self,
12825 input: &'a str,
12826 ) -> Result<Pin<Box<dyn Stream<Item = StreamChunk> + Send + 'a>>> {
12827 let RootTurnAdmission {
12828 guard: root_turn_guard,
12829 identity_stack,
12830 } = self.acquire_root_turn().await?;
12831 scope_runtime_gate_identity_stack(&identity_stack, self.init_storage()).await?;
12835 info!(input_len = input.len(), "Starting streaming chat");
12836 let terminal = new_runtime_stream_terminal_slot();
12837 let inner = self.run_loop_stream(input, terminal);
12838 let observation_context = self.build_observation_context(None);
12839 let stream: Pin<Box<dyn Stream<Item = StreamChunk> + Send + 'a>> =
12840 Box::pin(async_stream::stream! {
12841 let mut root_turn_guard = Some(root_turn_guard);
12842 let mut inner = inner;
12843 loop {
12844 let next = scope_runtime_gate_identity_stack(&identity_stack, async {
12845 if let Some(context) = observation_context.as_ref() {
12846 with_observation_context(context.clone(), inner.next()).await
12847 } else {
12848 inner.next().await
12849 }
12850 })
12851 .await;
12852 match next {
12853 Some(StreamChunk::Done {}) => {
12854 while scope_runtime_gate_identity_stack(&identity_stack, async {
12855 if let Some(context) = observation_context.as_ref() {
12856 with_observation_context(context.clone(), inner.next())
12857 .await
12858 .is_some()
12859 } else {
12860 inner.next().await.is_some()
12861 }
12862 })
12863 .await
12864 {}
12865 if observation_context.is_some() {
12866 scope_runtime_gate_identity_stack(
12867 &identity_stack,
12868 self.export_observability_if_configured(),
12869 )
12870 .await;
12871 }
12872 drop(root_turn_guard.take());
12873 yield StreamChunk::Done {};
12874 return;
12875 }
12876 Some(chunk) => yield chunk,
12877 None => {
12878 if observation_context.is_some() {
12879 scope_runtime_gate_identity_stack(
12880 &identity_stack,
12881 self.export_observability_if_configured(),
12882 )
12883 .await;
12884 }
12885 drop(root_turn_guard.take());
12886 return;
12887 }
12888 }
12889 }
12890 });
12891 Ok(stream)
12892 }
12893
12894 pub async fn chat_stream_events<'a>(
12898 &'a self,
12899 input: &'a str,
12900 ) -> Result<Pin<Box<dyn Stream<Item = AgentStreamEvent> + Send + 'a>>> {
12901 let RootTurnAdmission {
12902 guard: root_turn_guard,
12903 identity_stack,
12904 } = self.acquire_root_turn().await?;
12905 scope_runtime_gate_identity_stack(&identity_stack, self.init_storage()).await?;
12909 info!(input_len = input.len(), "Starting streaming chat events");
12910 let terminal = new_runtime_stream_terminal_slot();
12911 let mut inner = self.run_loop_stream(input, Arc::clone(&terminal));
12912 let observation_context = self.build_observation_context(None);
12913 let stream: Pin<Box<dyn Stream<Item = AgentStreamEvent> + Send + 'a>> =
12914 Box::pin(async_stream::stream! {
12915 let mut root_turn_guard = Some(root_turn_guard);
12916 loop {
12917 let next = scope_runtime_gate_identity_stack(&identity_stack, async {
12918 if let Some(context) = observation_context.as_ref() {
12919 with_observation_context(context.clone(), inner.next()).await
12920 } else {
12921 inner.next().await
12922 }
12923 })
12924 .await;
12925 match next {
12926 Some(StreamChunk::Done {}) => {
12927 let terminal_event = { terminal.write().take() };
12928 if let Some(response) = terminal_event {
12929 while scope_runtime_gate_identity_stack(&identity_stack, async {
12930 if let Some(context) = observation_context.as_ref() {
12931 with_observation_context(context.clone(), inner.next())
12932 .await
12933 .is_some()
12934 } else {
12935 inner.next().await.is_some()
12936 }
12937 })
12938 .await
12939 {}
12940 if observation_context.is_some() {
12941 scope_runtime_gate_identity_stack(
12942 &identity_stack,
12943 self.export_observability_if_configured(),
12944 )
12945 .await;
12946 }
12947 drop(root_turn_guard.take());
12948 yield AgentStreamEvent::Final(response);
12949 return;
12950 }
12951 }
12952 Some(StreamChunk::Error { message }) => {
12953 let finalized = { terminal.read().is_some() };
12954 if finalized {
12955 continue;
12956 }
12957 while scope_runtime_gate_identity_stack(&identity_stack, async {
12958 if let Some(context) = observation_context.as_ref() {
12959 with_observation_context(context.clone(), inner.next())
12960 .await
12961 .is_some()
12962 } else {
12963 inner.next().await.is_some()
12964 }
12965 })
12966 .await
12967 {}
12968 if observation_context.is_some() {
12969 scope_runtime_gate_identity_stack(
12970 &identity_stack,
12971 self.export_observability_if_configured(),
12972 )
12973 .await;
12974 }
12975 drop(root_turn_guard.take());
12976 yield AgentStreamEvent::Chunk(StreamChunk::Error { message });
12977 return;
12978 }
12979 Some(chunk) => yield AgentStreamEvent::Chunk(chunk),
12980 None => {
12981 if observation_context.is_some() {
12982 scope_runtime_gate_identity_stack(
12983 &identity_stack,
12984 self.export_observability_if_configured(),
12985 )
12986 .await;
12987 }
12988 drop(root_turn_guard.take());
12989 return;
12990 }
12991 }
12992 }
12993 });
12994 Ok(stream)
12995 }
12996}
12997
12998#[async_trait]
12999impl ToolInvoker for RuntimeAgent {
13000 async fn invoke_tool(&self, request: ToolExecutionRequest) -> Result<ToolExecutionRecord> {
13001 self.execute_tool_record(request).await
13002 }
13003}
13004
13005#[async_trait]
13006impl Agent for RuntimeAgent {
13007 async fn chat(&self, input: &str) -> Result<AgentResponse> {
13009 let RootTurnAdmission {
13010 guard,
13011 identity_stack,
13012 } = self.acquire_root_turn().await?;
13013 let result = scope_runtime_gate_identity_stack(&identity_stack, async {
13014 let result = if let Some(context) = self.build_observation_context(None) {
13015 with_observation_context(context, self.run_loop(input)).await
13016 } else {
13017 self.run_loop(input).await
13018 };
13019 self.export_observability_if_configured().await;
13020 result
13021 })
13022 .await;
13023 drop(guard);
13024 result
13025 }
13026
13027 fn info(&self) -> AgentInfo {
13028 self.info.clone()
13029 }
13030
13031 async fn reset(&self) -> Result<()> {
13033 self.reset_runtime_state().await
13034 }
13035}
13036
13037fn background_maintenance_tags(
13047 label: &str,
13048 stage: &str,
13049 reason: Option<&str>,
13050 policy: Option<&crate::optimization::config::MaintenanceTaskPolicy>,
13051) -> HashMap<String, String> {
13052 let mut tags = HashMap::new();
13053 tags.insert("runtime.background".to_string(), "true".to_string());
13054 tags.insert("runtime.maintenance".to_string(), label.to_string());
13055 tags.insert("runtime.maintenance_stage".to_string(), stage.to_string());
13056 if let Some(policy) = policy {
13057 tags.insert(
13058 "runtime.await_before_next_turn".to_string(),
13059 await_before_next_turn_label(policy.await_before_next_turn).to_string(),
13060 );
13061 tags.insert(
13062 "runtime.maintenance_mode".to_string(),
13063 maintenance_mode_label(policy.mode).to_string(),
13064 );
13065 }
13066 if let Some(reason) = reason {
13067 tags.insert("runtime.reason".to_string(), reason.to_string());
13068 }
13069 tags
13070}
13071
13072fn await_before_next_turn_label(policy: AwaitBeforeNextTurn) -> &'static str {
13073 match policy {
13074 AwaitBeforeNextTurn::Never => "never",
13075 AwaitBeforeNextTurn::SameActor => "same_actor",
13076 AwaitBeforeNextTurn::Always => "always",
13077 }
13078}
13079
13080fn maintenance_mode_label(mode: MaintenanceMode) -> &'static str {
13081 match mode {
13082 MaintenanceMode::InlineSerial => "inline_serial",
13083 MaintenanceMode::InlineParallel => "inline_parallel",
13084 MaintenanceMode::Background => "background",
13085 }
13086}
13087
13088fn record_background_maintenance_event(
13090 manager: Option<&Arc<ObservabilityManager>>,
13091 label: &str,
13092 status: EventStatus,
13093 duration_ms: u64,
13094 stage: &str,
13095 reason: Option<String>,
13096 policy: Option<&crate::optimization::config::MaintenanceTaskPolicy>,
13097) {
13098 if let Some(manager) = manager {
13099 manager.record_lifecycle_event(
13100 EventType::MemoryOperation {
13101 operation: format!("{}_background_{}", label, stage),
13102 },
13103 ObservationPurpose::Other(format!("{}_maintenance", label)),
13104 status,
13105 duration_ms,
13106 background_maintenance_tags(label, stage, reason.as_deref(), policy),
13107 None,
13108 );
13109 }
13110}
13111
13112fn effective_maintenance_mode(mode: MaintenanceMode, force_parallel: bool) -> MaintenanceMode {
13113 if force_parallel && matches!(mode, MaintenanceMode::InlineSerial) {
13114 MaintenanceMode::InlineParallel
13115 } else {
13116 mode
13117 }
13118}
13119
13120fn observation_purpose_for_process(hint: ProcessPurposeHint) -> ObservationPurpose {
13121 match hint {
13122 ProcessPurposeHint::Detect => ObservationPurpose::ProcessDetect,
13123 ProcessPurposeHint::Extract => ObservationPurpose::ProcessExtract,
13124 ProcessPurposeHint::Validate => ObservationPurpose::ProcessValidate,
13125 ProcessPurposeHint::Transform | ProcessPurposeHint::Other => {
13126 ObservationPurpose::ProcessTransform
13127 }
13128 }
13129}
13130
13131fn new_tool_resource_locks() -> ToolResourceLocks {
13132 Arc::new(RwLock::new(HashMap::new()))
13133}
13134
13135fn tool_resource_lock_keys(
13140 _canonical_id: &str,
13141 args: &Value,
13142 bindings: &ai_agents_core::ToolPolicyBindings,
13143 classification: &ai_agents_core::ToolCallClassification,
13144) -> Vec<String> {
13145 if classification.concurrency_safe {
13146 return Vec::new();
13147 }
13148
13149 let mut keys = Vec::new();
13150 let mut has_path_resource = false;
13151 for binding in &bindings.path_fields {
13152 let value = value_at_argument_path(args, &binding.field)
13153 .cloned()
13154 .or_else(|| {
13155 binding
13156 .default_path
13157 .as_ref()
13158 .map(|path| Value::String(path.clone()))
13159 });
13160 if let Some(value) = value {
13161 collect_resource_strings(&value, |_| {
13162 has_path_resource = true;
13163 });
13164 }
13165 }
13166 for binding in &bindings.domain_fields {
13167 if let Some(value) = value_at_argument_path(args, &binding.field) {
13168 collect_resource_strings(value, |domain| {
13169 let normalized = if binding.is_url {
13170 normalized_url_resource_key(domain)
13171 } else {
13172 domain.trim().trim_end_matches('.').to_ascii_lowercase()
13173 };
13174 keys.push(format!("domain:{}", normalized));
13175 });
13176 }
13177 }
13178 for binding in &bindings.command_fields {
13179 if !matches!(binding.kind, ai_agents_core::CommandBindingKind::Cwd) {
13180 continue;
13181 }
13182 if let Some(value) = value_at_argument_path(args, &binding.field) {
13183 collect_resource_strings(value, |_| {
13184 has_path_resource = true;
13185 });
13186 }
13187 }
13188 if has_path_resource {
13189 keys.push("path-mutation:global".to_string());
13190 }
13191 if keys.is_empty() {
13192 keys.push("side-effect:unbound".to_string());
13193 }
13194 keys.sort();
13195 keys.dedup();
13196 keys
13197}
13198
13199fn value_at_argument_path<'a>(value: &'a Value, field: &str) -> Option<&'a Value> {
13200 let mut current = value;
13201 for segment in field.split('.') {
13202 if segment.is_empty() {
13203 return None;
13204 }
13205 current = current.get(segment)?;
13206 }
13207 Some(current)
13208}
13209
13210fn collect_resource_strings(value: &Value, mut collect: impl FnMut(&str)) {
13211 match value {
13212 Value::String(value) => collect(value),
13213 Value::Array(values) => {
13214 for value in values {
13215 if let Some(value) = value.as_str() {
13216 collect(value);
13217 }
13218 }
13219 }
13220 _ => {}
13221 }
13222}
13223
13224fn normalized_url_resource_key(value: &str) -> String {
13225 let value = value.trim();
13226 let Some((scheme, remainder)) = value.split_once("://") else {
13227 return value.to_ascii_lowercase();
13228 };
13229 let authority_end = remainder.find(['/', '?', '#']).unwrap_or(remainder.len());
13230 let (authority, suffix) = remainder.split_at(authority_end);
13231 format!(
13232 "{}://{}{}",
13233 scheme.to_ascii_lowercase(),
13234 authority.to_ascii_lowercase(),
13235 suffix
13236 )
13237}
13238
13239fn render_concurrent_template(
13240 template: &str,
13241 user_input: &str,
13242 context_values: &std::collections::HashMap<String, serde_json::Value>,
13243) -> Result<String> {
13244 let mut env = minijinja::Environment::new();
13245 env.add_template("concurrent", template)
13246 .map_err(|e| AgentError::Other(format!("Concurrent template parse error: {}", e)))?;
13247
13248 let mut ctx = std::collections::BTreeMap::new();
13249 ctx.insert("user_input".to_string(), minijinja::Value::from(user_input));
13250
13251 let context_obj = minijinja::Value::from_serialize(context_values);
13253 ctx.insert("context".to_string(), context_obj);
13254
13255 let tmpl = env
13256 .get_template("concurrent")
13257 .map_err(|e| AgentError::Other(format!("Concurrent template error: {}", e)))?;
13258
13259 tmpl.render(minijinja::Value::from_serialize(&ctx))
13260 .map_err(|e| AgentError::Other(format!("Concurrent template render error: {}", e)))
13261}
13262
13263#[cfg(test)]
13264mod tests {
13265 use super::*;
13266 use crate::AgentBuilder;
13267 use ai_agents_core::{LLMChunk, LLMConfig, LLMError, LLMFeature, Tool};
13268 use ai_agents_llm::mock::MockLLMProvider;
13269 use ai_agents_skills::{SkillDefinition, SkillStep};
13270 use ai_agents_tools::{
13271 CalculatorTool, CopyPathTool, DeletePathTool, FileWriteTool, MovePathTool, ToolAliases,
13272 ToolDescriptor, ToolProvider, ToolProviderError, ToolProviderType, WebFetchResolver,
13273 WebFetchTool, WebFetchTransport, WebFetchTransportRequest, WebFetchTransportResponse,
13274 };
13275
13276 fn mock_with_response(response: &str) -> MockLLMProvider {
13277 let mut mock = MockLLMProvider::new("test");
13278 mock.set_response(response);
13279 mock
13280 }
13281
13282 fn mock_with_responses(responses: Vec<&str>) -> MockLLMProvider {
13283 let mut mock = MockLLMProvider::new("test");
13284 mock.set_responses(responses.into_iter().map(String::from).collect(), true);
13285 mock
13286 }
13287
13288 async fn collect_stream_events(
13290 agent: &RuntimeAgent,
13291 input: &str,
13292 ) -> (String, Vec<StreamChunk>, Option<AgentResponse>) {
13293 use futures::StreamExt;
13294 let mut events = agent.chat_stream_events(input).await.expect("stream opens");
13295 let mut content = String::new();
13296 let mut chunks = Vec::new();
13297 let mut final_response = None;
13298 while let Some(event) = events.next().await {
13299 match event {
13300 AgentStreamEvent::Chunk(chunk) => {
13301 if let StreamChunk::Content { text } = &chunk {
13302 content.push_str(text);
13303 }
13304 chunks.push(chunk);
13305 }
13306 AgentStreamEvent::Final(response) => final_response = Some(response),
13307 }
13308 }
13309 (content, chunks, final_response)
13310 }
13311
13312 fn metadata_keys(response: &AgentResponse) -> std::collections::BTreeSet<String> {
13313 response
13314 .metadata
13315 .as_ref()
13316 .map(|m| m.keys().cloned().collect())
13317 .unwrap_or_default()
13318 }
13319
13320 async fn assert_blocking_streaming_parity<F>(
13323 build: F,
13324 input: &str,
13325 ) -> (AgentResponse, AgentResponse, Vec<StreamChunk>)
13326 where
13327 F: Fn() -> RuntimeAgent,
13328 {
13329 let blocking_agent = build();
13330 let streaming_agent = build();
13331
13332 let blocking = blocking_agent
13333 .chat(input)
13334 .await
13335 .expect("blocking chat succeeds");
13336 let (_, chunks, final_response) = collect_stream_events(&streaming_agent, input).await;
13337 let streamed = final_response.expect("streaming must emit Final when blocking succeeds");
13338
13339 assert_eq!(
13340 blocking.content, streamed.content,
13341 "committed content differs"
13342 );
13343 assert_eq!(
13344 metadata_keys(&blocking),
13345 metadata_keys(&streamed),
13346 "metadata key sets differ"
13347 );
13348 assert_eq!(
13349 blocking.tool_calls.as_ref().map(Vec::len),
13350 streamed.tool_calls.as_ref().map(Vec::len),
13351 "tool call counts differ"
13352 );
13353 assert_eq!(
13354 blocking_agent.current_state(),
13355 streaming_agent.current_state(),
13356 "final states differ"
13357 );
13358 (blocking, streamed, chunks)
13359 }
13360
13361 fn signed_calculator_response(
13362 exchange_id: &str,
13363 call_id: &str,
13364 expression: &str,
13365 ) -> LLMResponse {
13366 let call = ToolCall {
13367 id: call_id.to_string(),
13368 name: "calculator".to_string(),
13369 arguments: serde_json::json!({"expression": expression}),
13370 };
13371 let state = ai_agents_core::NativeProviderState::new(
13372 exchange_id,
13373 "fixture",
13374 "native-tools",
13375 ai_agents_core::NativeProviderTarget::new("https://fixture.invalid/", "fixture-model")
13376 .unwrap(),
13377 serde_json::json!({
13378 "role": "model",
13379 "parts": [{
13380 "functionCall": {"name": "calculator", "args": {"expression": expression}},
13381 "thoughtSignature": format!("signature-{exchange_id}")
13382 }]
13383 }),
13384 vec![ai_agents_core::NativeCallBinding::new(call_id, 0).unwrap()],
13385 )
13386 .unwrap();
13387 LLMResponse::new("", FinishReason::ToolCall)
13388 .with_provider_state(state)
13389 .unwrap()
13390 .with_tool_calls(vec![call])
13391 .unwrap()
13392 }
13393
13394 struct TerminalHistoryProvider {
13395 calls: Arc<std::sync::atomic::AtomicU32>,
13396 }
13397
13398 struct DroppingSignedAssistantMemory {
13399 messages: RwLock<Vec<ChatMessage>>,
13400 }
13401
13402 struct DroppingEarlierSequentialMemory {
13403 messages: RwLock<Vec<ChatMessage>>,
13404 signed_seen: std::sync::atomic::AtomicUsize,
13405 }
13406
13407 #[async_trait]
13408 impl ai_agents_core::Memory for DroppingSignedAssistantMemory {
13409 async fn add_message(&self, message: ChatMessage) -> Result<()> {
13410 let signed = message.role == ai_agents_core::Role::Assistant
13411 && ai_agents_core::decode_native_tool_call_markers(&message.content)
13412 .map_err(|error| AgentError::LLM(error.to_string()))?
13413 .is_some_and(|batch| batch.provider_state().is_some());
13414 if !signed {
13415 self.messages.write().push(message);
13416 }
13417 Ok(())
13418 }
13419
13420 async fn get_messages(&self, limit: Option<usize>) -> Result<Vec<ChatMessage>> {
13421 let messages = self.messages.read();
13422 let start = limit
13423 .map(|limit| messages.len().saturating_sub(limit))
13424 .unwrap_or(0);
13425 Ok(messages[start..].to_vec())
13426 }
13427
13428 async fn clear(&self) -> Result<()> {
13429 self.messages.write().clear();
13430 Ok(())
13431 }
13432
13433 fn len(&self) -> usize {
13434 self.messages.read().len()
13435 }
13436
13437 async fn restore(&self, snapshot: ai_agents_core::MemorySnapshot) -> Result<()> {
13438 *self.messages.write() = snapshot.messages;
13439 Ok(())
13440 }
13441 }
13442
13443 #[async_trait]
13444 impl ai_agents_memory::Memory for DroppingSignedAssistantMemory {}
13445
13446 #[async_trait]
13447 impl ai_agents_core::Memory for DroppingEarlierSequentialMemory {
13448 async fn add_message(&self, message: ChatMessage) -> Result<()> {
13449 let signed = message.role == ai_agents_core::Role::Assistant
13450 && ai_agents_core::decode_native_tool_call_markers(&message.content)
13451 .map_err(|error| AgentError::LLM(error.to_string()))?
13452 .is_some_and(|batch| batch.provider_state().is_some());
13453 let mut messages = self.messages.write();
13454 if signed && self.signed_seen.fetch_add(1, Ordering::SeqCst) == 1 {
13455 messages.retain(|stored| !stored.content.contains("seq-call-1"));
13456 }
13457 messages.push(message);
13458 Ok(())
13459 }
13460
13461 async fn get_messages(&self, limit: Option<usize>) -> Result<Vec<ChatMessage>> {
13462 let messages = self.messages.read();
13463 let start = limit
13464 .map(|limit| messages.len().saturating_sub(limit))
13465 .unwrap_or(0);
13466 Ok(messages[start..].to_vec())
13467 }
13468
13469 async fn clear(&self) -> Result<()> {
13470 self.messages.write().clear();
13471 self.signed_seen.store(0, Ordering::SeqCst);
13472 Ok(())
13473 }
13474
13475 fn len(&self) -> usize {
13476 self.messages.read().len()
13477 }
13478
13479 async fn restore(&self, snapshot: ai_agents_core::MemorySnapshot) -> Result<()> {
13480 *self.messages.write() = snapshot.messages;
13481 self.signed_seen.store(0, Ordering::SeqCst);
13482 Ok(())
13483 }
13484 }
13485
13486 #[async_trait]
13487 impl ai_agents_memory::Memory for DroppingEarlierSequentialMemory {}
13488
13489 #[async_trait]
13490 impl LLMProvider for TerminalHistoryProvider {
13491 async fn complete(
13492 &self,
13493 _messages: &[ChatMessage],
13494 _config: Option<&LLMConfig>,
13495 ) -> std::result::Result<LLMResponse, LLMError> {
13496 self.calls.fetch_add(1, Ordering::SeqCst);
13497 Err(LLMError::Serialization(
13498 "native history integrity failure".to_string(),
13499 ))
13500 }
13501
13502 async fn complete_stream(
13503 &self,
13504 _messages: &[ChatMessage],
13505 _config: Option<&LLMConfig>,
13506 ) -> std::result::Result<
13507 Box<dyn Stream<Item = std::result::Result<LLMChunk, LLMError>> + Unpin + Send>,
13508 LLMError,
13509 > {
13510 Err(LLMError::Serialization(
13511 "native history integrity failure".to_string(),
13512 ))
13513 }
13514
13515 fn provider_name(&self) -> &str {
13516 "terminal-history"
13517 }
13518
13519 fn supports(&self, _feature: LLMFeature) -> bool {
13520 false
13521 }
13522
13523 fn is_terminal_error(&self, error: &LLMError) -> bool {
13524 matches!(error, LLMError::Serialization(_))
13525 }
13526 }
13527
13528 fn disambiguation_state_machine(
13530 state_enabled: Option<bool>,
13531 require_confirmation: bool,
13532 ) -> Arc<StateMachine> {
13533 let definition = ai_agents_state::StateDefinition {
13534 prompt: Some("Handle the resolved request.".to_string()),
13535 disambiguation: Some(ai_agents_disambiguation::StateDisambiguationOverride {
13536 enabled: state_enabled,
13537 require_confirmation,
13538 ..Default::default()
13539 }),
13540 ..Default::default()
13541 };
13542 let review = ai_agents_state::StateDefinition {
13543 prompt: Some("Review a fresh request.".to_string()),
13544 ..Default::default()
13545 };
13546 Arc::new(
13547 StateMachine::new(ai_agents_state::StateConfig {
13548 initial: "active".to_string(),
13549 states: std::collections::HashMap::from([
13550 ("active".to_string(), definition),
13551 ("review".to_string(), review),
13552 ]),
13553 global_transitions: Vec::new(),
13554 fallback: None,
13555 max_no_transition: None,
13556 regenerate_on_transition: true,
13557 })
13558 .unwrap(),
13559 )
13560 }
13561
13562 fn state_disambiguation_agent(
13564 responses: Vec<&str>,
13565 manager_enabled: bool,
13566 state_enabled: Option<bool>,
13567 require_confirmation: bool,
13568 ) -> (RuntimeAgent, MockLLMProvider) {
13569 state_disambiguation_agent_with_skills(
13570 responses,
13571 manager_enabled,
13572 state_enabled,
13573 require_confirmation,
13574 Vec::new(),
13575 )
13576 }
13577
13578 fn state_disambiguation_agent_with_skills(
13580 responses: Vec<&str>,
13581 manager_enabled: bool,
13582 state_enabled: Option<bool>,
13583 require_confirmation: bool,
13584 skills: Vec<SkillDefinition>,
13585 ) -> (RuntimeAgent, MockLLMProvider) {
13586 let mut mock = MockLLMProvider::new("state-confirmation");
13587 mock.set_responses(responses.into_iter().map(String::from).collect(), false);
13588 let observed = mock.clone();
13589 let agent = AgentBuilder::new()
13590 .system_prompt("Handle requests.")
13591 .llm(Arc::new(mock.clone()))
13592 .llm_alias("router", Arc::new(mock))
13593 .state_machine(disambiguation_state_machine(
13594 state_enabled,
13595 require_confirmation,
13596 ))
13597 .skills(skills)
13598 .build()
13599 .unwrap()
13600 .with_disambiguation(DisambiguationConfig {
13601 enabled: manager_enabled,
13602 ..Default::default()
13603 });
13604 (agent, observed)
13605 }
13606
13607 fn confirmation_skill() -> SkillDefinition {
13609 SkillDefinition {
13610 id: "send_report".to_string(),
13611 description: "Send a report after clarification".to_string(),
13612 trigger: "When the user asks to send a report".to_string(),
13613 steps: vec![SkillStep::Prompt {
13614 prompt: "Execute confirmed report skill for: {{ input }}".to_string(),
13615 llm: None,
13616 }],
13617 reasoning: None,
13618 reflection: None,
13619 disambiguation: Some(ai_agents_disambiguation::SkillDisambiguationOverride {
13620 enabled: Some(true),
13621 ..Default::default()
13622 }),
13623 }
13624 }
13625
13626 fn confirmation_skill_call_count(observed: &MockLLMProvider) -> usize {
13628 observed
13629 .call_history()
13630 .iter()
13631 .filter(|call| {
13632 call.messages
13633 .iter()
13634 .any(|message| message.content.contains("Execute confirmed report skill"))
13635 })
13636 .count()
13637 }
13638
13639 struct BlockingRuntimeConfirmationObserver {
13640 entered: tokio::sync::Barrier,
13641 release: tokio::sync::Notify,
13642 }
13643
13644 impl BlockingRuntimeConfirmationObserver {
13645 fn new() -> Self {
13646 Self {
13647 entered: tokio::sync::Barrier::new(2),
13648 release: tokio::sync::Notify::new(),
13649 }
13650 }
13651 }
13652
13653 struct ResetOnTransitionHooks {
13654 agent: parking_lot::Mutex<Option<Weak<RuntimeAgent>>>,
13655 invoked: AtomicBool,
13656 }
13657
13658 #[async_trait]
13659 impl AgentHooks for ResetOnTransitionHooks {
13660 async fn on_state_transition(&self, _from: Option<&str>, _to: &str, _reason: &str) {
13661 if self.invoked.swap(true, Ordering::SeqCst) {
13662 return;
13663 }
13664 let agent = self.agent.lock().as_ref().and_then(Weak::upgrade);
13665 if let Some(agent) = agent {
13666 agent.reset().await.unwrap();
13667 }
13668 }
13669 }
13670
13671 impl ClarificationObserver for BlockingRuntimeConfirmationObserver {
13672 fn observe_question<'a>(
13673 &'a self,
13674 future: ClarificationQuestionFuture<'a>,
13675 ) -> ClarificationQuestionFuture<'a> {
13676 future
13677 }
13678
13679 fn observe_parse<'a>(
13680 &'a self,
13681 future: ClarificationParseFuture<'a>,
13682 ) -> ClarificationParseFuture<'a> {
13683 future
13684 }
13685
13686 fn observe_confirmation_parse<'a>(
13687 &'a self,
13688 future: ConfirmationParseFuture<'a>,
13689 ) -> ConfirmationParseFuture<'a> {
13690 Box::pin(async move {
13691 self.entered.wait().await;
13692 self.release.notified().await;
13693 future.await
13694 })
13695 }
13696 }
13697
13698 #[tokio::test]
13699 async fn state_confirmation_blocks_redispatch_until_explicit_agreement() {
13700 let (agent, observed) = state_disambiguation_agent(
13701 vec![
13702 r#"{"is_ambiguous":true,"confidence":0.2,"ambiguity_type":"missing_target","reasoning":"target missing","what_is_unclear":["target"],"detected_language":"en"}"#,
13703 r#"{"question":"What should I send?","options":null}"#,
13704 r#"{"status":"answered","selected_option":null,"enriched_input":"Send the report to Ada","resolved":{"intent":"send_report"}}"#,
13705 r#"{"question":"Should I send the report to Ada?"}"#,
13706 r#"{"status":"confirmed"}"#,
13707 "Request executed.",
13708 ],
13709 true,
13710 None,
13711 true,
13712 );
13713
13714 let clarification = agent.chat("Send it").await.unwrap();
13715 assert_eq!(clarification.content, "What should I send?");
13716 assert_eq!(observed.call_count(), 2);
13717
13718 let confirmation = agent.chat("The report to Ada").await.unwrap();
13719 assert_eq!(confirmation.content, "Should I send the report to Ada?");
13720 assert_eq!(
13721 confirmation
13722 .metadata
13723 .as_ref()
13724 .and_then(|metadata| metadata.get("disambiguation"))
13725 .and_then(|metadata| metadata.get("status"))
13726 .and_then(Value::as_str),
13727 Some("awaiting_confirmation")
13728 );
13729 assert_eq!(observed.call_count(), 4);
13730
13731 let completed = agent.chat("Yes").await.unwrap();
13732 assert_eq!(completed.content, "Request executed.");
13733 assert_eq!(observed.call_count(), 6);
13734 }
13735
13736 #[tokio::test]
13737 async fn streaming_state_confirmation_ends_the_turn_before_redispatch() {
13738 let (agent, observed) = state_disambiguation_agent(
13739 vec![
13740 r#"{"is_ambiguous":true,"confidence":0.2,"ambiguity_type":"missing_target","reasoning":"target missing","what_is_unclear":["target"],"detected_language":"en"}"#,
13741 r#"{"question":"What should I send?","options":null}"#,
13742 r#"{"status":"answered","selected_option":null,"enriched_input":"Send the report to Ada","resolved":{"intent":"send_report"}}"#,
13743 r#"{"question":"Should I send the report to Ada?"}"#,
13744 r#"{"status":"confirmed"}"#,
13745 "Request executed.",
13746 ],
13747 true,
13748 None,
13749 true,
13750 );
13751
13752 let mut clarification_stream = agent.chat_stream("Send it").await.unwrap();
13753 let mut clarification = String::new();
13754 while let Some(chunk) = clarification_stream.next().await {
13755 match chunk {
13756 StreamChunk::Content { text } => clarification.push_str(&text),
13757 StreamChunk::Done {} => break,
13758 StreamChunk::Error { message } => panic!("unexpected stream error: {message}"),
13759 _ => {}
13760 }
13761 }
13762 assert_eq!(clarification, "What should I send?");
13763 assert_eq!(observed.call_count(), 2);
13764
13765 let mut confirmation_stream = agent.chat_stream_events("The report to Ada").await.unwrap();
13766 let mut confirmation = None;
13767 while let Some(event) = confirmation_stream.next().await {
13768 match event {
13769 AgentStreamEvent::Final(response) => confirmation = Some(response),
13770 AgentStreamEvent::Chunk(StreamChunk::Error { message }) => {
13771 panic!("unexpected stream error: {message}")
13772 }
13773 AgentStreamEvent::Chunk(_) => {}
13774 }
13775 }
13776 let confirmation = confirmation.expect("confirmation must finalize");
13777 assert_eq!(confirmation.content, "Should I send the report to Ada?");
13778 assert_eq!(
13779 confirmation
13780 .metadata
13781 .as_ref()
13782 .and_then(|metadata| metadata.get("disambiguation"))
13783 .and_then(|metadata| metadata.get("status"))
13784 .and_then(Value::as_str),
13785 Some("awaiting_confirmation")
13786 );
13787 assert_eq!(observed.call_count(), 4);
13788
13789 let mut completed_stream = agent.chat_stream("Yes").await.unwrap();
13790 let mut completed = String::new();
13791 while let Some(chunk) = completed_stream.next().await {
13792 match chunk {
13793 StreamChunk::Content { text } => completed.push_str(&text),
13794 StreamChunk::Done {} => break,
13795 StreamChunk::Error { message } => panic!("unexpected stream error: {message}"),
13796 _ => {}
13797 }
13798 }
13799 assert_eq!(completed, "Request executed.");
13800 assert_eq!(observed.call_count(), 6);
13801 }
13802
13803 #[tokio::test]
13805 async fn root_turn_gate_serializes_blocking_and_streaming_entry_points() {
13806 let (complete_entered, mut complete_events) = tokio::sync::mpsc::unbounded_channel();
13807 let agent = Arc::new(
13808 AgentBuilder::new()
13809 .system_prompt("Serialize root turns.")
13810 .llm(Arc::new(RootTurnProbeProvider { complete_entered }))
13811 .build()
13812 .unwrap(),
13813 );
13814 let blocking_agent = Arc::clone(&agent);
13815
13816 let legacy_stream = agent.chat_stream("stream owner").await.unwrap();
13817 assert!(agent.root_turn_gate.try_lock().is_err());
13818 let blocking = tokio::spawn(async move { blocking_agent.chat("blocked").await.unwrap() });
13819 assert!(
13820 tokio::time::timeout(std::time::Duration::from_millis(50), complete_events.recv())
13821 .await
13822 .is_err(),
13823 "blocking turn reached the provider while the legacy stream owned the root gate"
13824 );
13825
13826 drop(legacy_stream);
13827 assert_eq!(
13828 tokio::time::timeout(std::time::Duration::from_secs(2), complete_events.recv())
13829 .await
13830 .expect("blocking turn did not enter after stream drop"),
13831 Some(())
13832 );
13833 let response = tokio::time::timeout(std::time::Duration::from_secs(2), blocking)
13834 .await
13835 .expect("blocking turn did not finish after stream drop")
13836 .unwrap();
13837 assert_eq!(response.content, "blocking complete");
13838
13839 let mut event_stream = agent.chat_stream_events("event terminal").await.unwrap();
13840 assert!(agent.root_turn_gate.try_lock().is_err());
13841 let mut saw_final = false;
13842 while let Some(event) = event_stream.next().await {
13843 if matches!(event, AgentStreamEvent::Final(_)) {
13844 saw_final = true;
13845 break;
13846 }
13847 }
13848 assert!(saw_final);
13849 assert!(
13850 agent.root_turn_gate.try_lock().is_ok(),
13851 "authoritative terminal event retained the root gate"
13852 );
13853 }
13854
13855 #[tokio::test]
13857 async fn response_hook_rejects_same_runtime_chat_reentry() {
13858 let hooks = Arc::new(ResponseChatHooks {
13859 target: parking_lot::Mutex::new(None),
13860 invoked: AtomicBool::new(false),
13861 nested_result: parking_lot::Mutex::new(None),
13862 });
13863 let agent = Arc::new(
13864 AgentBuilder::new()
13865 .system_prompt("Reject response hook reentry.")
13866 .llm(Arc::new(mock_with_response("outer response")))
13867 .hooks(hooks.clone())
13868 .build()
13869 .unwrap(),
13870 );
13871 *hooks.target.lock() = Some(Arc::downgrade(&agent));
13872
13873 let response = tokio::time::timeout(
13874 std::time::Duration::from_secs(2),
13875 agent.chat("outer request"),
13876 )
13877 .await
13878 .expect("same-runtime response hook reentry must fail without deadlocking")
13879 .unwrap();
13880
13881 assert_eq!(response.content, "outer response");
13882 let nested_result = hooks
13883 .nested_result
13884 .lock()
13885 .clone()
13886 .expect("response hook must record its nested call");
13887 let error = nested_result.expect_err("same-runtime nested chat must be rejected");
13888 assert!(error.contains("reentrant root turn ownership"));
13889 }
13890
13891 #[tokio::test]
13893 async fn root_turn_gate_allows_nested_runtime_and_rejects_cycles() {
13894 let agent_a = AgentBuilder::new()
13895 .system_prompt("Runtime A.")
13896 .llm(Arc::new(mock_with_response("response A")))
13897 .build()
13898 .unwrap();
13899 let agent_b = AgentBuilder::new()
13900 .system_prompt("Runtime B.")
13901 .llm(Arc::new(mock_with_response("response B")))
13902 .build()
13903 .unwrap();
13904 let RootTurnAdmission {
13905 guard: guard_a,
13906 identity_stack: stack_a,
13907 } = agent_a.acquire_root_turn().await.unwrap();
13908
13909 let cycle_error = scope_runtime_gate_identity_stack(&stack_a, async {
13910 let RootTurnAdmission {
13911 guard: guard_b,
13912 identity_stack: stack_b,
13913 } = agent_b
13914 .acquire_root_turn()
13915 .await
13916 .expect("runtime B must acquire a different gate");
13917 let result =
13918 scope_runtime_gate_identity_stack(&stack_b, agent_a.acquire_root_turn()).await;
13919 drop(guard_b);
13920 match result {
13921 Err(error) => error,
13922 Ok(_) => panic!("runtime A accepted a repeated gate identity"),
13923 }
13924 })
13925 .await;
13926 drop(guard_a);
13927
13928 assert!(
13929 cycle_error
13930 .to_string()
13931 .contains("reentrant root turn ownership")
13932 );
13933 }
13934
13935 #[tokio::test]
13937 async fn concurrent_orchestration_propagates_root_gate_ancestry() {
13938 let registry = Arc::new(crate::spawner::AgentRegistry::new());
13939 let hooks_a = Arc::new(ConcurrentResponseHooks {
13940 registry: Arc::downgrade(®istry),
13941 child_id: "runtime-b".to_string(),
13942 invoked: AtomicBool::new(false),
13943 nested_result: parking_lot::Mutex::new(None),
13944 });
13945 let hooks_b = Arc::new(ResponseChatHooks {
13946 target: parking_lot::Mutex::new(None),
13947 invoked: AtomicBool::new(false),
13948 nested_result: parking_lot::Mutex::new(None),
13949 });
13950 let agent_a = AgentBuilder::new()
13951 .system_prompt("Runtime A dispatches runtime B concurrently.")
13952 .llm(Arc::new(mock_with_response("response A")))
13953 .hooks(hooks_a.clone())
13954 .build()
13955 .unwrap();
13956 let agent_b = AgentBuilder::new()
13957 .system_prompt("Runtime B attempts to re-enter runtime A.")
13958 .llm(Arc::new(mock_with_response("response B")))
13959 .hooks(hooks_b.clone())
13960 .build()
13961 .unwrap();
13962 let spec_a = crate::spec::AgentSpec {
13963 name: "runtime-a".to_string(),
13964 system_prompt: "Runtime A dispatches runtime B concurrently.".to_string(),
13965 ..crate::spec::AgentSpec::default()
13966 };
13967 let spec_b = crate::spec::AgentSpec {
13968 name: "runtime-b".to_string(),
13969 system_prompt: "Runtime B attempts to re-enter runtime A.".to_string(),
13970 ..crate::spec::AgentSpec::default()
13971 };
13972 registry
13973 .register(crate::spawner::SpawnedAgent::from_runtime(
13974 "runtime-a".to_string(),
13975 agent_a,
13976 spec_a,
13977 ))
13978 .await
13979 .unwrap();
13980 registry
13981 .register(crate::spawner::SpawnedAgent::from_runtime(
13982 "runtime-b".to_string(),
13983 agent_b,
13984 spec_b,
13985 ))
13986 .await
13987 .unwrap();
13988 let runtime_a = registry.get("runtime-a").unwrap();
13989 *hooks_b.target.lock() = Some(Arc::downgrade(&runtime_a));
13990
13991 let response = tokio::time::timeout(
13992 std::time::Duration::from_secs(2),
13993 runtime_a.chat("outer concurrent request"),
13994 )
13995 .await
13996 .expect("concurrent orchestration cycle must fail without deadlocking")
13997 .unwrap();
13998
13999 assert_eq!(response.content, "response A");
14000 let child_result = hooks_a
14001 .nested_result
14002 .lock()
14003 .clone()
14004 .expect("runtime A hook must record runtime B completion");
14005 assert_eq!(child_result.unwrap(), "response B");
14006 let cycle_result = hooks_b
14007 .nested_result
14008 .lock()
14009 .clone()
14010 .expect("runtime B hook must record runtime A reentry");
14011 assert!(
14012 cycle_result
14013 .expect_err("runtime A accepted a repeated gate identity")
14014 .contains("reentrant root turn ownership")
14015 );
14016 }
14017
14018 #[tokio::test]
14020 async fn confirmed_skill_route_executes_exactly_once() {
14021 let (agent, observed) = state_disambiguation_agent_with_skills(
14022 vec![
14023 r#"{"is_ambiguous":false,"confidence":0.99,"ambiguity_type":null,"reasoning":"top-level clear","what_is_unclear":[],"detected_language":"en"}"#,
14024 "send_report",
14025 r#"{"is_ambiguous":true,"confidence":0.2,"ambiguity_type":"missing_target","reasoning":"target missing","what_is_unclear":["target"],"detected_language":"en"}"#,
14026 r#"{"question":"What should I send?","options":null}"#,
14027 r#"{"status":"answered","selected_option":null,"enriched_input":"Send the report to Ada","resolved":{"intent":"send_report"}}"#,
14028 r#"{"question":"Should I send the report to Ada?"}"#,
14029 r#"{"status":"confirmed"}"#,
14030 r#"{"is_ambiguous":false,"confidence":0.99,"ambiguity_type":null,"reasoning":"resolved","what_is_unclear":[],"detected_language":"en"}"#,
14031 "Report skill executed.",
14032 ],
14033 true,
14034 None,
14035 true,
14036 vec![confirmation_skill()],
14037 );
14038
14039 let clarification = agent.chat("Send it").await.unwrap();
14040 assert_eq!(clarification.content, "What should I send?");
14041 assert_eq!(confirmation_skill_call_count(&observed), 0);
14042
14043 let confirmation = agent.chat("The report to Ada").await.unwrap();
14044 assert_eq!(confirmation.content, "Should I send the report to Ada?");
14045 assert_eq!(
14046 confirmation
14047 .metadata
14048 .as_ref()
14049 .and_then(|metadata| metadata.get("disambiguation"))
14050 .and_then(|metadata| metadata.get("status"))
14051 .and_then(Value::as_str),
14052 Some("awaiting_confirmation")
14053 );
14054 assert_eq!(confirmation_skill_call_count(&observed), 0);
14055
14056 let completed = agent.chat("Yes").await.unwrap();
14057 assert_eq!(completed.content, "Report skill executed.");
14058 assert_eq!(confirmation_skill_call_count(&observed), 1);
14059 assert!(agent.pending_skill_id.read().is_none());
14060 let messages = agent.memory.get_messages(None).await.unwrap();
14061 assert!(!messages.iter().any(|message| message.content == "Yes"));
14062 }
14063
14064 #[tokio::test]
14066 async fn confirmed_skill_recheck_preserves_new_clarification_metadata() {
14067 let (agent, observed) = state_disambiguation_agent_with_skills(
14068 vec![
14069 r#"{"is_ambiguous":false,"confidence":0.99,"ambiguity_type":null,"reasoning":"top-level clear","what_is_unclear":[],"detected_language":"en"}"#,
14070 "send_report",
14071 r#"{"is_ambiguous":true,"confidence":0.2,"ambiguity_type":"missing_target","reasoning":"target missing","what_is_unclear":["target"],"detected_language":"en"}"#,
14072 r#"{"question":"What should I send?","options":null}"#,
14073 r#"{"status":"answered","selected_option":null,"enriched_input":"Send the report to Ada","resolved":{"intent":"send_report"}}"#,
14074 r#"{"question":"Should I send the report to Ada?"}"#,
14075 r#"{"status":"confirmed"}"#,
14076 r#"{"is_ambiguous":true,"confidence":0.3,"ambiguity_type":"missing_parameters","reasoning":"timing missing","what_is_unclear":["timing"],"detected_language":"en"}"#,
14077 r#"{"question":"When should I send it?","options":null}"#,
14078 ],
14079 true,
14080 None,
14081 true,
14082 vec![confirmation_skill()],
14083 );
14084
14085 agent.chat("Send it").await.unwrap();
14086 agent.chat("The report to Ada").await.unwrap();
14087 let follow_up = agent.chat("Yes").await.unwrap();
14088
14089 assert_eq!(follow_up.content, "When should I send it?");
14090 let metadata = follow_up
14091 .metadata
14092 .as_ref()
14093 .and_then(|metadata| metadata.get("disambiguation"))
14094 .unwrap();
14095 assert_eq!(
14096 metadata.get("status").and_then(Value::as_str),
14097 Some("awaiting_clarification")
14098 );
14099 assert_eq!(
14100 metadata.get("skill_id").and_then(Value::as_str),
14101 Some("send_report")
14102 );
14103 assert!(metadata.get("detection").is_some());
14104 assert_eq!(confirmation_skill_call_count(&observed), 0);
14105 }
14106
14107 #[tokio::test]
14109 async fn rejected_skill_confirmation_never_executes() {
14110 let (agent, observed) = state_disambiguation_agent_with_skills(
14111 vec![
14112 r#"{"is_ambiguous":false,"confidence":0.99,"ambiguity_type":null,"reasoning":"top-level clear","what_is_unclear":[],"detected_language":"en"}"#,
14113 "send_report",
14114 r#"{"is_ambiguous":true,"confidence":0.2,"ambiguity_type":"missing_target","reasoning":"target missing","what_is_unclear":["target"],"detected_language":"en"}"#,
14115 r#"{"question":"What should I send?","options":null}"#,
14116 r#"{"status":"answered","selected_option":null,"enriched_input":"Send the report to Ada","resolved":{"intent":"send_report"}}"#,
14117 r#"{"question":"Should I send the report to Ada?"}"#,
14118 r#"{"status":"rejected"}"#,
14119 "Confirmation rejected.",
14120 ],
14121 true,
14122 None,
14123 true,
14124 vec![confirmation_skill()],
14125 );
14126
14127 agent.chat("Send it").await.unwrap();
14128 agent.chat("The report to Ada").await.unwrap();
14129 let rejected = agent.chat("No").await.unwrap();
14130
14131 assert_eq!(rejected.content, "Confirmation rejected.");
14132 assert_eq!(confirmation_skill_call_count(&observed), 0);
14133 assert!(agent.pending_skill_id.read().is_none());
14134 }
14135
14136 #[tokio::test]
14138 async fn reset_invalidates_pending_skill_confirmation_before_streaming_input() {
14139 let (agent, observed) = state_disambiguation_agent_with_skills(
14140 vec![
14141 r#"{"is_ambiguous":false,"confidence":0.99,"ambiguity_type":null,"reasoning":"top-level clear","what_is_unclear":[],"detected_language":"en"}"#,
14142 "send_report",
14143 r#"{"is_ambiguous":true,"confidence":0.2,"ambiguity_type":"missing_target","reasoning":"target missing","what_is_unclear":["target"],"detected_language":"en"}"#,
14144 r#"{"question":"What should I send?","options":null}"#,
14145 r#"{"status":"answered","selected_option":null,"enriched_input":"Send the report to Ada","resolved":{"intent":"send_report"}}"#,
14146 r#"{"question":"Should I send the report to Ada?"}"#,
14147 r#"{"is_ambiguous":false,"confidence":0.99,"ambiguity_type":null,"reasoning":"fresh input","what_is_unclear":[],"detected_language":"en"}"#,
14148 "none",
14149 "Fresh response.",
14150 ],
14151 true,
14152 None,
14153 true,
14154 vec![confirmation_skill()],
14155 );
14156
14157 agent.chat("Send it").await.unwrap();
14158 agent.chat("The report to Ada").await.unwrap();
14159 agent.reset().await.unwrap();
14160 assert!(agent.pending_skill_id.read().is_none());
14161 assert!(
14162 !agent
14163 .disambiguation_manager()
14164 .unwrap()
14165 .has_pending_clarification()
14166 .await
14167 );
14168
14169 let mut stream = agent.chat_stream("Yes").await.unwrap();
14170 let mut content = String::new();
14171 while let Some(chunk) = stream.next().await {
14172 match chunk {
14173 StreamChunk::Content { text } => content.push_str(&text),
14174 StreamChunk::Done {} => break,
14175 StreamChunk::Error { message } => panic!("unexpected stream error: {message}"),
14176 _ => {}
14177 }
14178 }
14179
14180 assert_eq!(content, "Fresh response.");
14181 assert_eq!(confirmation_skill_call_count(&observed), 0);
14182 }
14183
14184 #[tokio::test]
14186 async fn trait_reset_clears_pending_skill_confirmation() {
14187 let (agent, _) = state_disambiguation_agent_with_skills(
14188 vec![
14189 r#"{"is_ambiguous":false,"confidence":0.99,"ambiguity_type":null,"reasoning":"top-level clear","what_is_unclear":[],"detected_language":"en"}"#,
14190 "send_report",
14191 r#"{"is_ambiguous":true,"confidence":0.2,"ambiguity_type":"missing_target","reasoning":"target missing","what_is_unclear":["target"],"detected_language":"en"}"#,
14192 r#"{"question":"What should I send?","options":null}"#,
14193 r#"{"status":"answered","selected_option":null,"enriched_input":"Send the report to Ada","resolved":{"intent":"send_report"}}"#,
14194 r#"{"question":"Should I send the report to Ada?"}"#,
14195 ],
14196 true,
14197 None,
14198 true,
14199 vec![confirmation_skill()],
14200 );
14201
14202 agent.chat("Send it").await.unwrap();
14203 agent.chat("The report to Ada").await.unwrap();
14204 <RuntimeAgent as Agent>::reset(&agent).await.unwrap();
14205
14206 assert!(agent.pending_skill_id.read().is_none());
14207 assert!(
14208 !agent
14209 .disambiguation_manager()
14210 .unwrap()
14211 .has_pending_clarification()
14212 .await
14213 );
14214 }
14215
14216 #[tokio::test]
14218 async fn state_change_invalidates_pending_skill_confirmation() {
14219 let (agent, observed) = state_disambiguation_agent_with_skills(
14220 vec![
14221 r#"{"is_ambiguous":false,"confidence":0.99,"ambiguity_type":null,"reasoning":"top-level clear","what_is_unclear":[],"detected_language":"en"}"#,
14222 "send_report",
14223 r#"{"is_ambiguous":true,"confidence":0.2,"ambiguity_type":"missing_target","reasoning":"target missing","what_is_unclear":["target"],"detected_language":"en"}"#,
14224 r#"{"question":"What should I send?","options":null}"#,
14225 r#"{"status":"answered","selected_option":null,"enriched_input":"Send the report to Ada","resolved":{"intent":"send_report"}}"#,
14226 r#"{"question":"Should I send the report to Ada?"}"#,
14227 r#"{"is_ambiguous":false,"confidence":0.99,"ambiguity_type":null,"reasoning":"fresh input","what_is_unclear":[],"detected_language":"en"}"#,
14228 "none",
14229 "Fresh response.",
14230 ],
14231 true,
14232 None,
14233 true,
14234 vec![confirmation_skill()],
14235 );
14236
14237 agent.chat("Send it").await.unwrap();
14238 agent.chat("The report to Ada").await.unwrap();
14239 agent.transition_to("review").await.unwrap();
14240 let cancelled = agent.chat("Yes").await.unwrap();
14241
14242 assert_eq!(cancelled.content, "Fresh response.");
14243 assert_eq!(confirmation_skill_call_count(&observed), 0);
14244 assert!(agent.pending_skill_id.read().is_none());
14245 }
14246
14247 #[tokio::test]
14249 async fn in_flight_confirmation_cannot_redispatch_after_reset() {
14250 let (mut agent, observed) = state_disambiguation_agent_with_skills(
14251 vec![
14252 r#"{"is_ambiguous":false,"confidence":0.99,"ambiguity_type":null,"reasoning":"top-level clear","what_is_unclear":[],"detected_language":"en"}"#,
14253 "send_report",
14254 r#"{"is_ambiguous":true,"confidence":0.2,"ambiguity_type":"missing_target","reasoning":"target missing","what_is_unclear":["target"],"detected_language":"en"}"#,
14255 r#"{"question":"What should I send?","options":null}"#,
14256 r#"{"status":"answered","selected_option":null,"enriched_input":"Send the report to Ada","resolved":{"intent":"send_report"}}"#,
14257 r#"{"question":"Should I send the report to Ada?"}"#,
14258 r#"{"status":"confirmed"}"#,
14259 "Confirmation cancelled.",
14260 ],
14261 true,
14262 None,
14263 true,
14264 vec![confirmation_skill()],
14265 );
14266 let observer = Arc::new(BlockingRuntimeConfirmationObserver::new());
14267 let manager = agent
14268 .disambiguation_manager
14269 .take()
14270 .unwrap()
14271 .with_clarification_observer(observer.clone());
14272 agent.disambiguation_manager = Some(manager);
14273 let agent = Arc::new(agent);
14274
14275 agent.chat("Send it").await.unwrap();
14276 agent.chat("The report to Ada").await.unwrap();
14277
14278 let confirming_agent = Arc::clone(&agent);
14279 let confirmation = tokio::spawn(async move { confirming_agent.chat("Yes").await });
14280 observer.entered.wait().await;
14281 agent.reset().await.unwrap();
14282 observer.release.notify_one();
14283
14284 let response = confirmation.await.unwrap().unwrap();
14285 assert_eq!(response.content, "Confirmation cancelled.");
14286 assert_eq!(confirmation_skill_call_count(&observed), 0);
14287 assert!(agent.pending_skill_id.read().is_none());
14288 }
14289
14290 #[tokio::test]
14292 async fn queued_reset_prevents_stale_confirmation_question_publication() {
14293 let (agent, observed) = state_disambiguation_agent(
14294 vec![
14295 r#"{"is_ambiguous":true,"confidence":0.2,"ambiguity_type":"missing_target","reasoning":"target missing","what_is_unclear":["target"],"detected_language":"en"}"#,
14296 r#"{"question":"What should I send?","options":null}"#,
14297 r#"{"status":"answered","selected_option":null,"enriched_input":"Send the report to Ada","resolved":{"intent":"send_report"}}"#,
14298 r#"{"question":"Should I send the report to Ada?"}"#,
14299 ],
14300 true,
14301 None,
14302 true,
14303 );
14304 let agent = Arc::new(agent);
14305 agent.chat("Send it").await.unwrap();
14306
14307 let admission = agent.disambiguation_admission.write().await;
14308 let (started_tx, started_rx) = tokio::sync::oneshot::channel();
14309 let resetting_agent = Arc::clone(&agent);
14310 let reset = tokio::spawn(async move {
14311 let _ = started_tx.send(());
14312 resetting_agent.reset().await
14313 });
14314 started_rx.await.unwrap();
14315 tokio::task::yield_now().await;
14316
14317 let responding_agent = Arc::clone(&agent);
14318 let response =
14319 tokio::spawn(async move { responding_agent.chat("The report to Ada").await });
14320 tokio::time::timeout(std::time::Duration::from_secs(2), async {
14321 while observed.call_count() < 4 {
14322 tokio::task::yield_now().await;
14323 }
14324 })
14325 .await
14326 .expect("clarification processing must reach terminal publication");
14327 drop(admission);
14328
14329 reset.await.unwrap().unwrap();
14330 let error = response.await.unwrap().unwrap_err();
14331 assert!(error.to_string().contains("ownership changed"));
14332 assert!(
14333 !agent
14334 .disambiguation_manager()
14335 .unwrap()
14336 .has_pending_clarification()
14337 .await
14338 );
14339 assert!(agent.memory.get_messages(None).await.unwrap().is_empty());
14340 }
14341
14342 #[tokio::test]
14344 async fn queued_reset_prevents_stale_skill_clarification_publication() {
14345 let (agent, observed) = state_disambiguation_agent_with_skills(
14346 vec![
14347 r#"{"is_ambiguous":false,"confidence":0.99,"ambiguity_type":null,"reasoning":"top-level clear","what_is_unclear":[],"detected_language":"en"}"#,
14348 "send_report",
14349 r#"{"is_ambiguous":true,"confidence":0.2,"ambiguity_type":"missing_target","reasoning":"target missing","what_is_unclear":["target"],"detected_language":"en"}"#,
14350 r#"{"question":"What should I send?","options":null}"#,
14351 ],
14352 true,
14353 None,
14354 true,
14355 vec![confirmation_skill()],
14356 );
14357 let agent = Arc::new(agent);
14358 let admission = agent.disambiguation_admission.write().await;
14359 let (started_tx, started_rx) = tokio::sync::oneshot::channel();
14360 let resetting_agent = Arc::clone(&agent);
14361 let reset = tokio::spawn(async move {
14362 let _ = started_tx.send(());
14363 resetting_agent.reset().await
14364 });
14365 started_rx.await.unwrap();
14366 tokio::task::yield_now().await;
14367
14368 let responding_agent = Arc::clone(&agent);
14369 let response = tokio::spawn(async move { responding_agent.chat("Send it").await });
14370 tokio::time::timeout(std::time::Duration::from_secs(2), async {
14371 while observed.call_count() < 4 {
14372 tokio::task::yield_now().await;
14373 }
14374 })
14375 .await
14376 .expect("skill clarification must reach terminal publication");
14377 drop(admission);
14378
14379 reset.await.unwrap().unwrap();
14380 let error = response.await.unwrap().unwrap_err();
14381 assert!(error.to_string().contains("ownership changed"));
14382 assert_eq!(confirmation_skill_call_count(&observed), 0);
14383 assert!(agent.pending_skill_id.read().is_none());
14384 assert!(agent.memory.get_messages(None).await.unwrap().is_empty());
14385 }
14386
14387 #[tokio::test]
14389 async fn transition_hook_can_reset_without_admission_deadlock() {
14390 let hooks = Arc::new(ResetOnTransitionHooks {
14391 agent: parking_lot::Mutex::new(None),
14392 invoked: AtomicBool::new(false),
14393 });
14394 let agent = Arc::new(
14395 AgentBuilder::new()
14396 .system_prompt("Test transition hook reentrancy.")
14397 .llm(Arc::new(mock_with_response("done")))
14398 .state_machine(disambiguation_state_machine(None, false))
14399 .build()
14400 .unwrap()
14401 .with_hooks(hooks.clone()),
14402 );
14403 *hooks.agent.lock() = Some(Arc::downgrade(&agent));
14404
14405 let transitioned = tokio::time::timeout(
14406 std::time::Duration::from_secs(2),
14407 agent.apply_transition_target("active", "review", "test transition", None),
14408 )
14409 .await
14410 .expect("transition hook reset must not deadlock")
14411 .unwrap();
14412
14413 assert!(transitioned);
14414 assert!(hooks.invoked.load(Ordering::SeqCst));
14415 assert_eq!(agent.current_state().as_deref(), Some("active"));
14416 }
14417
14418 #[tokio::test]
14420 async fn concurrent_transition_cannot_duplicate_exit_actions() {
14421 let gate = PathMutationGate::new();
14422 let active = ai_agents_state::StateDefinition {
14423 on_exit: vec![StateAction::Tool {
14424 tool: "transition_exit".to_string(),
14425 args: Some(serde_json::json!({"path": "./transition-exit.txt"})),
14426 }],
14427 ..Default::default()
14428 };
14429 let state_machine = Arc::new(
14430 StateMachine::new(ai_agents_state::StateConfig {
14431 initial: "active".to_string(),
14432 states: HashMap::from([
14433 ("active".to_string(), active),
14434 (
14435 "review".to_string(),
14436 ai_agents_state::StateDefinition::default(),
14437 ),
14438 ]),
14439 global_transitions: Vec::new(),
14440 fallback: None,
14441 max_no_transition: None,
14442 regenerate_on_transition: true,
14443 })
14444 .unwrap(),
14445 );
14446 let agent = Arc::new(
14447 AgentBuilder::new()
14448 .system_prompt("Test transition reservation.")
14449 .llm(Arc::new(mock_with_response("done")))
14450 .tool(Arc::new(BlockingPathMutationTool {
14451 id: "transition_exit",
14452 path_fields: vec![ai_agents_core::PathPolicyBinding::write("path")],
14453 gate: gate.clone(),
14454 }))
14455 .state_machine(state_machine)
14456 .build()
14457 .unwrap(),
14458 );
14459
14460 let first_agent = Arc::clone(&agent);
14461 let first = tokio::spawn(async move { first_agent.transition_to("review").await });
14462 tokio::time::timeout(std::time::Duration::from_secs(2), gate.wait_until_entered())
14463 .await
14464 .expect("reserved transition must enter its exit action");
14465
14466 let second = tokio::time::timeout(
14467 std::time::Duration::from_secs(2),
14468 agent.transition_to("review"),
14469 )
14470 .await
14471 .expect("competing transition must fail without waiting for the exit action")
14472 .unwrap_err();
14473 assert!(second.to_string().contains("already in progress"));
14474
14475 gate.release();
14476 first.await.unwrap().unwrap();
14477 assert_eq!(agent.current_state().as_deref(), Some("review"));
14478 }
14479
14480 #[tokio::test]
14482 async fn concurrent_transition_cannot_overtake_enter_actions() {
14483 let gate = PathMutationGate::new();
14484 let review = ai_agents_state::StateDefinition {
14485 on_enter: vec![StateAction::Tool {
14486 tool: "transition_enter".to_string(),
14487 args: Some(serde_json::json!({"path": "./transition-enter.txt"})),
14488 }],
14489 ..Default::default()
14490 };
14491 let state_machine = Arc::new(
14492 StateMachine::new(ai_agents_state::StateConfig {
14493 initial: "active".to_string(),
14494 states: HashMap::from([
14495 (
14496 "active".to_string(),
14497 ai_agents_state::StateDefinition::default(),
14498 ),
14499 ("review".to_string(), review),
14500 ]),
14501 global_transitions: Vec::new(),
14502 fallback: None,
14503 max_no_transition: None,
14504 regenerate_on_transition: true,
14505 })
14506 .unwrap(),
14507 );
14508 let agent = Arc::new(
14509 AgentBuilder::new()
14510 .system_prompt("Test transition lifecycle reservation.")
14511 .llm(Arc::new(mock_with_response("done")))
14512 .tool(Arc::new(BlockingPathMutationTool {
14513 id: "transition_enter",
14514 path_fields: vec![ai_agents_core::PathPolicyBinding::write("path")],
14515 gate: gate.clone(),
14516 }))
14517 .state_machine(state_machine)
14518 .build()
14519 .unwrap(),
14520 );
14521
14522 let first_agent = Arc::clone(&agent);
14523 let first = tokio::spawn(async move { first_agent.transition_to("review").await });
14524 tokio::time::timeout(std::time::Duration::from_secs(2), gate.wait_until_entered())
14525 .await
14526 .expect("committed transition must enter its destination action");
14527
14528 let second = agent.transition_to("active").await.unwrap_err();
14529 assert!(second.to_string().contains("already in progress"));
14530 assert!(agent.reset().await.is_err());
14531
14532 gate.release();
14533 first.await.unwrap().unwrap();
14534 assert_eq!(agent.current_state().as_deref(), Some("review"));
14535 }
14536
14537 #[tokio::test]
14539 async fn same_state_restore_invalidates_pending_skill_confirmation() {
14540 let (agent, observed) = state_disambiguation_agent_with_skills(
14541 vec![
14542 r#"{"is_ambiguous":false,"confidence":0.99,"ambiguity_type":null,"reasoning":"top-level clear","what_is_unclear":[],"detected_language":"en"}"#,
14543 "send_report",
14544 r#"{"is_ambiguous":true,"confidence":0.2,"ambiguity_type":"missing_target","reasoning":"target missing","what_is_unclear":["target"],"detected_language":"en"}"#,
14545 r#"{"question":"What should I send?","options":null}"#,
14546 r#"{"status":"answered","selected_option":null,"enriched_input":"Send the report to Ada","resolved":{"intent":"send_report"}}"#,
14547 r#"{"question":"Should I send the report to Ada?"}"#,
14548 ],
14549 true,
14550 None,
14551 true,
14552 vec![confirmation_skill()],
14553 );
14554
14555 agent.chat("Send it").await.unwrap();
14556 agent.chat("The report to Ada").await.unwrap();
14557 let snapshot = agent.save_state().await.unwrap();
14558 assert_eq!(agent.current_state().as_deref(), Some("active"));
14559
14560 agent.restore_state(snapshot).await.unwrap();
14561
14562 assert_eq!(agent.current_state().as_deref(), Some("active"));
14563 assert!(agent.pending_skill_id.read().is_none());
14564 assert!(
14565 !agent
14566 .disambiguation_manager()
14567 .unwrap()
14568 .has_pending_clarification()
14569 .await
14570 );
14571 assert_eq!(confirmation_skill_call_count(&observed), 0);
14572 }
14573
14574 #[tokio::test]
14576 async fn direct_state_generation_change_invalidates_confirmation() {
14577 let (agent, observed) = state_disambiguation_agent_with_skills(
14578 vec![
14579 r#"{"is_ambiguous":false,"confidence":0.99,"ambiguity_type":null,"reasoning":"top-level clear","what_is_unclear":[],"detected_language":"en"}"#,
14580 "send_report",
14581 r#"{"is_ambiguous":true,"confidence":0.2,"ambiguity_type":"missing_target","reasoning":"target missing","what_is_unclear":["target"],"detected_language":"en"}"#,
14582 r#"{"question":"What should I send?","options":null}"#,
14583 r#"{"status":"answered","selected_option":null,"enriched_input":"Send the report to Ada","resolved":{"intent":"send_report"}}"#,
14584 r#"{"question":"Should I send the report to Ada?"}"#,
14585 "Confirmation cancelled.",
14586 ],
14587 true,
14588 None,
14589 true,
14590 vec![confirmation_skill()],
14591 );
14592
14593 agent.chat("Send it").await.unwrap();
14594 agent.chat("The report to Ada").await.unwrap();
14595 let state_machine = agent.state_machine().unwrap();
14596 state_machine
14597 .transition_to("review", "external test")
14598 .unwrap();
14599 state_machine
14600 .transition_to("active", "external test")
14601 .unwrap();
14602
14603 let response = agent.chat("Yes").await.unwrap();
14604
14605 assert_eq!(response.content, "Confirmation cancelled.");
14606 assert_eq!(confirmation_skill_call_count(&observed), 0);
14607 assert!(agent.pending_skill_id.read().is_none());
14608 }
14609
14610 #[tokio::test]
14611 async fn state_confirmation_does_not_add_a_question_for_clear_input() {
14612 let (agent, observed) = state_disambiguation_agent(
14613 vec![
14614 r#"{"is_ambiguous":false,"confidence":0.99,"ambiguity_type":null,"reasoning":"clear","what_is_unclear":[],"detected_language":"en"}"#,
14615 "Request executed.",
14616 ],
14617 true,
14618 None,
14619 true,
14620 );
14621
14622 let response = agent.chat("Send the report to Ada").await.unwrap();
14623
14624 assert_eq!(response.content, "Request executed.");
14625 assert_eq!(observed.call_count(), 2);
14626 }
14627
14628 #[tokio::test]
14629 async fn state_override_cannot_activate_a_disabled_top_level_manager() {
14630 let (agent, observed) =
14631 state_disambiguation_agent(vec!["Request executed."], false, Some(true), true);
14632
14633 assert!(!agent.has_disambiguation());
14634 let response = agent.chat("Send it").await.unwrap();
14635
14636 assert_eq!(response.content, "Request executed.");
14637 assert_eq!(observed.call_count(), 1);
14638 }
14639
14640 #[tokio::test]
14641 async fn native_required_choice_executes_through_the_shared_tool_path() {
14642 let mut mock = MockLLMProvider::new("native-required");
14643 mock.set_tool_choice(Some(ToolChoice::Required));
14644 let native_call = ToolCall {
14645 id: "provider-call-1".to_string(),
14646 name: "calculator".to_string(),
14647 arguments: serde_json::json!({"expression": "2 + 2"}),
14648 };
14649 let provider_state = ai_agents_core::NativeProviderState::new(
14650 "fixture-exchange-1",
14651 "fixture",
14652 "native-tools",
14653 ai_agents_core::NativeProviderTarget::new("https://fixture.invalid/", "fixture-model")
14654 .unwrap(),
14655 serde_json::json!({
14656 "role": "model",
14657 "parts": [{
14658 "functionCall": {"name": "calculator", "args": {"expression": "2 + 2"}},
14659 "thoughtSignature": "fixture-signature"
14660 }]
14661 }),
14662 vec![ai_agents_core::NativeCallBinding::new("provider-call-1", 0).unwrap()],
14663 )
14664 .unwrap();
14665 mock.add_response(
14666 LLMResponse::new("", FinishReason::ToolCall)
14667 .with_provider_state(provider_state)
14668 .unwrap()
14669 .with_tool_calls(vec![native_call])
14670 .unwrap(),
14671 );
14672 mock.add_response(LLMResponse::new("The answer is 4.", FinishReason::Stop));
14673 let observed = mock.clone();
14674 let agent = AgentBuilder::new()
14675 .system_prompt("Use the calculator when needed.")
14676 .llm(Arc::new(mock))
14677 .tool(Arc::new(CalculatorTool::new()))
14678 .build()
14679 .unwrap();
14680
14681 let response = agent.chat("What is 2 + 2?").await.unwrap();
14682
14683 assert_eq!(response.content, "The answer is 4.");
14684 assert_eq!(
14685 response.tool_calls.as_ref().unwrap()[0].id,
14686 "provider-call-1"
14687 );
14688 let calls = observed.call_history();
14689 assert_eq!(calls.len(), 2);
14690 assert!(matches!(
14691 calls[0].request.as_ref().map(|request| &request.choice),
14692 Some(ToolChoice::Required)
14693 ));
14694 assert!(matches!(
14695 calls[1].request.as_ref().map(|request| &request.choice),
14696 Some(ToolChoice::Auto)
14697 ));
14698 let replay_batch = calls[1]
14699 .messages
14700 .iter()
14701 .find_map(|message| {
14702 ai_agents_core::decode_native_tool_call_markers(&message.content).unwrap()
14703 })
14704 .expect("signed native call marker must be replayed");
14705 assert_eq!(
14706 replay_batch.provider_state().unwrap().exchange_id(),
14707 "fixture-exchange-1"
14708 );
14709 assert!(calls[1].messages.iter().any(|message| {
14710 ai_agents_core::decode_native_tool_result_markers(&message.content)
14711 .is_ok_and(|results| results.is_some())
14712 }));
14713 }
14714
14715 #[tokio::test]
14716 async fn custom_memory_loss_stops_before_signed_tool_execution() {
14717 let mut mock = MockLLMProvider::new("native-custom-memory");
14718 mock.set_tool_choice(Some(ToolChoice::Required));
14719 let call = ToolCall {
14720 id: "provider-call-drop".to_string(),
14721 name: "calculator".to_string(),
14722 arguments: serde_json::json!({"expression": "3 + 4"}),
14723 };
14724 let state = ai_agents_core::NativeProviderState::new(
14725 "fixture-exchange-drop",
14726 "fixture",
14727 "native-tools",
14728 ai_agents_core::NativeProviderTarget::new("https://fixture.invalid/", "fixture-model")
14729 .unwrap(),
14730 serde_json::json!({
14731 "role": "model",
14732 "parts": [{
14733 "functionCall": {"name": "calculator", "args": {"expression": "3 + 4"}},
14734 "thoughtSignature": "fixture-signature-drop"
14735 }]
14736 }),
14737 vec![ai_agents_core::NativeCallBinding::new("provider-call-drop", 0).unwrap()],
14738 )
14739 .unwrap();
14740 mock.add_response(
14741 LLMResponse::new("", FinishReason::ToolCall)
14742 .with_provider_state(state)
14743 .unwrap()
14744 .with_tool_calls(vec![call])
14745 .unwrap(),
14746 );
14747 let agent = AgentBuilder::new()
14748 .system_prompt("Use the calculator.")
14749 .llm(Arc::new(mock))
14750 .memory(Arc::new(DroppingSignedAssistantMemory {
14751 messages: RwLock::new(Vec::new()),
14752 }))
14753 .tool(Arc::new(CalculatorTool::new()))
14754 .build()
14755 .unwrap();
14756
14757 let error = agent.chat("What is 3 + 4?").await.unwrap_err();
14758
14759 assert!(
14760 error
14761 .to_string()
14762 .contains("removed before provider continuation")
14763 );
14764 assert!(agent.tool_call_history.read().is_empty());
14765 }
14766
14767 #[tokio::test]
14768 async fn sequential_signed_history_validates_every_prior_exchange() {
14769 let mut mock = MockLLMProvider::new("native-sequential-memory");
14770 mock.set_tool_choice(Some(ToolChoice::Required));
14771 mock.add_response(signed_calculator_response(
14772 "seq-exchange-1",
14773 "seq-call-1",
14774 "1 + 1",
14775 ));
14776 mock.add_response(signed_calculator_response(
14777 "seq-exchange-2",
14778 "seq-call-2",
14779 "2 + 2",
14780 ));
14781 let agent = AgentBuilder::new()
14782 .system_prompt("Use the calculator sequentially.")
14783 .llm(Arc::new(mock))
14784 .memory(Arc::new(DroppingEarlierSequentialMemory {
14785 messages: RwLock::new(Vec::new()),
14786 signed_seen: std::sync::atomic::AtomicUsize::new(0),
14787 }))
14788 .tool(Arc::new(CalculatorTool::new()))
14789 .build()
14790 .unwrap();
14791
14792 let error = agent.chat("Calculate twice.").await.unwrap_err();
14793
14794 assert!(error.to_string().contains("seq-exchange-1"));
14795 assert_eq!(agent.tool_call_history.read().len(), 1);
14796 }
14797
14798 #[tokio::test]
14799 async fn post_transition_signed_hitl_rejection_stops_before_continuation() {
14800 let mut native = MockLLMProvider::new("post-transition-native");
14801 native.set_tool_choice(Some(ToolChoice::Auto));
14802 let call = ToolCall {
14803 id: "post-transition-call".to_string(),
14804 name: "echo".to_string(),
14805 arguments: serde_json::json!({"message": "hello"}),
14806 };
14807 let state = ai_agents_core::NativeProviderState::new(
14808 "post-transition-exchange",
14809 "fixture",
14810 "native-tools",
14811 ai_agents_core::NativeProviderTarget::new("https://fixture.invalid/", "fixture-model")
14812 .unwrap(),
14813 serde_json::json!({
14814 "role": "model",
14815 "parts": [{
14816 "functionCall": {"name": "echo", "args": {"message": "hello"}},
14817 "thoughtSignature": "post-transition-signature"
14818 }]
14819 }),
14820 vec![ai_agents_core::NativeCallBinding::new("post-transition-call", 0).unwrap()],
14821 )
14822 .unwrap();
14823 native.add_response(
14824 LLMResponse::new("", FinishReason::ToolCall)
14825 .with_provider_state(state)
14826 .unwrap()
14827 .with_tool_calls(vec![call])
14828 .unwrap(),
14829 );
14830 let observed_native = native.clone();
14831 let yaml = r#"
14832name: PostTransitionNativeReject
14833system_prompt: test
14834tools: [echo]
14835hitl:
14836 tools:
14837 echo:
14838 require_approval: true
14839states:
14840 initial: intake
14841 states:
14842 intake:
14843 prompt: intake
14844 transitions:
14845 - to: active
14846 guard:
14847 context:
14848 route:
14849 eq: active
14850 active:
14851 prompt: active
14852 llm: native
14853"#;
14854 let agent = AgentBuilder::from_yaml(yaml)
14855 .unwrap()
14856 .llm(Arc::new(mock_with_response("stale intake response")))
14857 .llm_alias("native", Arc::new(native))
14858 .auto_configure_features()
14859 .unwrap()
14860 .build()
14861 .unwrap();
14862 agent
14863 .set_context("route", serde_json::json!("active"))
14864 .unwrap();
14865
14866 let error = agent.chat("move to active").await.unwrap_err();
14867
14868 assert!(matches!(error, AgentError::HITLRejected(_)));
14869 assert_eq!(observed_native.call_count(), 1);
14870 }
14871
14872 #[test]
14873 fn runtime_overflow_removes_a_past_signed_user_turn_as_one_prefix() {
14874 let call = ToolCall {
14875 id: "overflow-call".to_string(),
14876 name: "calculator".to_string(),
14877 arguments: serde_json::json!({"expression": "1 + 1"}),
14878 };
14879 let state = ai_agents_core::NativeProviderState::new(
14880 "overflow-exchange",
14881 "google",
14882 "generateContent",
14883 ai_agents_core::NativeProviderTarget::new("https://example.invalid/", "gemini-3")
14884 .unwrap(),
14885 serde_json::json!({
14886 "role": "model",
14887 "parts": [{
14888 "functionCall": {"name": "calculator", "args": {"expression": "1 + 1"}},
14889 "thoughtSignature": "overflow-signature"
14890 }]
14891 }),
14892 vec![ai_agents_core::NativeCallBinding::new("overflow-call", 0).unwrap()],
14893 )
14894 .unwrap();
14895 let call_marker = ai_agents_core::encode_native_tool_call_markers(
14896 std::slice::from_ref(&call),
14897 Some(&state),
14898 )
14899 .unwrap();
14900 let result_marker = ai_agents_core::encode_native_tool_result_marker(
14901 &call,
14902 serde_json::json!({"result": 2}),
14903 )
14904 .unwrap();
14905 let history = vec![
14906 ChatMessage::user("old question"),
14907 ChatMessage::assistant(call_marker),
14908 ChatMessage::function("calculator", result_marker),
14909 ChatMessage::assistant("old answer"),
14910 ChatMessage::user("new question"),
14911 ];
14912
14913 let removable = RuntimeAgent::native_safe_prefix_at_least(&history, 1).unwrap();
14914
14915 assert_eq!(removable, 4);
14916 }
14917
14918 #[test]
14919 fn auxiliary_projection_does_not_interpret_user_marker_text() {
14920 let user_text = serde_json::json!({
14921 "_ai_agents_native_tool_call": true,
14922 "id": "",
14923 "tool": "user-data",
14924 "arguments": {}
14925 })
14926 .to_string();
14927
14928 let projected =
14929 RuntimeAgent::readable_native_messages(vec![ChatMessage::user(&user_text)]).unwrap();
14930
14931 assert_eq!(projected[0].content, user_text);
14932 }
14933
14934 #[tokio::test]
14935 async fn terminal_provider_history_error_skips_retry_and_static_fallback() {
14936 let calls = Arc::new(std::sync::atomic::AtomicU32::new(0));
14937 let recovery = RecoveryManager::new(ai_agents_recovery::ErrorRecoveryConfig {
14938 default: ai_agents_recovery::RetryConfig {
14939 max_retries: 3,
14940 ..Default::default()
14941 },
14942 llm: ai_agents_recovery::LLMRecoveryConfig {
14943 on_failure: LLMFailureAction::FallbackResponse {
14944 message: "must not be returned".to_string(),
14945 },
14946 ..Default::default()
14947 },
14948 ..Default::default()
14949 });
14950 let agent = AgentBuilder::new()
14951 .system_prompt("Reject corrupted native history.")
14952 .llm(Arc::new(TerminalHistoryProvider {
14953 calls: Arc::clone(&calls),
14954 }))
14955 .recovery_manager(recovery)
14956 .build()
14957 .unwrap();
14958
14959 let error = agent.chat("continue").await.unwrap_err();
14960
14961 assert!(
14962 error
14963 .to_string()
14964 .contains("native history integrity failure")
14965 );
14966 assert_eq!(calls.load(Ordering::SeqCst), 1);
14967 }
14968
14969 #[tokio::test]
14970 async fn prompt_fallback_uses_one_corrective_retry() {
14971 let mut mock = MockLLMProvider::new("prompt-required");
14972 mock.set_tool_choice(Some(ToolChoice::Required));
14973 mock.set_native_tool_support(false);
14974 mock.set_responses(
14975 vec![
14976 "I can calculate that.".to_string(),
14977 r#"{"tool":"calculator","arguments":{"expression":"2 + 2"}}"#.to_string(),
14978 "The answer is 4.".to_string(),
14979 ],
14980 false,
14981 );
14982 let observed = mock.clone();
14983 let agent = AgentBuilder::new()
14984 .system_prompt("Use tools.")
14985 .llm(Arc::new(mock))
14986 .tool(Arc::new(CalculatorTool::new()))
14987 .build()
14988 .unwrap();
14989
14990 let response = agent.chat("What is 2 + 2?").await.unwrap();
14991
14992 assert_eq!(response.content, "The answer is 4.");
14993 assert_eq!(observed.call_count(), 3);
14994 let corrective = &observed.call_history()[1].messages;
14995 assert!(
14996 corrective
14997 .last()
14998 .unwrap()
14999 .content
15000 .contains("previous response")
15001 );
15002 }
15003
15004 #[tokio::test]
15005 async fn prompt_fallback_fails_after_one_noncompliant_retry() {
15006 let mut mock = MockLLMProvider::new("prompt-required-failure");
15007 mock.set_tool_choice(Some(ToolChoice::Required));
15008 mock.set_native_tool_support(false);
15009 mock.set_responses(
15010 vec!["No tool.".to_string(), "Still no tool.".to_string()],
15011 false,
15012 );
15013 let observed = mock.clone();
15014 let agent = AgentBuilder::new()
15015 .system_prompt("Use tools.")
15016 .llm(Arc::new(mock))
15017 .tool(Arc::new(CalculatorTool::new()))
15018 .build()
15019 .unwrap();
15020
15021 let error = agent.chat("What is 2 + 2?").await.unwrap_err();
15022
15023 assert!(error.to_string().contains("one corrective retry"));
15024 assert_eq!(observed.call_count(), 2);
15025 }
15026
15027 #[tokio::test]
15028 async fn specific_choice_cannot_widen_the_effective_grant() {
15029 let mut mock = MockLLMProvider::new("specific-outside-grant");
15030 mock.set_tool_choice(Some(ToolChoice::Specific("random".to_string())));
15031 let observed = mock.clone();
15032 let agent = AgentBuilder::new()
15033 .system_prompt("Use tools.")
15034 .llm(Arc::new(mock))
15035 .tool(Arc::new(CalculatorTool::new()))
15036 .build()
15037 .unwrap();
15038
15039 let error = agent.chat("Generate a value.").await.unwrap_err();
15040
15041 assert!(error.to_string().contains("is not registered"));
15042 assert_eq!(observed.call_count(), 0);
15043 }
15044
15045 #[tokio::test]
15046 async fn none_choice_exposes_no_tool_protocol() {
15047 let mut mock = MockLLMProvider::new("no-tools");
15048 mock.set_tool_choice(Some(ToolChoice::None));
15049 mock.set_response(r#"{"tool":"calculator","arguments":{"expression":"2 + 2"}}"#);
15050 let observed = mock.clone();
15051 let agent = AgentBuilder::new()
15052 .system_prompt("Answer directly.")
15053 .llm(Arc::new(mock))
15054 .tool(Arc::new(CalculatorTool::new()))
15055 .build()
15056 .unwrap();
15057
15058 let response = agent.chat("Hello").await.unwrap();
15059
15060 assert!(response.tool_calls.is_none());
15061 assert_eq!(observed.call_count(), 1);
15062 let call = observed.last_call().unwrap();
15063 assert!(call.request.is_none());
15064 assert!(
15065 call.messages
15066 .iter()
15067 .all(|message| !message.content.contains("Available tools:"))
15068 );
15069 }
15070
15071 struct RuntimeStorage {
15072 capabilities: Box<[StorageCapability]>,
15073 snapshots: RwLock<HashMap<String, AgentSnapshot>>,
15074 metadata: RwLock<HashMap<String, ai_agents_core::SessionMetadata>>,
15075 metadata_save_calls: AtomicU64,
15076 metadata_load_calls: AtomicU64,
15077 fail_metadata_save: AtomicBool,
15078 fail_metadata_load: AtomicBool,
15079 }
15080
15081 impl RuntimeStorage {
15082 fn new(capabilities: impl IntoIterator<Item = StorageCapability>) -> Self {
15083 Self {
15084 capabilities: capabilities.into_iter().collect(),
15085 snapshots: RwLock::new(HashMap::new()),
15086 metadata: RwLock::new(HashMap::new()),
15087 metadata_save_calls: AtomicU64::new(0),
15088 metadata_load_calls: AtomicU64::new(0),
15089 fail_metadata_save: AtomicBool::new(false),
15090 fail_metadata_load: AtomicBool::new(false),
15091 }
15092 }
15093 }
15094
15095 #[async_trait]
15096 impl AgentStorage for RuntimeStorage {
15097 fn supports(&self, capability: StorageCapability) -> bool {
15098 self.capabilities.contains(&capability)
15099 }
15100
15101 async fn save(&self, session_id: &str, snapshot: &AgentSnapshot) -> Result<()> {
15102 self.snapshots
15103 .write()
15104 .insert(session_id.to_string(), snapshot.clone());
15105 Ok(())
15106 }
15107
15108 async fn load(&self, session_id: &str) -> Result<Option<AgentSnapshot>> {
15109 Ok(self.snapshots.read().get(session_id).cloned())
15110 }
15111
15112 async fn delete(&self, session_id: &str) -> Result<()> {
15113 self.snapshots.write().remove(session_id);
15114 Ok(())
15115 }
15116
15117 async fn list_sessions(&self) -> Result<Vec<String>> {
15118 Ok(self.snapshots.read().keys().cloned().collect())
15119 }
15120
15121 async fn save_snapshot_with_metadata(
15122 &self,
15123 session_id: &str,
15124 snapshot: &AgentSnapshot,
15125 metadata: &ai_agents_core::SessionMetadata,
15126 ) -> Result<()> {
15127 self.metadata_save_calls.fetch_add(1, Ordering::SeqCst);
15128 if self.fail_metadata_save.load(Ordering::SeqCst) {
15129 return Err(AgentError::Persistence("metadata save failed".into()));
15130 }
15131 self.snapshots
15132 .write()
15133 .insert(session_id.to_string(), snapshot.clone());
15134 self.metadata
15135 .write()
15136 .insert(session_id.to_string(), metadata.clone());
15137 Ok(())
15138 }
15139
15140 async fn save_metadata(
15141 &self,
15142 session_id: &str,
15143 metadata: &ai_agents_core::SessionMetadata,
15144 ) -> Result<()> {
15145 self.metadata_save_calls.fetch_add(1, Ordering::SeqCst);
15146 if self.fail_metadata_save.load(Ordering::SeqCst) {
15147 return Err(AgentError::Persistence("metadata save failed".into()));
15148 }
15149 self.metadata
15150 .write()
15151 .insert(session_id.to_string(), metadata.clone());
15152 Ok(())
15153 }
15154
15155 async fn load_metadata(
15156 &self,
15157 session_id: &str,
15158 ) -> Result<Option<ai_agents_core::SessionMetadata>> {
15159 self.metadata_load_calls.fetch_add(1, Ordering::SeqCst);
15160 if self.fail_metadata_load.load(Ordering::SeqCst) {
15161 return Err(AgentError::Persistence("metadata load failed".into()));
15162 }
15163 Ok(self.metadata.read().get(session_id).cloned())
15164 }
15165 }
15166
15167 fn runtime_storage_agent() -> RuntimeAgent {
15168 AgentBuilder::new()
15169 .system_prompt("Test runtime storage integration.")
15170 .llm(Arc::new(mock_with_response("done")))
15171 .build()
15172 .unwrap()
15173 }
15174
15175 fn restore_spec(id: &str) -> crate::spec::AgentSpec {
15176 crate::spec::AgentSpec {
15177 name: id.to_string(),
15178 system_prompt: format!("Restore child {id}."),
15179 ..crate::spec::AgentSpec::default()
15180 }
15181 }
15182
15183 fn restore_entry(id: &str) -> ai_agents_core::SpawnedAgentEntry {
15184 ai_agents_core::SpawnedAgentEntry {
15185 id: id.to_string(),
15186 name: id.to_string(),
15187 spec_yaml: serde_yaml::to_string(&restore_spec(id)).unwrap(),
15188 }
15189 }
15190
15191 fn restore_spawner(
15192 storage: Arc<RuntimeStorage>,
15193 max_agents: usize,
15194 ) -> (
15195 Arc<crate::spawner::AgentSpawner>,
15196 Arc<crate::spawner::AgentRegistry>,
15197 ) {
15198 let mut llms = LLMRegistry::new();
15199 llms.register("default", Arc::new(mock_with_response("done")));
15200 (
15201 Arc::new(
15202 crate::spawner::AgentSpawner::new()
15203 .with_shared_llms(llms)
15204 .with_shared_storage(storage)
15205 .with_max_agents(max_agents),
15206 ),
15207 Arc::new(crate::spawner::AgentRegistry::new()),
15208 )
15209 }
15210
15211 async fn save_restore_target(
15212 parent: &RuntimeAgent,
15213 storage: &RuntimeStorage,
15214 session_id: &str,
15215 entries: Vec<ai_agents_core::SpawnedAgentEntry>,
15216 ) {
15217 let mut snapshot = parent.save_state().await.unwrap();
15218 snapshot.spawned_agents = Some(entries);
15219 storage.save(session_id, &snapshot).await.unwrap();
15220 storage
15221 .save_metadata(session_id, &ai_agents_core::SessionMetadata::default())
15222 .await
15223 .unwrap();
15224 }
15225
15226 #[tokio::test]
15227 async fn storage_init_requires_storage_for_actor_facts() {
15228 let facts = ai_agents_facts::FactsConfig {
15229 enabled: true,
15230 ..Default::default()
15231 };
15232 let agent = runtime_storage_agent().with_facts_config(None, Some(facts));
15233
15234 let error = agent.init_storage().await.unwrap_err();
15235 assert!(matches!(
15236 error,
15237 AgentError::Config(message)
15238 if message.contains("actor facts or actor memory")
15239 && message.contains("none is configured or injected")
15240 ));
15241 }
15242
15243 #[tokio::test]
15244 async fn storage_init_validates_actor_facts_capability() {
15245 let storage = Arc::new(RuntimeStorage::new([StorageCapability::Snapshot]));
15246 let actor_memory = ai_agents_facts::ActorMemoryConfig {
15247 enabled: true,
15248 ..Default::default()
15249 };
15250 let agent = runtime_storage_agent()
15251 .with_storage(storage)
15252 .with_facts_config(Some(actor_memory), None);
15253
15254 assert!(matches!(
15255 agent.init_storage().await,
15256 Err(AgentError::UnsupportedStorageCapability(
15257 StorageCapability::ActorFacts
15258 ))
15259 ));
15260 }
15261
15262 #[tokio::test]
15263 async fn blocking_chat_rejects_unsupported_required_storage() {
15264 let storage = Arc::new(RuntimeStorage::new([StorageCapability::Snapshot]));
15265 let facts = ai_agents_facts::FactsConfig {
15266 enabled: true,
15267 ..Default::default()
15268 };
15269 let agent = runtime_storage_agent()
15270 .with_storage(storage)
15271 .with_facts_config(None, Some(facts));
15272
15273 assert!(matches!(
15274 agent.chat("hello").await,
15275 Err(AgentError::UnsupportedStorageCapability(
15276 StorageCapability::ActorFacts
15277 ))
15278 ));
15279 }
15280
15281 #[tokio::test]
15282 async fn streaming_chat_rejects_unsupported_required_storage_before_stream_creation() {
15283 let storage = Arc::new(RuntimeStorage::new([StorageCapability::Snapshot]));
15284 let config = ai_agents_relationships::RelationshipConfig {
15285 enabled: true,
15286 ..Default::default()
15287 };
15288 let manager = Arc::new(RelationshipManager::from_config(config).unwrap());
15289 let agent = runtime_storage_agent()
15290 .with_storage(storage)
15291 .with_relationships(manager);
15292
15293 assert!(matches!(
15294 agent.chat_stream("hello").await,
15295 Err(AgentError::UnsupportedStorageCapability(
15296 StorageCapability::ActorRelationships
15297 ))
15298 ));
15299 }
15300
15301 #[tokio::test]
15302 async fn storage_init_completes_facts_for_injected_storage() {
15303 let storage = Arc::new(RuntimeStorage::new([
15304 StorageCapability::Snapshot,
15305 StorageCapability::ActorFacts,
15306 ]));
15307 let facts = ai_agents_facts::FactsConfig {
15308 enabled: true,
15309 ..Default::default()
15310 };
15311 let agent = runtime_storage_agent()
15312 .with_storage(storage)
15313 .with_facts_config(None, Some(facts));
15314
15315 agent.init_storage().await.unwrap();
15316 assert!(agent.fact_store().is_some());
15317 }
15318
15319 #[tokio::test]
15320 async fn storage_init_requires_storage_for_persistent_relationships() {
15321 let config = ai_agents_relationships::RelationshipConfig {
15322 enabled: true,
15323 ..Default::default()
15324 };
15325 let manager = Arc::new(RelationshipManager::from_config(config).unwrap());
15326 let agent = runtime_storage_agent().with_relationships(manager);
15327
15328 let error = agent.init_storage().await.unwrap_err();
15329 assert!(matches!(
15330 error,
15331 AgentError::Config(message)
15332 if message.contains("persistent relationships")
15333 && message.contains("none is configured or injected")
15334 ));
15335 }
15336
15337 #[tokio::test]
15338 async fn storage_init_validates_persistent_relationships_capability() {
15339 let storage = Arc::new(RuntimeStorage::new([StorageCapability::Snapshot]));
15340 let config = ai_agents_relationships::RelationshipConfig {
15341 enabled: true,
15342 ..Default::default()
15343 };
15344 let manager = Arc::new(RelationshipManager::from_config(config).unwrap());
15345 let agent = runtime_storage_agent()
15346 .with_storage(storage)
15347 .with_relationships(manager);
15348
15349 assert!(matches!(
15350 agent.init_storage().await,
15351 Err(AgentError::UnsupportedStorageCapability(
15352 StorageCapability::ActorRelationships
15353 ))
15354 ));
15355 }
15356
15357 #[tokio::test]
15358 async fn session_restore_updates_identity_and_clears_stale_actor_binding() {
15359 let storage = Arc::new(RuntimeStorage::new([
15360 StorageCapability::Snapshot,
15361 StorageCapability::SessionMetadata,
15362 ]));
15363 let agent = runtime_storage_agent().with_storage(storage.clone());
15364 agent.set_actor_id("old-actor").unwrap();
15365 agent.save_session("old").await.unwrap();
15366 storage
15367 .save("target", &agent.save_state().await.unwrap())
15368 .await
15369 .unwrap();
15370 storage
15371 .save_metadata("target", &ai_agents_core::SessionMetadata::default())
15372 .await
15373 .unwrap();
15374
15375 assert!(agent.load_session("target").await.unwrap());
15376
15377 assert_eq!(agent.current_session_id.read().as_deref(), Some("target"));
15378 assert_eq!(agent.actor_id(), None);
15379 }
15380
15381 #[tokio::test]
15382 async fn complete_restore_reconciles_growth_shrink_and_empty_topologies() {
15383 let storage = Arc::new(RuntimeStorage::new([
15384 StorageCapability::Snapshot,
15385 StorageCapability::SessionMetadata,
15386 ]));
15387 let (spawner, registry) = restore_spawner(storage.clone(), 3);
15388 let parent = runtime_storage_agent()
15389 .with_storage(storage.clone())
15390 .with_spawner_handles(Arc::clone(&spawner), Arc::clone(®istry));
15391
15392 for id in ["a", "b"] {
15393 let spawned = spawner
15394 .spawn_with_id(id.to_string(), restore_spec(id))
15395 .await
15396 .unwrap();
15397 spawned.agent.save_session("grow").await.unwrap();
15398 registry.register(spawned).await.unwrap();
15399 }
15400 let staged_c = crate::spawner::storage::NamespacedStorage::new(storage.clone(), "c");
15401 staged_c
15402 .save("grow", &AgentSnapshot::new("c".into()))
15403 .await
15404 .unwrap();
15405 staged_c
15406 .save_metadata("grow", &ai_agents_core::SessionMetadata::default())
15407 .await
15408 .unwrap();
15409 save_restore_target(
15410 &parent,
15411 storage.as_ref(),
15412 "grow",
15413 vec![restore_entry("a"), restore_entry("b"), restore_entry("c")],
15414 )
15415 .await;
15416
15417 assert_eq!(parent.restore_session_full("grow").await.unwrap(), 3);
15418 assert_eq!(registry.count(), 3);
15419 assert!(registry.contains("c"));
15420 assert_eq!(spawner.spawned_count(), 3);
15421
15422 for id in ["a", "b"] {
15423 registry
15424 .get(id)
15425 .unwrap()
15426 .save_session("shrink")
15427 .await
15428 .unwrap();
15429 }
15430 save_restore_target(
15431 &parent,
15432 storage.as_ref(),
15433 "shrink",
15434 vec![restore_entry("a"), restore_entry("b")],
15435 )
15436 .await;
15437
15438 assert_eq!(parent.restore_session_full("shrink").await.unwrap(), 2);
15439 assert_eq!(registry.count(), 2);
15440 assert!(!registry.contains("c"));
15441 assert_eq!(spawner.spawned_count(), 2);
15442
15443 save_restore_target(&parent, storage.as_ref(), "empty", Vec::new()).await;
15444
15445 assert_eq!(parent.restore_session_full("empty").await.unwrap(), 0);
15446 assert_eq!(registry.count(), 0);
15447 assert_eq!(spawner.spawned_count(), 0);
15448 assert_eq!(parent.current_session_id.read().as_deref(), Some("empty"));
15449 }
15450
15451 #[tokio::test]
15452 async fn storage_session_metadata_is_called_only_when_advertised() {
15453 let storage = Arc::new(RuntimeStorage::new([StorageCapability::Snapshot]));
15454 storage.fail_metadata_save.store(true, Ordering::SeqCst);
15455 storage.fail_metadata_load.store(true, Ordering::SeqCst);
15456 let agent = runtime_storage_agent().with_storage(storage.clone());
15457
15458 agent.save_session("session").await.unwrap();
15459 assert!(agent.load_session("session").await.unwrap());
15460 assert_eq!(storage.metadata_save_calls.load(Ordering::SeqCst), 0);
15461 assert_eq!(storage.metadata_load_calls.load(Ordering::SeqCst), 0);
15462 }
15463
15464 #[cfg(feature = "sqlite")]
15465 #[tokio::test]
15466 async fn sqlite_runtime_save_filter_reopen_and_reload_stay_consistent() {
15467 let directory =
15468 std::env::temp_dir().join(format!("ai-agents-runtime-sqlite-{}", uuid::Uuid::new_v4()));
15469 let path = directory.join("sessions.sqlite");
15470 let path_string = path.to_string_lossy().into_owned();
15471 let storage = Arc::new(
15472 ai_agents_storage::SqliteStorage::new(&path_string)
15473 .await
15474 .unwrap(),
15475 );
15476 let agent = runtime_storage_agent().with_storage(storage.clone());
15477 agent.set_session_metadata(ai_agents_core::SessionMetadata {
15478 tags: vec!["initial".into()],
15479 ..Default::default()
15480 });
15481 agent.chat("persist this turn").await.unwrap();
15482 agent.save_session("session").await.unwrap();
15483
15484 agent.set_session_metadata(ai_agents_core::SessionMetadata {
15485 tags: vec!["updated".into()],
15486 ..Default::default()
15487 });
15488 agent.save_session("session").await.unwrap();
15489 assert!(
15490 agent
15491 .list_sessions_filtered(&ai_agents_core::SessionFilter {
15492 tags: Some(vec!["initial".into()]),
15493 ..Default::default()
15494 })
15495 .await
15496 .unwrap()
15497 .is_empty()
15498 );
15499 assert_eq!(
15500 agent
15501 .list_sessions_filtered(&ai_agents_core::SessionFilter {
15502 tags: Some(vec!["updated".into()]),
15503 ..Default::default()
15504 })
15505 .await
15506 .unwrap()
15507 .len(),
15508 1
15509 );
15510 drop(agent);
15511 storage.close().await;
15512 drop(storage);
15513
15514 let reopened_storage = Arc::new(
15515 ai_agents_storage::SqliteStorage::new(&path_string)
15516 .await
15517 .unwrap(),
15518 );
15519 let restored = runtime_storage_agent().with_storage(reopened_storage.clone());
15520 assert!(restored.load_session("session").await.unwrap());
15521 assert_eq!(restored.session_metadata().tags, vec!["updated"]);
15522 assert_eq!(
15523 restored.current_session_id.read().as_deref(),
15524 Some("session")
15525 );
15526 assert!(restored.save_state().await.unwrap().memory.messages.len() >= 2);
15527 assert_eq!(
15528 restored
15529 .list_sessions_filtered(&ai_agents_core::SessionFilter {
15530 tags: Some(vec!["updated".into()]),
15531 ..Default::default()
15532 })
15533 .await
15534 .unwrap()
15535 .len(),
15536 1
15537 );
15538
15539 drop(restored);
15540 reopened_storage.close().await;
15541 drop(reopened_storage);
15542 crate::remove_sqlite_test_directory(&directory)
15543 .await
15544 .unwrap();
15545 }
15546
15547 #[tokio::test]
15548 async fn storage_session_metadata_backend_failures_propagate() {
15549 let storage = Arc::new(RuntimeStorage::new([
15550 StorageCapability::Snapshot,
15551 StorageCapability::SessionMetadata,
15552 ]));
15553 let agent = runtime_storage_agent().with_storage(storage.clone());
15554
15555 agent.save_session("session").await.unwrap();
15556 storage
15557 .save("target", &agent.save_state().await.unwrap())
15558 .await
15559 .unwrap();
15560 storage.fail_metadata_load.store(true, Ordering::SeqCst);
15561 assert!(matches!(
15562 agent.load_session("target").await,
15563 Err(AgentError::Persistence(message)) if message == "metadata load failed"
15564 ));
15565 assert_eq!(agent.current_session_id.read().as_deref(), Some("session"));
15566
15567 storage.fail_metadata_save.store(true, Ordering::SeqCst);
15568 assert!(matches!(
15569 agent.save_session("session").await,
15570 Err(AgentError::Persistence(message)) if message == "metadata save failed"
15571 ));
15572 }
15573
15574 struct ProviderFutureDropSignal {
15575 dropped: Arc<AtomicBool>,
15576 }
15577
15578 impl Drop for ProviderFutureDropSignal {
15579 fn drop(&mut self) {
15580 self.dropped.store(true, Ordering::SeqCst);
15581 }
15582 }
15583
15584 struct BufferedLockingProvider {
15585 lock: Arc<tokio::sync::Mutex<()>>,
15586 stream_started: Arc<tokio::sync::Notify>,
15587 stream_dropped: Arc<AtomicBool>,
15588 committed_after_drop: Arc<AtomicBool>,
15589 }
15590
15591 #[async_trait]
15592 impl LLMProvider for BufferedLockingProvider {
15593 async fn complete(
15594 &self,
15595 _messages: &[ChatMessage],
15596 _config: Option<&LLMConfig>,
15597 ) -> std::result::Result<LLMResponse, LLMError> {
15598 let _guard = self.lock.lock().await;
15599 self.committed_after_drop
15600 .store(self.stream_dropped.load(Ordering::SeqCst), Ordering::SeqCst);
15601 Ok(LLMResponse::new(
15602 "Committed technical response.",
15603 FinishReason::Stop,
15604 ))
15605 }
15606
15607 async fn complete_stream(
15608 &self,
15609 _messages: &[ChatMessage],
15610 _config: Option<&LLMConfig>,
15611 ) -> std::result::Result<
15612 Box<dyn Stream<Item = std::result::Result<LLMChunk, LLMError>> + Unpin + Send>,
15613 LLMError,
15614 > {
15615 let _guard = self.lock.lock().await;
15616 let _drop_signal = ProviderFutureDropSignal {
15617 dropped: Arc::clone(&self.stream_dropped),
15618 };
15619 self.stream_started.notify_one();
15620 std::future::pending().await
15621 }
15622
15623 fn provider_name(&self) -> &str {
15624 "buffered-locking"
15625 }
15626
15627 fn supports(&self, _feature: LLMFeature) -> bool {
15628 false
15629 }
15630 }
15631
15632 struct PendingDropStream {
15633 dropped: Arc<AtomicBool>,
15634 dropped_notify: Arc<tokio::sync::Notify>,
15635 }
15636
15637 impl Stream for PendingDropStream {
15638 type Item = std::result::Result<LLMChunk, LLMError>;
15639
15640 fn poll_next(
15641 self: Pin<&mut Self>,
15642 _cx: &mut std::task::Context<'_>,
15643 ) -> std::task::Poll<Option<Self::Item>> {
15644 std::task::Poll::Pending
15645 }
15646 }
15647
15648 impl Drop for PendingDropStream {
15649 fn drop(&mut self) {
15650 self.dropped.store(true, Ordering::SeqCst);
15651 self.dropped_notify.notify_one();
15652 }
15653 }
15654
15655 struct EstablishedStreamProvider {
15656 stream_started: Arc<tokio::sync::Notify>,
15657 stream_dropped: Arc<AtomicBool>,
15658 stream_dropped_notify: Arc<tokio::sync::Notify>,
15659 committed_after_drop: Arc<AtomicBool>,
15660 }
15661
15662 #[async_trait]
15663 impl LLMProvider for EstablishedStreamProvider {
15664 async fn complete(
15665 &self,
15666 _messages: &[ChatMessage],
15667 _config: Option<&LLMConfig>,
15668 ) -> std::result::Result<LLMResponse, LLMError> {
15669 if !self.stream_dropped.load(Ordering::SeqCst) {
15670 self.stream_dropped_notify.notified().await;
15671 }
15672 self.committed_after_drop
15673 .store(self.stream_dropped.load(Ordering::SeqCst), Ordering::SeqCst);
15674 Ok(LLMResponse::new(
15675 "Committed technical response.",
15676 FinishReason::Stop,
15677 ))
15678 }
15679
15680 async fn complete_stream(
15681 &self,
15682 _messages: &[ChatMessage],
15683 _config: Option<&LLMConfig>,
15684 ) -> std::result::Result<
15685 Box<dyn Stream<Item = std::result::Result<LLMChunk, LLMError>> + Unpin + Send>,
15686 LLMError,
15687 > {
15688 self.stream_started.notify_one();
15689 Ok(Box::new(PendingDropStream {
15690 dropped: Arc::clone(&self.stream_dropped),
15691 dropped_notify: Arc::clone(&self.stream_dropped_notify),
15692 }))
15693 }
15694
15695 fn provider_name(&self) -> &str {
15696 "established-stream"
15697 }
15698
15699 fn supports(&self, _feature: LLMFeature) -> bool {
15700 false
15701 }
15702 }
15703
15704 struct FirstCallLockingProvider {
15705 lock: Arc<tokio::sync::Mutex<()>>,
15706 first_started: Arc<tokio::sync::Notify>,
15707 first_dropped: Arc<AtomicBool>,
15708 committed_after_drop: Arc<AtomicBool>,
15709 calls: AtomicU64,
15710 }
15711
15712 #[async_trait]
15713 impl LLMProvider for FirstCallLockingProvider {
15714 async fn complete(
15715 &self,
15716 _messages: &[ChatMessage],
15717 _config: Option<&LLMConfig>,
15718 ) -> std::result::Result<LLMResponse, LLMError> {
15719 let _guard = self.lock.lock().await;
15720 let call = self.calls.fetch_add(1, Ordering::SeqCst);
15721 if call == 0 {
15722 let _drop_signal = ProviderFutureDropSignal {
15723 dropped: Arc::clone(&self.first_dropped),
15724 };
15725 self.first_started.notify_one();
15726 return std::future::pending().await;
15727 }
15728 self.committed_after_drop
15729 .store(self.first_dropped.load(Ordering::SeqCst), Ordering::SeqCst);
15730 Ok(LLMResponse::new(
15731 "Committed technical response.",
15732 FinishReason::Stop,
15733 ))
15734 }
15735
15736 async fn complete_stream(
15737 &self,
15738 _messages: &[ChatMessage],
15739 _config: Option<&LLMConfig>,
15740 ) -> std::result::Result<
15741 Box<dyn Stream<Item = std::result::Result<LLMChunk, LLMError>> + Unpin + Send>,
15742 LLMError,
15743 > {
15744 Err(LLMError::Other(
15745 "streaming is not used in this test".to_string(),
15746 ))
15747 }
15748
15749 fn provider_name(&self) -> &str {
15750 "first-call-locking"
15751 }
15752
15753 fn supports(&self, _feature: LLMFeature) -> bool {
15754 false
15755 }
15756 }
15757
15758 struct RoutingAfterProviderStart {
15759 provider_started: Arc<tokio::sync::Notify>,
15760 }
15761
15762 #[async_trait]
15763 impl LLMProvider for RoutingAfterProviderStart {
15764 async fn complete(
15765 &self,
15766 _messages: &[ChatMessage],
15767 _config: Option<&LLMConfig>,
15768 ) -> std::result::Result<LLMResponse, LLMError> {
15769 self.provider_started.notified().await;
15770 Ok(LLMResponse::new("1", FinishReason::Stop))
15771 }
15772
15773 async fn complete_stream(
15774 &self,
15775 _messages: &[ChatMessage],
15776 _config: Option<&LLMConfig>,
15777 ) -> std::result::Result<
15778 Box<dyn Stream<Item = std::result::Result<LLMChunk, LLMError>> + Unpin + Send>,
15779 LLMError,
15780 > {
15781 Err(LLMError::Other(
15782 "streaming is not used in this test".to_string(),
15783 ))
15784 }
15785
15786 fn provider_name(&self) -> &str {
15787 "routing-after-start"
15788 }
15789
15790 fn supports(&self, _feature: LLMFeature) -> bool {
15791 false
15792 }
15793 }
15794
15795 struct ResponseCountingHooks {
15797 responses: Arc<std::sync::atomic::AtomicUsize>,
15798 }
15799
15800 struct RootTurnProbeProvider {
15802 complete_entered: tokio::sync::mpsc::UnboundedSender<()>,
15803 }
15804
15805 struct ResponseChatHooks {
15807 target: parking_lot::Mutex<Option<Weak<RuntimeAgent>>>,
15808 invoked: AtomicBool,
15809 nested_result: parking_lot::Mutex<Option<std::result::Result<String, String>>>,
15810 }
15811
15812 struct ConcurrentResponseHooks {
15814 registry: Weak<crate::spawner::AgentRegistry>,
15815 child_id: String,
15816 invoked: AtomicBool,
15817 nested_result: parking_lot::Mutex<Option<std::result::Result<String, String>>>,
15818 }
15819
15820 struct RetryDeadlineTool {
15822 calls: Arc<std::sync::atomic::AtomicUsize>,
15823 deadlines: Arc<parking_lot::Mutex<Vec<chrono::DateTime<chrono::Utc>>>>,
15824 remaining_ms: Arc<parking_lot::Mutex<Vec<i64>>>,
15825 }
15826
15827 struct ToolLifecycleRecordingHooks {
15829 events: parking_lot::Mutex<Vec<String>>,
15830 records: parking_lot::Mutex<Vec<ToolExecutionRecord>>,
15831 }
15832
15833 impl ToolLifecycleRecordingHooks {
15834 fn new() -> Self {
15836 Self {
15837 events: parking_lot::Mutex::new(Vec::new()),
15838 records: parking_lot::Mutex::new(Vec::new()),
15839 }
15840 }
15841
15842 fn events(&self) -> Vec<String> {
15844 self.events.lock().clone()
15845 }
15846
15847 fn records(&self) -> Vec<ToolExecutionRecord> {
15849 self.records.lock().clone()
15850 }
15851 }
15852
15853 struct ContextEchoTool;
15855
15856 #[async_trait]
15857 impl LLMProvider for RootTurnProbeProvider {
15858 async fn complete(
15859 &self,
15860 _messages: &[ChatMessage],
15861 _config: Option<&LLMConfig>,
15862 ) -> std::result::Result<LLMResponse, LLMError> {
15863 let _ = self.complete_entered.send(());
15864 Ok(LLMResponse::new("blocking complete", FinishReason::Stop))
15865 }
15866
15867 async fn complete_stream(
15868 &self,
15869 _messages: &[ChatMessage],
15870 _config: Option<&LLMConfig>,
15871 ) -> std::result::Result<
15872 Box<dyn Stream<Item = std::result::Result<LLMChunk, LLMError>> + Unpin + Send>,
15873 LLMError,
15874 > {
15875 Ok(Box::new(futures::stream::iter(vec![Ok(
15876 LLMChunk::final_chunk("stream complete", FinishReason::Stop, None),
15877 )])))
15878 }
15879
15880 fn provider_name(&self) -> &str {
15881 "root-turn-probe"
15882 }
15883
15884 fn supports(&self, feature: LLMFeature) -> bool {
15885 matches!(feature, LLMFeature::Streaming)
15886 }
15887 }
15888
15889 #[async_trait]
15890 impl ai_agents_core::Tool for ContextEchoTool {
15891 fn id(&self) -> &str {
15892 "context_echo"
15893 }
15894
15895 fn name(&self) -> &str {
15896 "Context Echo"
15897 }
15898
15899 fn description(&self) -> &str {
15900 "Returns selected execution context fields."
15901 }
15902
15903 fn input_schema(&self) -> Value {
15904 serde_json::json!({"type": "object"})
15905 }
15906
15907 fn policy_bindings(&self) -> ai_agents_core::ToolPolicyBindings {
15908 ai_agents_core::ToolPolicyBindings {
15909 path_fields: vec![ai_agents_core::PathPolicyBinding::read("path")],
15910 result_limit_fields: vec![ai_agents_core::ResultLimitBinding::new(
15911 "max_results",
15912 ai_agents_core::ResultLimitKind::MaxResults,
15913 )],
15914 ..Default::default()
15915 }
15916 }
15917
15918 async fn execute(
15919 &self,
15920 _args: Value,
15921 ctx: ai_agents_core::ToolExecutionContext,
15922 ) -> ToolResult {
15923 ToolResult::ok(
15924 serde_json::json!({
15925 "requested_name": ctx.requested_name,
15926 "canonical_id": ctx.canonical_id,
15927 "display_name": ctx.display_name,
15928 "max_results": ctx.limits.max_results,
15929 "custom_config": ctx.custom_config,
15930 })
15931 .to_string(),
15932 )
15933 }
15934 }
15935
15936 #[async_trait]
15937 impl ai_agents_core::Tool for RetryDeadlineTool {
15938 fn id(&self) -> &str {
15939 "retry_deadline"
15940 }
15941
15942 fn name(&self) -> &str {
15943 "Retry Deadline"
15944 }
15945
15946 fn description(&self) -> &str {
15947 "Records one deadline per retry invocation."
15948 }
15949
15950 fn input_schema(&self) -> Value {
15951 serde_json::json!({"type": "object"})
15952 }
15953
15954 fn safety_metadata(&self) -> ai_agents_core::ToolSafetyMetadata {
15955 ai_agents_core::ToolSafetyMetadata::compute()
15956 }
15957
15958 fn classify_call(&self, _args: &Value) -> ai_agents_core::ToolCallClassification {
15959 let mut classification =
15960 ai_agents_core::ToolCallClassification::from_metadata(&self.safety_metadata());
15961 classification.timeout_ms = Some(1_000);
15962 classification.safely_retryable = true;
15963 classification
15964 }
15965
15966 async fn execute(
15968 &self,
15969 _args: Value,
15970 ctx: ai_agents_core::ToolExecutionContext,
15971 ) -> ToolResult {
15972 let deadline = ctx
15973 .deadline
15974 .expect("each invocation must receive a deadline");
15975 self.remaining_ms.lock().push(
15976 deadline
15977 .signed_duration_since(chrono::Utc::now())
15978 .num_milliseconds(),
15979 );
15980 self.deadlines.lock().push(deadline);
15981 let call = self.calls.fetch_add(1, Ordering::SeqCst);
15982 if call == 0 {
15983 tokio::time::sleep(std::time::Duration::from_millis(20)).await;
15984 ToolResult::error("retry")
15985 } else {
15986 ToolResult::ok("done")
15987 }
15988 }
15989 }
15990
15991 struct ClassifiedTimeoutTool {
15993 id: &'static str,
15994 calls: Arc<std::sync::atomic::AtomicUsize>,
15995 timeout_ms: u64,
15996 sleep_ms: u64,
15997 requires_approval: bool,
15998 remaining_ms: Arc<parking_lot::Mutex<Vec<i64>>>,
15999 }
16000
16001 struct ApprovalModifiedTimeoutTool {
16003 calls: Arc<std::sync::atomic::AtomicUsize>,
16004 }
16005
16006 struct SlowTool;
16008
16009 struct FlakyWriteTool {
16011 calls: Arc<std::sync::atomic::AtomicUsize>,
16012 }
16013
16014 struct LockedWriteTool {
16016 active: Arc<std::sync::atomic::AtomicUsize>,
16017 max_active: Arc<std::sync::atomic::AtomicUsize>,
16018 }
16019
16020 struct MultiResourceWriteTool {
16021 active: Arc<std::sync::atomic::AtomicUsize>,
16022 max_active: Arc<std::sync::atomic::AtomicUsize>,
16023 }
16024
16025 #[derive(Clone)]
16026 struct PathMutationGate {
16027 entered: Arc<AtomicBool>,
16028 entered_notify: Arc<tokio::sync::Notify>,
16029 release: Arc<tokio::sync::Notify>,
16030 }
16031
16032 impl PathMutationGate {
16033 fn new() -> Self {
16034 Self {
16035 entered: Arc::new(AtomicBool::new(false)),
16036 entered_notify: Arc::new(tokio::sync::Notify::new()),
16037 release: Arc::new(tokio::sync::Notify::new()),
16038 }
16039 }
16040
16041 async fn wait_until_entered(&self) {
16042 if !self.entered.load(Ordering::SeqCst) {
16043 self.entered_notify.notified().await;
16044 }
16045 }
16046
16047 fn release(&self) {
16048 self.release.notify_one();
16049 }
16050 }
16051
16052 struct BlockingPathMutationTool {
16053 id: &'static str,
16054 path_fields: Vec<ai_agents_core::PathPolicyBinding>,
16055 gate: PathMutationGate,
16056 }
16057
16058 struct NoBindingWriteTool {
16059 active: Arc<std::sync::atomic::AtomicUsize>,
16060 max_active: Arc<std::sync::atomic::AtomicUsize>,
16061 }
16062
16063 struct RecoveryTestTool {
16064 id: String,
16065 succeeds: bool,
16066 calls: Arc<std::sync::atomic::AtomicUsize>,
16067 max_output_chars: Option<usize>,
16068 }
16069
16070 struct BlockingApprovalHandler {
16071 entered: Arc<tokio::sync::Barrier>,
16072 release: Arc<tokio::sync::Notify>,
16073 result: ApprovalResult,
16074 }
16075
16076 struct CountingApprovalHandler {
16077 calls: Arc<std::sync::atomic::AtomicUsize>,
16078 }
16079
16080 struct DriftingFallbackProvider {
16082 refreshed: AtomicBool,
16083 primary_calls: Arc<std::sync::atomic::AtomicUsize>,
16084 secondary_calls: Arc<std::sync::atomic::AtomicUsize>,
16085 }
16086
16087 struct RefreshFallbackProviderHooks {
16089 agent: parking_lot::Mutex<Option<Weak<RuntimeAgent>>>,
16090 lifecycle: Arc<ToolLifecycleRecordingHooks>,
16091 }
16092
16093 struct RuntimeWebFetchTransport {
16094 calls: Arc<std::sync::atomic::AtomicUsize>,
16095 }
16096
16097 struct RuntimeWebFetchResolver;
16098
16099 struct ReentrantToolHooks {
16100 agent: parking_lot::Mutex<Option<Weak<RuntimeAgent>>>,
16101 invoked: AtomicBool,
16102 nested_success: AtomicBool,
16103 }
16104
16105 #[async_trait]
16106 impl ai_agents_core::Tool for ClassifiedTimeoutTool {
16107 fn id(&self) -> &str {
16109 self.id
16110 }
16111
16112 fn name(&self) -> &str {
16114 "Classified Timeout"
16115 }
16116
16117 fn description(&self) -> &str {
16119 "Records and waits under one call-level timeout."
16120 }
16121
16122 fn input_schema(&self) -> Value {
16124 serde_json::json!({"type": "object"})
16125 }
16126
16127 fn classify_call(&self, _args: &Value) -> ai_agents_core::ToolCallClassification {
16129 let mut classification =
16130 ai_agents_core::ToolCallClassification::from_metadata(&self.safety_metadata());
16131 classification.timeout_ms = Some(self.timeout_ms);
16132 classification.requires_approval = self.requires_approval;
16133 classification
16134 }
16135
16136 async fn execute(
16138 &self,
16139 _args: Value,
16140 ctx: ai_agents_core::ToolExecutionContext,
16141 ) -> ToolResult {
16142 self.calls.fetch_add(1, Ordering::SeqCst);
16143 let deadline = ctx
16144 .deadline
16145 .expect("each invocation must receive a deadline");
16146 self.remaining_ms.lock().push(
16147 deadline
16148 .signed_duration_since(chrono::Utc::now())
16149 .num_milliseconds(),
16150 );
16151 tokio::time::sleep(Duration::from_millis(self.sleep_ms)).await;
16152 ToolResult::ok("done")
16153 }
16154 }
16155
16156 #[async_trait]
16157 impl ai_agents_core::Tool for ApprovalModifiedTimeoutTool {
16158 fn id(&self) -> &str {
16160 "approval_modified_timeout"
16161 }
16162
16163 fn name(&self) -> &str {
16165 "Approval Modified Timeout"
16166 }
16167
16168 fn description(&self) -> &str {
16170 "Becomes invalid only after approval modifies its arguments."
16171 }
16172
16173 fn input_schema(&self) -> Value {
16175 serde_json::json!({"type": "object"})
16176 }
16177
16178 fn policy_bindings(&self) -> ai_agents_core::ToolPolicyBindings {
16180 ai_agents_core::ToolPolicyBindings {
16181 path_fields: vec![ai_agents_core::PathPolicyBinding::write("path")],
16182 ..Default::default()
16183 }
16184 }
16185
16186 fn safety_metadata(&self) -> ai_agents_core::ToolSafetyMetadata {
16188 ai_agents_core::ToolSafetyMetadata {
16189 read_only: false,
16190 concurrency_safe: false,
16191 operation: ai_agents_core::ToolOperationKind::Write,
16192 side_effect_level: ai_agents_core::ToolSideEffectLevel::LocalWrite,
16193 requires_network: false,
16194 destructive: false,
16195 open_world: false,
16196 host_dependent: false,
16197 requires_user_interaction: false,
16198 supports_cancellation: true,
16199 default_requires_approval: true,
16200 should_defer_schema: false,
16201 max_output_chars: Some(1024),
16202 max_result_size_chars: Some(1024),
16203 }
16204 }
16205
16206 fn classify_call(&self, args: &Value) -> ai_agents_core::ToolCallClassification {
16208 let mut classification =
16209 ai_agents_core::ToolCallClassification::from_metadata(&self.safety_metadata());
16210 classification.timeout_ms = Some(if args["invalid_timeout"].as_bool() == Some(true) {
16211 u64::MAX
16212 } else {
16213 1_000
16214 });
16215 classification
16216 }
16217
16218 async fn execute(
16220 &self,
16221 _args: Value,
16222 _ctx: ai_agents_core::ToolExecutionContext,
16223 ) -> ToolResult {
16224 self.calls.fetch_add(1, Ordering::SeqCst);
16225 ToolResult::ok("unexpected")
16226 }
16227 }
16228
16229 #[async_trait]
16230 impl ai_agents_core::Tool for SlowTool {
16231 fn id(&self) -> &str {
16232 "slow"
16233 }
16234
16235 fn name(&self) -> &str {
16236 "Slow"
16237 }
16238
16239 fn description(&self) -> &str {
16240 "Waits until cancelled or timed out."
16241 }
16242
16243 fn input_schema(&self) -> Value {
16244 serde_json::json!({"type": "object"})
16245 }
16246
16247 async fn execute(
16248 &self,
16249 _args: Value,
16250 _ctx: ai_agents_core::ToolExecutionContext,
16251 ) -> ToolResult {
16252 tokio::time::sleep(std::time::Duration::from_secs(5)).await;
16253 ToolResult::ok("done")
16254 }
16255 }
16256
16257 #[async_trait]
16258 impl ai_agents_core::Tool for FlakyWriteTool {
16259 fn id(&self) -> &str {
16260 "flaky_write"
16261 }
16262
16263 fn name(&self) -> &str {
16264 "Flaky Write"
16265 }
16266
16267 fn description(&self) -> &str {
16268 "Fails on the first write attempt."
16269 }
16270
16271 fn input_schema(&self) -> Value {
16272 serde_json::json!({"type": "object", "properties": {"path": {"type": "string"}}})
16273 }
16274
16275 fn policy_bindings(&self) -> ai_agents_core::ToolPolicyBindings {
16276 ai_agents_core::ToolPolicyBindings {
16277 path_fields: vec![ai_agents_core::PathPolicyBinding::write("path")],
16278 ..Default::default()
16279 }
16280 }
16281
16282 fn safety_metadata(&self) -> ai_agents_core::ToolSafetyMetadata {
16283 ai_agents_core::ToolSafetyMetadata {
16284 read_only: false,
16285 concurrency_safe: false,
16286 operation: ai_agents_core::ToolOperationKind::Write,
16287 side_effect_level: ai_agents_core::ToolSideEffectLevel::LocalWrite,
16288 requires_network: false,
16289 destructive: false,
16290 open_world: false,
16291 host_dependent: false,
16292 requires_user_interaction: false,
16293 supports_cancellation: true,
16294 default_requires_approval: false,
16295 should_defer_schema: false,
16296 max_output_chars: Some(1024),
16297 max_result_size_chars: Some(1024),
16298 }
16299 }
16300
16301 fn classify_call(&self, _args: &Value) -> ai_agents_core::ToolCallClassification {
16302 let mut classification =
16303 ai_agents_core::ToolCallClassification::from_metadata(&self.safety_metadata());
16304 classification.safely_retryable = false;
16305 classification
16306 }
16307
16308 async fn execute(
16309 &self,
16310 _args: Value,
16311 _ctx: ai_agents_core::ToolExecutionContext,
16312 ) -> ToolResult {
16313 let call = self.calls.fetch_add(1, Ordering::SeqCst);
16314 if call == 0 {
16315 ToolResult::error("first failure")
16316 } else {
16317 ToolResult::ok("second success")
16318 }
16319 }
16320 }
16321
16322 #[async_trait]
16323 impl ai_agents_core::Tool for LockedWriteTool {
16324 fn id(&self) -> &str {
16325 "locked_write"
16326 }
16327
16328 fn name(&self) -> &str {
16329 "Locked Write"
16330 }
16331
16332 fn description(&self) -> &str {
16333 "Tracks concurrent execution on one resource."
16334 }
16335
16336 fn input_schema(&self) -> Value {
16337 serde_json::json!({"type": "object", "properties": {"path": {"type": "string"}}})
16338 }
16339
16340 fn policy_bindings(&self) -> ai_agents_core::ToolPolicyBindings {
16341 ai_agents_core::ToolPolicyBindings {
16342 path_fields: vec![ai_agents_core::PathPolicyBinding::write("path")],
16343 ..Default::default()
16344 }
16345 }
16346
16347 fn safety_metadata(&self) -> ai_agents_core::ToolSafetyMetadata {
16348 ai_agents_core::ToolSafetyMetadata {
16349 read_only: false,
16350 concurrency_safe: false,
16351 operation: ai_agents_core::ToolOperationKind::Write,
16352 side_effect_level: ai_agents_core::ToolSideEffectLevel::LocalWrite,
16353 requires_network: false,
16354 destructive: false,
16355 open_world: false,
16356 host_dependent: false,
16357 requires_user_interaction: false,
16358 supports_cancellation: true,
16359 default_requires_approval: false,
16360 should_defer_schema: false,
16361 max_output_chars: Some(1024),
16362 max_result_size_chars: Some(1024),
16363 }
16364 }
16365
16366 async fn execute(
16367 &self,
16368 _args: Value,
16369 _ctx: ai_agents_core::ToolExecutionContext,
16370 ) -> ToolResult {
16371 let active = self.active.fetch_add(1, Ordering::SeqCst) + 1;
16372 loop {
16373 let current_max = self.max_active.load(Ordering::SeqCst);
16374 if active <= current_max {
16375 break;
16376 }
16377 if self
16378 .max_active
16379 .compare_exchange(current_max, active, Ordering::SeqCst, Ordering::SeqCst)
16380 .is_ok()
16381 {
16382 break;
16383 }
16384 }
16385 tokio::time::sleep(std::time::Duration::from_millis(100)).await;
16386 self.active.fetch_sub(1, Ordering::SeqCst);
16387 ToolResult::ok("done")
16388 }
16389 }
16390
16391 #[async_trait]
16392 impl ai_agents_core::Tool for MultiResourceWriteTool {
16393 fn id(&self) -> &str {
16394 "multi_resource_write"
16395 }
16396
16397 fn name(&self) -> &str {
16398 "Multi Resource Write"
16399 }
16400
16401 fn description(&self) -> &str {
16402 "Tracks concurrent execution across source and destination resources."
16403 }
16404
16405 fn input_schema(&self) -> Value {
16406 serde_json::json!({"type": "object"})
16407 }
16408
16409 fn policy_bindings(&self) -> ai_agents_core::ToolPolicyBindings {
16410 ai_agents_core::ToolPolicyBindings {
16411 path_fields: vec![
16412 ai_agents_core::PathPolicyBinding::read_write("source_path"),
16413 ai_agents_core::PathPolicyBinding::write("destination_path"),
16414 ],
16415 ..Default::default()
16416 }
16417 }
16418
16419 fn safety_metadata(&self) -> ai_agents_core::ToolSafetyMetadata {
16420 LockedWriteTool {
16421 active: Arc::clone(&self.active),
16422 max_active: Arc::clone(&self.max_active),
16423 }
16424 .safety_metadata()
16425 }
16426
16427 async fn execute(
16428 &self,
16429 _args: Value,
16430 _ctx: ai_agents_core::ToolExecutionContext,
16431 ) -> ToolResult {
16432 let active = self.active.fetch_add(1, Ordering::SeqCst) + 1;
16433 self.max_active.fetch_max(active, Ordering::SeqCst);
16434 tokio::time::sleep(std::time::Duration::from_millis(75)).await;
16435 self.active.fetch_sub(1, Ordering::SeqCst);
16436 ToolResult::ok("done")
16437 }
16438 }
16439
16440 #[async_trait]
16441 impl ai_agents_core::Tool for BlockingPathMutationTool {
16442 fn id(&self) -> &str {
16443 self.id
16444 }
16445
16446 fn name(&self) -> &str {
16447 self.id
16448 }
16449
16450 fn description(&self) -> &str {
16451 "Blocks a path mutation until the test releases it."
16452 }
16453
16454 fn input_schema(&self) -> Value {
16455 serde_json::json!({"type": "object"})
16456 }
16457
16458 fn policy_bindings(&self) -> ai_agents_core::ToolPolicyBindings {
16459 ai_agents_core::ToolPolicyBindings {
16460 path_fields: self.path_fields.clone(),
16461 ..Default::default()
16462 }
16463 }
16464
16465 fn safety_metadata(&self) -> ai_agents_core::ToolSafetyMetadata {
16466 ai_agents_core::ToolSafetyMetadata {
16467 read_only: false,
16468 concurrency_safe: false,
16469 operation: ai_agents_core::ToolOperationKind::Write,
16470 side_effect_level: ai_agents_core::ToolSideEffectLevel::LocalWrite,
16471 requires_network: false,
16472 destructive: false,
16473 open_world: false,
16474 host_dependent: false,
16475 requires_user_interaction: false,
16476 supports_cancellation: true,
16477 default_requires_approval: false,
16478 should_defer_schema: false,
16479 max_output_chars: Some(1024),
16480 max_result_size_chars: Some(1024),
16481 }
16482 }
16483
16484 async fn execute(
16485 &self,
16486 _args: Value,
16487 _ctx: ai_agents_core::ToolExecutionContext,
16488 ) -> ToolResult {
16489 self.gate.entered.store(true, Ordering::SeqCst);
16490 self.gate.entered_notify.notify_one();
16491 self.gate.release.notified().await;
16492 ToolResult::ok("done")
16493 }
16494 }
16495
16496 #[async_trait]
16497 impl ai_agents_core::Tool for NoBindingWriteTool {
16498 fn id(&self) -> &str {
16499 "no_binding_write"
16500 }
16501
16502 fn name(&self) -> &str {
16503 "No Binding Write"
16504 }
16505
16506 fn description(&self) -> &str {
16507 "Tracks concurrent execution without resource bindings."
16508 }
16509
16510 fn input_schema(&self) -> Value {
16511 serde_json::json!({"type": "object"})
16512 }
16513
16514 fn safety_metadata(&self) -> ai_agents_core::ToolSafetyMetadata {
16515 LockedWriteTool {
16516 active: Arc::clone(&self.active),
16517 max_active: Arc::clone(&self.max_active),
16518 }
16519 .safety_metadata()
16520 }
16521
16522 async fn execute(
16523 &self,
16524 _args: Value,
16525 _ctx: ai_agents_core::ToolExecutionContext,
16526 ) -> ToolResult {
16527 let active = self.active.fetch_add(1, Ordering::SeqCst) + 1;
16528 self.max_active.fetch_max(active, Ordering::SeqCst);
16529 tokio::time::sleep(std::time::Duration::from_millis(75)).await;
16530 self.active.fetch_sub(1, Ordering::SeqCst);
16531 ToolResult::ok("done")
16532 }
16533 }
16534
16535 #[async_trait]
16536 impl ai_agents_core::Tool for RecoveryTestTool {
16537 fn id(&self) -> &str {
16538 &self.id
16539 }
16540
16541 fn name(&self) -> &str {
16542 &self.id
16543 }
16544
16545 fn description(&self) -> &str {
16546 "Records recovery execution and returns a configured result."
16547 }
16548
16549 fn input_schema(&self) -> Value {
16550 serde_json::json!({"type": "object"})
16551 }
16552
16553 fn policy_bindings(&self) -> ai_agents_core::ToolPolicyBindings {
16554 ai_agents_core::ToolPolicyBindings {
16555 path_fields: vec![ai_agents_core::PathPolicyBinding::write("path")],
16556 ..Default::default()
16557 }
16558 }
16559
16560 fn safety_metadata(&self) -> ai_agents_core::ToolSafetyMetadata {
16562 ai_agents_core::ToolSafetyMetadata {
16563 read_only: false,
16564 concurrency_safe: false,
16565 operation: ai_agents_core::ToolOperationKind::Write,
16566 side_effect_level: ai_agents_core::ToolSideEffectLevel::LocalWrite,
16567 requires_network: false,
16568 destructive: false,
16569 open_world: false,
16570 host_dependent: false,
16571 requires_user_interaction: false,
16572 supports_cancellation: true,
16573 default_requires_approval: false,
16574 should_defer_schema: false,
16575 max_output_chars: Some(self.max_output_chars.unwrap_or(1024)),
16576 max_result_size_chars: Some(1024),
16577 }
16578 }
16579
16580 async fn execute(
16582 &self,
16583 _args: Value,
16584 _ctx: ai_agents_core::ToolExecutionContext,
16585 ) -> ToolResult {
16586 self.calls.fetch_add(1, Ordering::SeqCst);
16587 let mut result = if self.succeeds {
16588 ToolResult::ok(format!("{} succeeded", self.id))
16589 } else {
16590 ToolResult::error(format!("{} failed", self.id))
16591 };
16592 result.metadata = Some(HashMap::from([(
16593 "recovery_test_tool".to_string(),
16594 Value::String(self.id.clone()),
16595 )]));
16596 result
16597 }
16598 }
16599
16600 #[async_trait]
16601 impl WebFetchTransport for RuntimeWebFetchTransport {
16602 async fn send(
16604 &self,
16605 _request: WebFetchTransportRequest,
16606 ) -> std::result::Result<WebFetchTransportResponse, String> {
16607 Err("validated addresses are required".to_string())
16608 }
16609
16610 async fn send_validated(
16612 &self,
16613 _request: WebFetchTransportRequest,
16614 _addresses: &[std::net::SocketAddr],
16615 ) -> std::result::Result<WebFetchTransportResponse, String> {
16616 self.calls.fetch_add(1, Ordering::SeqCst);
16617 Ok(WebFetchTransportResponse {
16618 status: 200,
16619 content_type: Some("text/plain".to_string()),
16620 location: None,
16621 body: b"approved".to_vec(),
16622 })
16623 }
16624 }
16625
16626 #[async_trait]
16627 impl WebFetchResolver for RuntimeWebFetchResolver {
16628 async fn resolve(
16630 &self,
16631 _host: &str,
16632 _port: u16,
16633 ) -> std::result::Result<Vec<std::net::IpAddr>, String> {
16634 Ok(vec![std::net::IpAddr::V4(std::net::Ipv4Addr::new(
16635 93, 184, 216, 34,
16636 ))])
16637 }
16638 }
16639
16640 #[async_trait]
16641 impl ToolProvider for DriftingFallbackProvider {
16642 fn id(&self) -> &str {
16644 "drifting_fallback"
16645 }
16646
16647 fn name(&self) -> &str {
16649 "Drifting Fallback"
16650 }
16651
16652 fn provider_type(&self) -> ToolProviderType {
16654 ToolProviderType::Custom
16655 }
16656
16657 async fn list_tools(&self) -> Vec<ToolDescriptor> {
16659 let alias = ToolAliases::new().with_name("en", "fallback alias");
16660 let mut primary = ToolDescriptor::new(
16661 "primary",
16662 "Primary",
16663 "Fails before fallback.",
16664 serde_json::json!({"type": "object"}),
16665 );
16666 let mut secondary = ToolDescriptor::new(
16667 "secondary",
16668 "Secondary",
16669 "Must not execute after final canonical drift.",
16670 serde_json::json!({"type": "object"}),
16671 );
16672 if self.refreshed.load(Ordering::SeqCst) {
16673 primary = primary.with_aliases(alias);
16674 } else {
16675 secondary = secondary.with_aliases(alias);
16676 }
16677 vec![primary, secondary]
16678 }
16679
16680 async fn get_tool(&self, tool_id: &str) -> Option<Arc<dyn Tool>> {
16682 let calls = match tool_id {
16683 "primary" => Arc::clone(&self.primary_calls),
16684 "secondary" => Arc::clone(&self.secondary_calls),
16685 _ => return None,
16686 };
16687 Some(Arc::new(RecoveryTestTool {
16688 id: tool_id.to_string(),
16689 succeeds: false,
16690 calls,
16691 max_output_chars: None,
16692 }))
16693 }
16694
16695 fn supports_refresh(&self) -> bool {
16697 true
16698 }
16699
16700 async fn refresh(&self) -> std::result::Result<(), ToolProviderError> {
16702 self.refreshed.store(true, Ordering::SeqCst);
16703 Ok(())
16704 }
16705 }
16706
16707 #[async_trait]
16708 impl AgentHooks for RefreshFallbackProviderHooks {
16709 async fn on_tool_start(&self, tool: &str, args: &Value) {
16711 self.lifecycle.on_tool_start(tool, args).await;
16712 if tool != "secondary" {
16713 return;
16714 }
16715 let agent = self.agent.lock().as_ref().and_then(Weak::upgrade);
16716 if let Some(agent) = agent {
16717 agent
16718 .tools
16719 .refresh_provider("drifting_fallback")
16720 .await
16721 .unwrap();
16722 }
16723 }
16724
16725 async fn on_tool_complete(&self, tool: &str, result: &ToolResult, duration_ms: u64) {
16726 self.lifecycle
16727 .on_tool_complete(tool, result, duration_ms)
16728 .await;
16729 }
16730
16731 async fn on_tool_execution_record(&self, record: &ToolExecutionRecord) {
16732 self.lifecycle.on_tool_execution_record(record).await;
16733 }
16734
16735 async fn on_error(&self, error: &AgentError) {
16736 self.lifecycle.on_error(error).await;
16737 }
16738 }
16739
16740 #[async_trait]
16741 impl ApprovalHandler for BlockingApprovalHandler {
16742 async fn request_approval(
16743 &self,
16744 _request: ai_agents_hitl::ApprovalRequest,
16745 ) -> ApprovalResult {
16746 self.entered.wait().await;
16747 self.release.notified().await;
16748 self.result.clone()
16749 }
16750 }
16751
16752 #[async_trait]
16753 impl ApprovalHandler for CountingApprovalHandler {
16754 async fn request_approval(
16755 &self,
16756 _request: ai_agents_hitl::ApprovalRequest,
16757 ) -> ApprovalResult {
16758 self.calls.fetch_add(1, Ordering::SeqCst);
16759 ApprovalResult::Approved
16760 }
16761 }
16762
16763 #[async_trait]
16764 impl AgentHooks for ReentrantToolHooks {
16765 async fn on_tool_complete(&self, tool: &str, _result: &ToolResult, _duration_ms: u64) {
16766 if tool != "reentrant_write" || self.invoked.swap(true, Ordering::SeqCst) {
16767 return;
16768 }
16769 let agent = self.agent.lock().as_ref().and_then(Weak::upgrade);
16770 if let Some(agent) = agent {
16771 let result = agent
16772 .invoke_tool(ToolExecutionRequest::new(
16773 "nested-hook-call",
16774 "reentrant_write",
16775 serde_json::json!({"path": "./hook.txt"}),
16776 ToolCallSource::Manual,
16777 ))
16778 .await;
16779 self.nested_success
16780 .store(result.is_ok_and(|record| record.success), Ordering::SeqCst);
16781 }
16782 }
16783 }
16784
16785 #[async_trait]
16786 impl AgentHooks for ResponseCountingHooks {
16787 async fn on_response(&self, _response: &AgentResponse) {
16788 self.responses.fetch_add(1, Ordering::SeqCst);
16789 }
16790 }
16791
16792 #[async_trait]
16793 impl AgentHooks for ResponseChatHooks {
16794 async fn on_response(&self, _response: &AgentResponse) {
16796 if self.invoked.swap(true, Ordering::SeqCst) {
16797 return;
16798 }
16799 let target = self.target.lock().as_ref().and_then(Weak::upgrade);
16800 let result = if let Some(target) = target {
16801 target
16802 .chat("nested response hook call")
16803 .await
16804 .map(|response| response.content)
16805 .map_err(|error| error.to_string())
16806 } else {
16807 Err("response hook target is unavailable".to_string())
16808 };
16809 *self.nested_result.lock() = Some(result);
16810 }
16811 }
16812
16813 #[async_trait]
16814 impl AgentHooks for ConcurrentResponseHooks {
16815 async fn on_response(&self, _response: &AgentResponse) {
16817 if self.invoked.swap(true, Ordering::SeqCst) {
16818 return;
16819 }
16820 let Some(registry) = self.registry.upgrade() else {
16821 *self.nested_result.lock() =
16822 Some(Err("concurrent registry is unavailable".to_string()));
16823 return;
16824 };
16825 let agents = [ai_agents_state::ConcurrentAgentRef::Id(
16826 self.child_id.clone(),
16827 )];
16828 let aggregation = ai_agents_state::AggregationConfig {
16829 strategy: ai_agents_state::AggregationStrategy::FirstWins,
16830 synthesizer_llm: None,
16831 synthesizer_prompt: None,
16832 vote: None,
16833 };
16834 let result = crate::orchestration::concurrent(
16835 ®istry,
16836 "nested concurrent response hook call",
16837 &agents,
16838 &aggregation,
16839 None,
16840 Some(1),
16841 None,
16842 ai_agents_state::PartialFailureAction::Abort,
16843 None,
16844 )
16845 .await
16846 .map(|result| result.response.content)
16847 .map_err(|error| error.to_string());
16848 *self.nested_result.lock() = Some(result);
16849 }
16850 }
16851
16852 #[async_trait]
16853 impl AgentHooks for ToolLifecycleRecordingHooks {
16854 async fn on_tool_start(&self, tool: &str, _args: &Value) {
16855 self.events.lock().push(format!("start:{tool}"));
16856 }
16857
16858 async fn on_tool_complete(&self, tool: &str, result: &ToolResult, _duration_ms: u64) {
16859 self.events
16860 .lock()
16861 .push(format!("complete:{tool}:{}", result.success));
16862 }
16863
16864 async fn on_tool_execution_record(&self, record: &ToolExecutionRecord) {
16865 self.events.lock().push(format!(
16866 "record:{}:{}",
16867 record.canonical_id, record.executed
16868 ));
16869 self.records.lock().push(record.clone());
16870 }
16871
16872 async fn on_error(&self, _error: &AgentError) {
16874 self.events.lock().push("error".to_string());
16875 }
16876 }
16877
16878 struct ApprovalRecordingHooks {
16879 events: parking_lot::Mutex<Vec<String>>,
16880 }
16881
16882 impl ApprovalRecordingHooks {
16883 fn new() -> Self {
16884 Self {
16885 events: parking_lot::Mutex::new(Vec::new()),
16886 }
16887 }
16888
16889 fn events(&self) -> Vec<String> {
16890 self.events.lock().clone()
16891 }
16892 }
16893
16894 #[async_trait]
16895 impl AgentHooks for ApprovalRecordingHooks {
16896 async fn on_approval_result(&self, request_id: &str, result: &ApprovalResult) {
16897 self.events.lock().push(format!(
16898 "raw:{}:{}",
16899 request_id,
16900 approval_result_name(result)
16901 ));
16902 }
16903
16904 async fn on_approval_resolved(
16905 &self,
16906 request: &ai_agents_hitl::ApprovalRequest,
16907 raw_result: &ApprovalResult,
16908 outcome: &ApprovalResolvedOutcome,
16909 ) {
16910 self.events.lock().push(format!(
16911 "resolved:{}:{}:{}",
16912 request.id,
16913 approval_result_name(raw_result),
16914 approval_outcome_name(outcome)
16915 ));
16916 }
16917 }
16918
16919 fn approval_result_name(result: &ApprovalResult) -> &'static str {
16920 match result {
16921 ApprovalResult::Approved => "approved",
16922 ApprovalResult::Rejected { .. } => "rejected",
16923 ApprovalResult::Modified { .. } => "modified",
16924 ApprovalResult::Timeout => "timeout",
16925 }
16926 }
16927
16928 fn approval_outcome_name(outcome: &ApprovalResolvedOutcome) -> &'static str {
16929 match outcome {
16930 ApprovalResolvedOutcome::Approved => "approved",
16931 ApprovalResolvedOutcome::Rejected { .. } => "rejected",
16932 ApprovalResolvedOutcome::Modified { .. } => "modified",
16933 ApprovalResolvedOutcome::Error { .. } => "error",
16934 }
16935 }
16936
16937 fn assert_correlated_approval_events(
16938 events: &[String],
16939 raw_status: &str,
16940 outcome_status: &str,
16941 ) {
16942 assert_eq!(events.len(), 2);
16943 let raw: Vec<_> = events[0].split(':').collect();
16944 let resolved: Vec<_> = events[1].split(':').collect();
16945 assert_eq!(raw[0], "raw");
16946 assert_eq!(resolved[0], "resolved");
16947 assert_eq!(raw[1], resolved[1]);
16948 assert_eq!(raw[2], raw_status);
16949 assert_eq!(resolved[2], raw_status);
16950 assert_eq!(resolved[3], outcome_status);
16951 }
16952
16953 fn approval_security_config(policy_enabled: bool) -> ToolSecurityConfig {
16954 let mut security = ToolSecurityConfig {
16955 enabled: true,
16956 fail_closed: true,
16957 ..Default::default()
16958 };
16959 let policy = ai_agents_tools::ToolPolicyConfig {
16960 enabled: policy_enabled,
16961 write_paths: vec![".".to_string()],
16962 require_confirmation: true,
16963 ..Default::default()
16964 };
16965 security.tools.insert("locked_write".to_string(), policy);
16966 security
16967 }
16968
16969 struct MutationTestWorkspace {
16970 root: std::path::PathBuf,
16971 }
16972
16973 impl MutationTestWorkspace {
16974 fn new() -> Self {
16975 let root = std::env::temp_dir().join(format!(
16976 "ai-agents-runtime-mutation-{}",
16977 uuid::Uuid::new_v4()
16978 ));
16979 std::fs::create_dir_all(&root).unwrap();
16980 Self { root }
16981 }
16982 }
16983
16984 impl Drop for MutationTestWorkspace {
16985 fn drop(&mut self) {
16986 let _ = std::fs::remove_dir_all(&self.root);
16987 }
16988 }
16989
16990 async fn wait_for_resource_lock_strong_count(locks: &ToolResourceLocks, minimum: usize) {
16991 tokio::time::timeout(std::time::Duration::from_secs(2), async {
16992 loop {
16993 let strong_count = locks
16994 .read()
16995 .get("path-mutation:global")
16996 .map_or(0, |lock| lock.strong_count());
16997 if strong_count >= minimum {
16998 break;
16999 }
17000 tokio::task::yield_now().await;
17001 }
17002 })
17003 .await
17004 .expect("path mutation call did not reach the shared lock");
17005 }
17006
17007 async fn assert_path_mutation_pair_serialized(
17008 first_id: &'static str,
17009 first_fields: Vec<ai_agents_core::PathPolicyBinding>,
17010 first_args: Value,
17011 second_id: &'static str,
17012 second_fields: Vec<ai_agents_core::PathPolicyBinding>,
17013 second_args: Value,
17014 ) {
17015 let locks = new_tool_resource_locks();
17016 let first_gate = PathMutationGate::new();
17017 let second_gate = PathMutationGate::new();
17018 second_gate.release();
17019 let agent = Arc::new(
17020 AgentBuilder::new()
17021 .system_prompt("Test global path mutation locking.")
17022 .llm(Arc::new(mock_with_response("done")))
17023 .tool(Arc::new(BlockingPathMutationTool {
17024 id: first_id,
17025 path_fields: first_fields,
17026 gate: first_gate.clone(),
17027 }))
17028 .tool(Arc::new(BlockingPathMutationTool {
17029 id: second_id,
17030 path_fields: second_fields,
17031 gate: second_gate.clone(),
17032 }))
17033 .build()
17034 .unwrap()
17035 .with_shared_resource_locks(Arc::clone(&locks)),
17036 );
17037
17038 let first = {
17039 let agent = Arc::clone(&agent);
17040 tokio::spawn(async move {
17041 agent
17042 .invoke_tool(ToolExecutionRequest::new(
17043 format!("{}-first", first_id),
17044 first_id,
17045 first_args,
17046 ToolCallSource::Manual,
17047 ))
17048 .await
17049 .unwrap()
17050 })
17051 };
17052 first_gate.wait_until_entered().await;
17053
17054 let second = {
17055 let agent = Arc::clone(&agent);
17056 tokio::spawn(async move {
17057 agent
17058 .invoke_tool(ToolExecutionRequest::new(
17059 format!("{}-second", second_id),
17060 second_id,
17061 second_args,
17062 ToolCallSource::Manual,
17063 ))
17064 .await
17065 .unwrap()
17066 })
17067 };
17068 wait_for_resource_lock_strong_count(&locks, 2).await;
17069 assert!(!second_gate.entered.load(Ordering::SeqCst));
17070 assert!(!second.is_finished());
17071
17072 first_gate.release();
17073 let (first, second) = tokio::time::timeout(std::time::Duration::from_secs(2), async {
17074 tokio::join!(first, second)
17075 })
17076 .await
17077 .expect("serialized path mutation calls did not finish");
17078 assert!(first.unwrap().success);
17079 assert!(second.unwrap().success);
17080 assert!(second_gate.entered.load(Ordering::SeqCst));
17081 assert!(locks.read().is_empty());
17082 }
17083
17084 #[derive(Clone, Copy)]
17085 enum MutationDenial {
17086 Policy,
17087 Approval,
17088 }
17089
17090 fn mutation_denial_security_config(
17091 tool_id: &str,
17092 workspace: &std::path::Path,
17093 denial: MutationDenial,
17094 ) -> ToolSecurityConfig {
17095 let workspace = workspace.to_string_lossy().into_owned();
17096 let mut policy = ai_agents_tools::ToolPolicyConfig {
17097 read_paths: vec![workspace.clone()],
17098 write_paths: vec![workspace.clone()],
17099 ..Default::default()
17100 };
17101 match denial {
17102 MutationDenial::Policy => policy.blocked_paths = vec![workspace],
17103 MutationDenial::Approval => policy.require_confirmation = true,
17104 }
17105
17106 let mut security = ToolSecurityConfig {
17107 enabled: true,
17108 fail_closed: true,
17109 ..Default::default()
17110 };
17111 security.tools.insert(tool_id.to_string(), policy);
17112 security
17113 }
17114
17115 async fn assert_path_mutation_denied(tool: Arc<dyn Tool>, denial: MutationDenial) {
17116 let workspace = MutationTestWorkspace::new();
17117 let tool_id = tool.id().to_string();
17118 let preserved = workspace.root.join(format!("{}-preserved.txt", tool_id));
17119 let destination = workspace.root.join(format!("{}-destination.txt", tool_id));
17120 std::fs::write(&preserved, "preserved").unwrap();
17121 let arguments = match tool_id.as_str() {
17122 "copy_path" | "move_path" => serde_json::json!({
17123 "source_path": preserved.to_string_lossy(),
17124 "destination_path": destination.to_string_lossy(),
17125 "dry_run": false
17126 }),
17127 "delete_path" => serde_json::json!({
17128 "path": preserved.to_string_lossy(),
17129 "recursive": false,
17130 "dry_run": false
17131 }),
17132 _ => panic!("unsupported mutation tool: {}", tool_id),
17133 };
17134 let security = mutation_denial_security_config(&tool_id, &workspace.root, denial);
17135 let builder = AgentBuilder::new()
17136 .system_prompt("Test mutation denial.")
17137 .llm(Arc::new(mock_with_response("done")))
17138 .tool(tool)
17139 .tool_security(ToolSecurityEngine::new(security));
17140 let builder = match denial {
17141 MutationDenial::Policy => builder,
17142 MutationDenial::Approval => builder
17143 .hitl_engine(HITLEngine::new(ai_agents_hitl::HITLConfig::default()))
17144 .approval_handler(Arc::new(RejectAllHandler::new())),
17145 };
17146 let agent = builder.build().unwrap();
17147
17148 let record = agent
17149 .invoke_tool(ToolExecutionRequest::new(
17150 format!("{}-denied", tool_id),
17151 tool_id.clone(),
17152 arguments,
17153 ToolCallSource::Manual,
17154 ))
17155 .await
17156 .unwrap();
17157
17158 assert!(!record.executed, "{} must not be invoked", tool_id);
17159 assert!(!record.success);
17160 match denial {
17161 MutationDenial::Policy => {
17162 assert_eq!(record.policy.outcome, PermissionOutcome::Deny);
17163 assert!(record.approval.as_ref().is_some_and(|approval| matches!(
17164 &approval.status,
17165 ToolApprovalStatus::NotRequired
17166 )));
17167 }
17168 MutationDenial::Approval => {
17169 assert_eq!(record.policy.outcome, PermissionOutcome::RequiresApproval);
17170 assert!(record.approval.as_ref().is_some_and(|approval| matches!(
17171 &approval.status,
17172 ToolApprovalStatus::Rejected
17173 )));
17174 }
17175 }
17176 assert_eq!(std::fs::read_to_string(&preserved).unwrap(), "preserved");
17177 assert!(!destination.exists());
17178 }
17179
17180 fn recovery_manager_with_fallbacks(
17181 fallbacks: impl IntoIterator<Item = (String, String)>,
17182 ) -> RecoveryManager {
17183 use ai_agents_recovery::{ErrorRecoveryConfig, ToolRecoveryConfig, ToolRetryConfig};
17184
17185 let per_tool = fallbacks
17186 .into_iter()
17187 .map(|(tool, fallback_tool)| {
17188 (
17189 tool,
17190 ToolRetryConfig {
17191 max_retries: 0,
17192 timeout_ms: Some(1_000),
17193 on_failure: ToolFailureAction::Fallback { fallback_tool },
17194 },
17195 )
17196 })
17197 .collect();
17198 RecoveryManager::new(ErrorRecoveryConfig {
17199 tools: ToolRecoveryConfig {
17200 per_tool,
17201 ..Default::default()
17202 },
17203 ..Default::default()
17204 })
17205 }
17206
17207 fn approval_check() -> HITLCheckResult {
17208 HITLCheckResult::required(
17209 ApprovalTrigger::tool("test", serde_json::json!({})),
17210 HashMap::new(),
17211 "Approve?",
17212 None,
17213 )
17214 }
17215
17216 fn agent_with_approval_result(
17217 raw_result: ApprovalResult,
17218 timeout_action: TimeoutAction,
17219 hooks: Arc<ApprovalRecordingHooks>,
17220 ) -> RuntimeAgent {
17221 use ai_agents_hitl::{CallbackHandler, HITLConfig};
17222
17223 let config = HITLConfig {
17224 on_timeout: timeout_action,
17225 ..Default::default()
17226 };
17227 let handler = CallbackHandler::new(move |_| raw_result.clone());
17228 AgentBuilder::new()
17229 .system_prompt("Test HITL hooks.")
17230 .llm(Arc::new(mock_with_response("done")))
17231 .build()
17232 .unwrap()
17233 .with_hooks(hooks)
17234 .with_hitl(HITLEngine::new(config), Arc::new(handler))
17235 }
17236
17237 #[tokio::test]
17238 async fn approval_hooks_expose_direct_effective_decisions_after_raw_results() {
17239 let cases = vec![
17240 (ApprovalResult::Approved, "approved"),
17241 (
17242 ApprovalResult::Rejected {
17243 reason: Some("denied".to_string()),
17244 },
17245 "rejected",
17246 ),
17247 (
17248 ApprovalResult::Modified {
17249 changes: HashMap::from([("value".to_string(), serde_json::json!(2))]),
17250 },
17251 "modified",
17252 ),
17253 ];
17254
17255 for (raw_result, expected) in cases {
17256 let hooks = Arc::new(ApprovalRecordingHooks::new());
17257 let agent =
17258 agent_with_approval_result(raw_result, TimeoutAction::Reject, hooks.clone());
17259
17260 let result = agent.request_hitl_approval(approval_check()).await.unwrap();
17261
17262 assert_eq!(approval_result_name(&result), expected);
17263 assert_correlated_approval_events(&hooks.events(), expected, expected);
17264 }
17265 }
17266
17267 #[tokio::test]
17268 async fn approval_hooks_expose_timeout_policy_decisions() {
17269 for (timeout_action, expected) in [
17270 (TimeoutAction::Approve, "approved"),
17271 (TimeoutAction::Reject, "rejected"),
17272 ] {
17273 let hooks = Arc::new(ApprovalRecordingHooks::new());
17274 let agent =
17275 agent_with_approval_result(ApprovalResult::Timeout, timeout_action, hooks.clone());
17276
17277 let result = agent.request_hitl_approval(approval_check()).await.unwrap();
17278
17279 assert_eq!(approval_result_name(&result), expected);
17280 assert_correlated_approval_events(&hooks.events(), "timeout", expected);
17281 }
17282 }
17283
17284 #[tokio::test]
17285 async fn timeout_error_fires_correlated_resolved_error_before_returning() {
17286 let hooks = Arc::new(ApprovalRecordingHooks::new());
17287 let agent = agent_with_approval_result(
17288 ApprovalResult::Timeout,
17289 TimeoutAction::Error,
17290 hooks.clone(),
17291 );
17292
17293 let error = agent
17294 .request_hitl_approval(approval_check())
17295 .await
17296 .unwrap_err();
17297
17298 assert!(error.to_string().contains("HITL approval timeout"));
17299 assert_correlated_approval_events(&hooks.events(), "timeout", "error");
17300 }
17301
17302 #[tokio::test]
17304 async fn test_integration_yaml_to_chat_basic() {
17305 let mock = mock_with_response("Hello! How can I help you?");
17306 let agent = AgentBuilder::new()
17307 .system_prompt("You are a test assistant.")
17308 .llm(Arc::new(mock))
17309 .build()
17310 .unwrap();
17311
17312 let response = agent.chat("Hi").await.unwrap();
17313 assert!(!response.content.is_empty());
17314 assert_eq!(response.content, "Hello! How can I help you?");
17315 }
17316
17317 #[tokio::test]
17318 async fn stream_events_emit_one_authoritative_final_without_legacy_done() {
17319 let agent = AgentBuilder::new()
17320 .system_prompt("You are a test assistant.")
17321 .llm(Arc::new(mock_with_response(
17322 "Hello from the final response.",
17323 )))
17324 .build()
17325 .unwrap();
17326
17327 let mut stream = agent.chat_stream_events("Hi").await.unwrap();
17328 let mut final_responses = Vec::new();
17329 let mut legacy_done = 0;
17330 while let Some(event) = stream.next().await {
17331 match event {
17332 AgentStreamEvent::Chunk(StreamChunk::Done {}) => legacy_done += 1,
17333 AgentStreamEvent::Chunk(StreamChunk::Error { message }) => {
17334 panic!("unexpected stream error: {message}")
17335 }
17336 AgentStreamEvent::Final(response) => final_responses.push(response),
17337 AgentStreamEvent::Chunk(_) => {}
17338 }
17339 }
17340
17341 assert_eq!(legacy_done, 0);
17342 assert_eq!(final_responses.len(), 1);
17343 let response = final_responses.pop().unwrap();
17344 assert_eq!(response.content, "Hello from the final response.");
17345 assert!(
17346 response
17347 .metadata
17348 .as_ref()
17349 .is_some_and(|metadata| { metadata.contains_key("reasoning") })
17350 );
17351 }
17352
17353 #[tokio::test]
17354 async fn stream_final_content_includes_output_processing_after_provisional_chunks() {
17355 let yaml = r#"
17356name: ProcessedStreamAgent
17357system_prompt: "Answer directly."
17358process:
17359 output:
17360 - type: format
17361 config:
17362 template: "{{ response }} [finalized]"
17363streaming:
17364 enabled: true
17365"#;
17366 let agent = AgentBuilder::from_yaml(yaml)
17367 .unwrap()
17368 .llm(Arc::new(mock_with_response("provisional answer")))
17369 .auto_configure_features()
17370 .unwrap()
17371 .build()
17372 .unwrap();
17373
17374 let mut stream = agent.chat_stream_events("Hi").await.unwrap();
17375 let mut provisional = String::new();
17376 let mut final_content = None;
17377 while let Some(event) = stream.next().await {
17378 match event {
17379 AgentStreamEvent::Chunk(StreamChunk::Content { text }) => {
17380 provisional.push_str(&text)
17381 }
17382 AgentStreamEvent::Chunk(StreamChunk::Error { message }) => {
17383 panic!("unexpected stream error: {message}")
17384 }
17385 AgentStreamEvent::Final(response) => final_content = Some(response.content),
17386 AgentStreamEvent::Chunk(_) => {}
17387 }
17388 }
17389
17390 assert_eq!(provisional, "provisional answer");
17391 assert_eq!(
17392 final_content.as_deref(),
17393 Some("provisional answer [finalized]")
17394 );
17395 }
17396
17397 #[tokio::test]
17398 async fn stream_events_preserve_tool_progress_and_final_tool_calls() {
17399 let agent = AgentBuilder::new()
17400 .system_prompt("Use the echo tool once, then answer.")
17401 .llm(Arc::new(mock_with_responses(vec![
17402 r#"{"tool":"echo","arguments":{"message":"hello"}}"#,
17403 "Echo completed.",
17404 ])))
17405 .tool(Arc::new(ai_agents_tools::EchoTool::new()))
17406 .build()
17407 .unwrap();
17408
17409 let mut stream = agent.chat_stream_events("echo hello").await.unwrap();
17410 let mut starts = 0;
17411 let mut results = 0;
17412 let mut ends = 0;
17413 let mut final_response = None;
17414 while let Some(event) = stream.next().await {
17415 match event {
17416 AgentStreamEvent::Chunk(StreamChunk::ToolCallStart { name, .. }) => {
17417 assert_eq!(name, "echo");
17418 starts += 1;
17419 }
17420 AgentStreamEvent::Chunk(StreamChunk::ToolResult { name, success, .. }) => {
17421 assert_eq!(name, "echo");
17422 assert!(success);
17423 results += 1;
17424 }
17425 AgentStreamEvent::Chunk(StreamChunk::ToolCallEnd { .. }) => ends += 1,
17426 AgentStreamEvent::Chunk(StreamChunk::Error { message }) => {
17427 panic!("unexpected stream error: {message}")
17428 }
17429 AgentStreamEvent::Final(response) => final_response = Some(response),
17430 AgentStreamEvent::Chunk(_) => {}
17431 }
17432 }
17433
17434 assert_eq!((starts, results, ends), (1, 1, 1));
17435 let response = final_response.expect("tool stream must finalize");
17436 assert_eq!(response.content, "Echo completed.");
17437 assert_eq!(
17438 response.tool_calls.as_ref().map(|calls| calls
17439 .iter()
17440 .map(|call| call.name.as_str())
17441 .collect::<Vec<_>>()),
17442 Some(vec!["echo"])
17443 );
17444 }
17445
17446 #[tokio::test]
17447 async fn legacy_stream_still_emits_one_done_chunk() {
17448 let agent = AgentBuilder::new()
17449 .system_prompt("You are a test assistant.")
17450 .llm(Arc::new(mock_with_response(
17451 "Hello from the legacy stream.",
17452 )))
17453 .build()
17454 .unwrap();
17455
17456 let mut stream = agent.chat_stream("Hi").await.unwrap();
17457 let mut done = 0;
17458 while let Some(chunk) = stream.next().await {
17459 match chunk {
17460 StreamChunk::Done {} => done += 1,
17461 StreamChunk::Error { message } => panic!("unexpected stream error: {message}"),
17462 _ => {}
17463 }
17464 }
17465
17466 assert_eq!(done, 1);
17467 }
17468
17469 #[tokio::test]
17471 async fn test_integration_multi_turn_conversation() {
17472 let mock = mock_with_responses(vec![
17473 "Hello! I'm your assistant.",
17474 "The weather is sunny today.",
17475 "Goodbye!",
17476 ]);
17477 let agent = AgentBuilder::new()
17478 .system_prompt("You are helpful.")
17479 .llm(Arc::new(mock))
17480 .build()
17481 .unwrap();
17482
17483 let r1 = agent.chat("Hi").await.unwrap();
17484 assert_eq!(r1.content, "Hello! I'm your assistant.");
17485
17486 let r2 = agent.chat("What's the weather?").await.unwrap();
17487 assert_eq!(r2.content, "The weather is sunny today.");
17488
17489 let r3 = agent.chat("Bye").await.unwrap();
17490 assert_eq!(r3.content, "Goodbye!");
17491
17492 let messages = agent.memory.get_messages(None).await.unwrap();
17494 assert_eq!(messages.len(), 6);
17496 }
17497
17498 #[test]
17499 fn later_approval_preserves_modified_evidence() {
17500 let arguments = serde_json::json!({"dry_run": true});
17501 let mut record = Some(ToolApprovalRecord {
17502 status: ToolApprovalStatus::Modified,
17503 reason: None,
17504 modified_arguments: Some(arguments.clone()),
17505 });
17506
17507 merge_approved_record(&mut record);
17508
17509 let record = record.unwrap();
17510 assert!(matches!(record.status, ToolApprovalStatus::Modified));
17511 assert_eq!(record.modified_arguments, Some(arguments));
17512 }
17513
17514 #[test]
17515 fn approval_binding_rejects_replaced_tool_implementation() {
17516 let reviewed_tool: Arc<dyn ai_agents_core::Tool> = Arc::new(ContextEchoTool);
17517 let same_tool = Arc::clone(&reviewed_tool);
17518 let replacement_tool: Arc<dyn ai_agents_core::Tool> = Arc::new(ContextEchoTool);
17519 let arguments = serde_json::json!({"path": "."});
17520 let versions = ToolDecisionVersions {
17521 policy: 2,
17522 registry: 3,
17523 runtime_control: 4,
17524 state: Some(5),
17525 };
17526 let binding = ToolApprovalBinding {
17527 canonical_id: "context_echo".to_string(),
17528 arguments: arguments.clone(),
17529 confirmation_required: true,
17530 policy_version: versions.policy,
17531 runtime_control_version: versions.runtime_control,
17532 state_generation: versions.state,
17533 reviewed_tool,
17534 };
17535
17536 assert!(!binding.is_stale("context_echo", &arguments, true, versions, &same_tool,));
17537 assert!(binding.is_stale(
17538 "context_echo",
17539 &arguments,
17540 true,
17541 versions,
17542 &replacement_tool,
17543 ));
17544 }
17545
17546 #[tokio::test]
17547 async fn approved_mutation_to_dry_run_remains_executable() {
17548 use ai_agents_hitl::CallbackHandler;
17549
17550 let handler = CallbackHandler::new(|_| ApprovalResult::Modified {
17551 changes: HashMap::from([("dry_run".to_string(), serde_json::json!(true))]),
17552 });
17553 let agent = AgentBuilder::new()
17554 .system_prompt("Test safer approval modifications.")
17555 .llm(Arc::new(mock_with_response("done")))
17556 .tool(Arc::new(ai_agents_tools::FileWriteTool::new()))
17557 .hitl_engine(HITLEngine::new(ai_agents_hitl::HITLConfig::default()))
17558 .approval_handler(Arc::new(handler))
17559 .build()
17560 .unwrap();
17561
17562 let record = agent
17563 .invoke_tool(ToolExecutionRequest::new(
17564 "approved-dry-run",
17565 "file_write",
17566 serde_json::json!({
17567 "path": "./approval-dry-run.txt",
17568 "content": "not written"
17569 }),
17570 ToolCallSource::Manual,
17571 ))
17572 .await
17573 .unwrap();
17574
17575 assert!(record.executed);
17576 assert!(record.success);
17577 assert_eq!(record.executed_arguments["dry_run"], true);
17578 assert!(matches!(
17579 record.approval.as_ref().map(|approval| &approval.status),
17580 Some(ToolApprovalStatus::Modified)
17581 ));
17582 let output: Value = serde_json::from_str(&record.output).unwrap();
17583 assert_eq!(output["mutation_performed"], false);
17584 }
17585
17586 #[tokio::test]
17588 async fn shared_executor_approval_reaches_web_fetch_transport() {
17589 use ai_agents_hitl::{CallbackHandler, HITLConfig};
17590 use ai_agents_tools::{DomainPolicyConfig, ToolPolicyConfig};
17591
17592 let calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
17593 let tool = WebFetchTool::with_transport_and_resolver(
17594 Arc::new(RuntimeWebFetchTransport {
17595 calls: Arc::clone(&calls),
17596 }),
17597 Arc::new(RuntimeWebFetchResolver),
17598 );
17599 let mut security = ToolSecurityConfig {
17600 enabled: true,
17601 fail_closed: true,
17602 ..Default::default()
17603 };
17604 security.tools.insert(
17605 "web_fetch".to_string(),
17606 ToolPolicyConfig {
17607 domains: DomainPolicyConfig {
17608 requires_approval: vec!["approval.test".to_string()],
17609 ..Default::default()
17610 },
17611 allowed_schemes: vec!["https".to_string()],
17612 allowed_ports: vec![443],
17613 ..Default::default()
17614 },
17615 );
17616 let handler = CallbackHandler::new(|_| ApprovalResult::Approved);
17617 let agent = AgentBuilder::new()
17618 .system_prompt("Test approved web fetch execution.")
17619 .llm(Arc::new(mock_with_response("done")))
17620 .tool(Arc::new(tool))
17621 .tool_security(ToolSecurityEngine::new(security))
17622 .build()
17623 .unwrap()
17624 .with_hitl(HITLEngine::new(HITLConfig::default()), Arc::new(handler));
17625
17626 let record = agent
17627 .invoke_tool(ToolExecutionRequest::new(
17628 "approved-web-fetch",
17629 "web_fetch",
17630 serde_json::json!({
17631 "url": "https://approval.test/page",
17632 "cache_ttl_seconds": 0
17633 }),
17634 ToolCallSource::Manual,
17635 ))
17636 .await
17637 .unwrap();
17638
17639 assert!(record.success);
17640 assert!(
17641 record
17642 .approval
17643 .as_ref()
17644 .is_some_and(|approval| matches!(approval.status, ToolApprovalStatus::Approved))
17645 );
17646 assert_eq!(calls.load(Ordering::SeqCst), 1);
17647 }
17648
17649 #[tokio::test]
17650 async fn context_preserves_requested_and_canonical_identity() {
17651 let mock = mock_with_response("hello");
17652 let mut tools = ai_agents_tools::ToolRegistry::new();
17653 tools.register(Arc::new(ContextEchoTool)).unwrap();
17654
17655 let mut security = ToolSecurityConfig {
17656 enabled: true,
17657 fail_closed: true,
17658 ..Default::default()
17659 };
17660 let mut policy = ai_agents_tools::ToolPolicyConfig {
17661 read_paths: vec![".".to_string()],
17662 max_results: Some(7),
17663 ..Default::default()
17664 };
17665 policy
17666 .config
17667 .insert("backend".to_string(), serde_json::json!("memory"));
17668 security.tools.insert("context_echo".to_string(), policy);
17669
17670 let agent = AgentBuilder::new()
17671 .system_prompt("You are helpful.")
17672 .llm(Arc::new(mock))
17673 .tools(tools)
17674 .tool_security(ToolSecurityEngine::new(security))
17675 .build()
17676 .unwrap();
17677
17678 let record = agent
17679 .invoke_tool(ToolExecutionRequest::new(
17680 "ctx-call",
17681 "Context Echo",
17682 serde_json::json!({"path": ".", "max_results": 99}),
17683 ToolCallSource::Manual,
17684 ))
17685 .await
17686 .unwrap();
17687
17688 assert!(record.success);
17689 assert!(matches!(&record.source, ToolCallSource::Manual));
17690 assert_eq!(record.requested_name, "Context Echo");
17691 assert_eq!(record.canonical_id, "context_echo");
17692 assert_eq!(record.policy.outcome, PermissionOutcome::Allow);
17693 assert_eq!(record.executed_arguments["max_results"], 7);
17694 let output: Value = serde_json::from_str(&record.output).unwrap();
17695 assert_eq!(output["requested_name"], "Context Echo");
17696 assert_eq!(output["canonical_id"], "context_echo");
17697 assert_eq!(output["max_results"], 7);
17698 assert_eq!(output["custom_config"]["backend"], "memory");
17699 assert!(record.metadata.contains_key("effective_limits"));
17700 assert!(record.metadata.contains_key("policy_snapshot"));
17701 }
17702
17703 #[tokio::test]
17704 async fn test_runtime_control_cancels_active_tool_call() {
17705 let mock = mock_with_response("hello");
17706 let agent = Arc::new(
17707 AgentBuilder::new()
17708 .system_prompt("You are helpful.")
17709 .llm(Arc::new(mock))
17710 .tool(Arc::new(SlowTool))
17711 .build()
17712 .unwrap(),
17713 );
17714 let control = agent.runtime_control();
17715 let running_agent = Arc::clone(&agent);
17716 let handle = tokio::spawn(async move {
17717 running_agent
17718 .invoke_tool(ToolExecutionRequest::new(
17719 "slow-call",
17720 "slow",
17721 serde_json::json!({}),
17722 ToolCallSource::Manual,
17723 ))
17724 .await
17725 .unwrap()
17726 });
17727
17728 tokio::time::sleep(std::time::Duration::from_millis(100)).await;
17729 control.cancel_all();
17730 let record = handle.await.unwrap();
17731
17732 assert!(record.executed);
17733 assert!(record.cancelled);
17734 assert!(!record.success);
17735 assert!(record.cancellation_reason.is_some());
17736 }
17737
17738 #[tokio::test]
17740 async fn cancelled_tool_does_not_enter_fallback() {
17741 let fallback_calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
17742 let agent = Arc::new(
17743 AgentBuilder::new()
17744 .system_prompt("Test cancellation before fallback.")
17745 .llm(Arc::new(mock_with_response("done")))
17746 .tool(Arc::new(SlowTool))
17747 .tool(Arc::new(RecoveryTestTool {
17748 id: "fallback".to_string(),
17749 succeeds: true,
17750 calls: Arc::clone(&fallback_calls),
17751 max_output_chars: None,
17752 }))
17753 .recovery_manager(recovery_manager_with_fallbacks([(
17754 "slow".to_string(),
17755 "fallback".to_string(),
17756 )]))
17757 .build()
17758 .unwrap(),
17759 );
17760 let control = agent.runtime_control();
17761 let running_agent = Arc::clone(&agent);
17762 let handle = tokio::spawn(async move {
17763 running_agent
17764 .invoke_tool(ToolExecutionRequest::new(
17765 "cancelled-fallback-call",
17766 "slow",
17767 serde_json::json!({}),
17768 ToolCallSource::Manual,
17769 ))
17770 .await
17771 .unwrap()
17772 });
17773
17774 tokio::time::sleep(Duration::from_millis(100)).await;
17775 control.cancel_all();
17776 let record = handle.await.unwrap();
17777
17778 assert!(record.executed);
17779 assert!(record.cancelled);
17780 assert!(!record.success);
17781 assert_eq!(record.canonical_id, "slow");
17782 assert_eq!(fallback_calls.load(Ordering::SeqCst), 0);
17783 assert_eq!(agent.tool_call_history().len(), 1);
17784 }
17785
17786 #[tokio::test]
17787 async fn non_idempotent_tool_calls_are_not_retried() {
17788 use ai_agents_recovery::{ErrorRecoveryConfig, ToolRecoveryConfig, ToolRetryConfig};
17789
17790 let mock = mock_with_response("hello");
17791 let calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
17792 let agent = AgentBuilder::new()
17793 .system_prompt("You are helpful.")
17794 .llm(Arc::new(mock))
17795 .tool(Arc::new(FlakyWriteTool {
17796 calls: Arc::clone(&calls),
17797 }))
17798 .recovery_manager(RecoveryManager::new(ErrorRecoveryConfig {
17799 tools: ToolRecoveryConfig {
17800 default: ToolRetryConfig {
17801 max_retries: 2,
17802 ..Default::default()
17803 },
17804 ..Default::default()
17805 },
17806 ..Default::default()
17807 }))
17808 .build()
17809 .unwrap();
17810
17811 let record = agent
17812 .invoke_tool(ToolExecutionRequest::new(
17813 "flaky-call",
17814 "flaky_write",
17815 serde_json::json!({"path": "./tmp.txt"}),
17816 ToolCallSource::Manual,
17817 ))
17818 .await
17819 .unwrap();
17820
17821 assert!(!record.success);
17822 assert_eq!(calls.load(Ordering::SeqCst), 1);
17823 }
17824
17825 #[tokio::test]
17826 async fn safely_retryable_tool_receives_a_fresh_deadline_per_attempt() {
17827 use ai_agents_recovery::{ErrorRecoveryConfig, ToolRecoveryConfig, ToolRetryConfig};
17828
17829 let calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
17830 let deadlines = Arc::new(parking_lot::Mutex::new(Vec::new()));
17831 let remaining_ms = Arc::new(parking_lot::Mutex::new(Vec::new()));
17832 let agent = AgentBuilder::new()
17833 .system_prompt("Test retry deadlines.")
17834 .llm(Arc::new(mock_with_response("done")))
17835 .tool(Arc::new(RetryDeadlineTool {
17836 calls: Arc::clone(&calls),
17837 deadlines: Arc::clone(&deadlines),
17838 remaining_ms: Arc::clone(&remaining_ms),
17839 }))
17840 .recovery_manager(RecoveryManager::new(ErrorRecoveryConfig {
17841 tools: ToolRecoveryConfig {
17842 per_tool: HashMap::from([(
17843 "retry_deadline".to_string(),
17844 ToolRetryConfig {
17845 max_retries: 1,
17846 ..Default::default()
17847 },
17848 )]),
17849 ..Default::default()
17850 },
17851 ..Default::default()
17852 }))
17853 .build()
17854 .unwrap();
17855
17856 let record = agent
17857 .invoke_tool(ToolExecutionRequest::new(
17858 "retry-deadline-call",
17859 "retry_deadline",
17860 serde_json::json!({}),
17861 ToolCallSource::Manual,
17862 ))
17863 .await
17864 .unwrap();
17865
17866 assert!(record.executed);
17867 assert!(record.success);
17868 assert_eq!(calls.load(Ordering::SeqCst), 2);
17869 let deadlines = deadlines.lock();
17870 assert_eq!(deadlines.len(), 2);
17871 assert!(
17872 deadlines[1] > deadlines[0],
17873 "retry inherited the first invocation deadline"
17874 );
17875 let remaining_ms = remaining_ms.lock();
17876 assert_eq!(remaining_ms.len(), 2);
17877 assert!(
17878 remaining_ms
17879 .iter()
17880 .all(|remaining| (800..=1_000).contains(remaining))
17881 );
17882 }
17883
17884 #[tokio::test]
17886 async fn call_classification_timeout_controls_deadline_and_timer() {
17887 let calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
17888 let remaining_ms = Arc::new(parking_lot::Mutex::new(Vec::new()));
17889 let agent = AgentBuilder::new()
17890 .system_prompt("Test call-level timeout.")
17891 .llm(Arc::new(mock_with_response("done")))
17892 .tool(Arc::new(ClassifiedTimeoutTool {
17893 id: "classified_timeout",
17894 calls: Arc::clone(&calls),
17895 timeout_ms: 100,
17896 sleep_ms: 150,
17897 requires_approval: false,
17898 remaining_ms: Arc::clone(&remaining_ms),
17899 }))
17900 .build()
17901 .unwrap();
17902
17903 let started = Instant::now();
17904 let record = agent
17905 .invoke_tool(ToolExecutionRequest::new(
17906 "classified-timeout-call",
17907 "classified_timeout",
17908 serde_json::json!({}),
17909 ToolCallSource::Manual,
17910 ))
17911 .await
17912 .unwrap();
17913
17914 assert!(record.executed);
17915 assert!(record.timed_out);
17916 assert!(!record.success);
17917 assert_eq!(calls.load(Ordering::SeqCst), 1);
17918 assert!(started.elapsed() < Duration::from_secs(1));
17919 let remaining_ms = remaining_ms.lock();
17920 assert_eq!(remaining_ms.len(), 1);
17921 assert!((1..=100).contains(&remaining_ms[0]));
17922 }
17923
17924 #[tokio::test]
17926 async fn recovery_timeout_only_lowers_call_and_policy_timeouts() {
17927 use ai_agents_recovery::{ErrorRecoveryConfig, ToolRecoveryConfig, ToolRetryConfig};
17928
17929 let calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
17930 let remaining_ms = Arc::new(parking_lot::Mutex::new(Vec::new()));
17931 let agent = AgentBuilder::new()
17932 .system_prompt("Test recovery timeout.")
17933 .llm(Arc::new(mock_with_response("done")))
17934 .tool(Arc::new(ClassifiedTimeoutTool {
17935 id: "recovery_timeout",
17936 calls: Arc::clone(&calls),
17937 timeout_ms: 1_000,
17938 sleep_ms: 150,
17939 requires_approval: false,
17940 remaining_ms: Arc::clone(&remaining_ms),
17941 }))
17942 .recovery_manager(RecoveryManager::new(ErrorRecoveryConfig {
17943 tools: ToolRecoveryConfig {
17944 per_tool: HashMap::from([(
17945 "recovery_timeout".to_string(),
17946 ToolRetryConfig {
17947 timeout_ms: Some(100),
17948 ..Default::default()
17949 },
17950 )]),
17951 ..Default::default()
17952 },
17953 ..Default::default()
17954 }))
17955 .build()
17956 .unwrap();
17957
17958 let started = Instant::now();
17959 let record = agent
17960 .invoke_tool(ToolExecutionRequest::new(
17961 "recovery-timeout-call",
17962 "recovery_timeout",
17963 serde_json::json!({}),
17964 ToolCallSource::Manual,
17965 ))
17966 .await
17967 .unwrap();
17968
17969 assert!(record.executed);
17970 assert!(record.timed_out);
17971 assert!(!record.success);
17972 assert_eq!(calls.load(Ordering::SeqCst), 1);
17973 assert!(started.elapsed() < Duration::from_secs(1));
17974 assert_eq!(record.metadata["effective_limits"]["timeout_ms"], 100);
17975 let remaining_ms = remaining_ms.lock();
17976 assert_eq!(remaining_ms.len(), 1);
17977 assert!((1..=100).contains(&remaining_ms[0]));
17978 }
17979
17980 #[tokio::test]
17982 async fn recovery_default_timeout_controls_deadline_and_timer() {
17983 use ai_agents_recovery::{ErrorRecoveryConfig, ToolRecoveryConfig, ToolRetryConfig};
17984
17985 let calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
17986 let remaining_ms = Arc::new(parking_lot::Mutex::new(Vec::new()));
17987 let agent = AgentBuilder::new()
17988 .system_prompt("Test default recovery timeout.")
17989 .llm(Arc::new(mock_with_response("done")))
17990 .tool(Arc::new(ClassifiedTimeoutTool {
17991 id: "default_recovery_timeout",
17992 calls: Arc::clone(&calls),
17993 timeout_ms: 1_000,
17994 sleep_ms: 150,
17995 requires_approval: false,
17996 remaining_ms: Arc::clone(&remaining_ms),
17997 }))
17998 .recovery_manager(RecoveryManager::new(ErrorRecoveryConfig {
17999 tools: ToolRecoveryConfig {
18000 default: ToolRetryConfig {
18001 timeout_ms: Some(100),
18002 ..Default::default()
18003 },
18004 ..Default::default()
18005 },
18006 ..Default::default()
18007 }))
18008 .build()
18009 .unwrap();
18010
18011 let started = Instant::now();
18012 let record = agent
18013 .invoke_tool(ToolExecutionRequest::new(
18014 "default-recovery-timeout-call",
18015 "default_recovery_timeout",
18016 serde_json::json!({}),
18017 ToolCallSource::Manual,
18018 ))
18019 .await
18020 .unwrap();
18021
18022 assert!(record.executed);
18023 assert!(record.timed_out);
18024 assert!(!record.success);
18025 assert_eq!(calls.load(Ordering::SeqCst), 1);
18026 assert!(started.elapsed() < Duration::from_secs(1));
18027 assert_eq!(record.metadata["effective_limits"]["timeout_ms"], 100);
18028 let remaining_ms = remaining_ms.lock();
18029 assert_eq!(remaining_ms.len(), 1);
18030 assert!((1..=100).contains(&remaining_ms[0]));
18031 }
18032
18033 #[test]
18035 fn recovery_timeout_cannot_widen_security_baseline() {
18036 let security_engine = ToolSecurityEngine::new(ToolSecurityConfig {
18037 default_timeout_ms: 100,
18038 ..Default::default()
18039 });
18040 let safety = ToolSafetyMetadata::compute();
18041 let mut classification = ToolCallClassification::from_metadata(&safety);
18042 classification.timeout_ms = Some(500);
18043
18044 let (limits, timeout) = RuntimeAgent::effective_tool_limits(
18045 &security_engine,
18046 "recovery_cannot_widen",
18047 &safety,
18048 &classification,
18049 Some(1_000),
18050 )
18051 .unwrap();
18052
18053 assert_eq!(limits.timeout_ms, Some(100));
18054 assert_eq!(timeout.timer, Duration::from_millis(100));
18055 }
18056
18057 #[tokio::test]
18059 async fn invalid_call_timeout_stops_before_approval_or_tool_invocation() {
18060 let tool_calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
18061 let approval_calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
18062 let remaining_ms = Arc::new(parking_lot::Mutex::new(Vec::new()));
18063 let mut security = ToolSecurityConfig {
18064 enabled: true,
18065 ..Default::default()
18066 };
18067 security.tools.insert(
18068 "invalid_call_timeout".to_string(),
18069 ai_agents_tools::ToolPolicyConfig {
18070 require_confirmation: true,
18071 ..Default::default()
18072 },
18073 );
18074 let agent = AgentBuilder::new()
18075 .system_prompt("Test invalid call timeout.")
18076 .llm(Arc::new(mock_with_response("done")))
18077 .tool(Arc::new(ClassifiedTimeoutTool {
18078 id: "invalid_call_timeout",
18079 calls: Arc::clone(&tool_calls),
18080 timeout_ms: u64::MAX,
18081 sleep_ms: 0,
18082 requires_approval: false,
18083 remaining_ms,
18084 }))
18085 .tool_security(ToolSecurityEngine::new(security))
18086 .hitl_engine(HITLEngine::new(ai_agents_hitl::HITLConfig::default()))
18087 .approval_handler(Arc::new(CountingApprovalHandler {
18088 calls: Arc::clone(&approval_calls),
18089 }))
18090 .build()
18091 .unwrap();
18092
18093 let error = agent
18094 .invoke_tool(ToolExecutionRequest::new(
18095 "invalid-call-timeout",
18096 "invalid_call_timeout",
18097 serde_json::json!({}),
18098 ToolCallSource::Manual,
18099 ))
18100 .await
18101 .unwrap_err();
18102
18103 assert!(error.to_string().contains(
18104 "effective tool timeout_ms must be no greater than 3153600000000000 milliseconds"
18105 ));
18106 assert_eq!(approval_calls.load(Ordering::SeqCst), 0);
18107 assert_eq!(tool_calls.load(Ordering::SeqCst), 0);
18108 }
18109
18110 #[tokio::test]
18112 async fn invalid_modified_call_timeout_stops_before_lock_or_invocation() {
18113 use ai_agents_hitl::CallbackHandler;
18114
18115 let blocker_gate = PathMutationGate::new();
18116 let tool_calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
18117 let hooks = Arc::new(ToolLifecycleRecordingHooks::new());
18118 let handler = CallbackHandler::new(|_| ApprovalResult::Modified {
18119 changes: HashMap::from([("invalid_timeout".to_string(), Value::Bool(true))]),
18120 });
18121 let agent = Arc::new(
18122 AgentBuilder::new()
18123 .system_prompt("Test final call timeout validation.")
18124 .llm(Arc::new(mock_with_response("done")))
18125 .tool(Arc::new(BlockingPathMutationTool {
18126 id: "timeout_lock_blocker",
18127 path_fields: vec![ai_agents_core::PathPolicyBinding::write("path")],
18128 gate: blocker_gate.clone(),
18129 }))
18130 .tool(Arc::new(ApprovalModifiedTimeoutTool {
18131 calls: Arc::clone(&tool_calls),
18132 }))
18133 .hitl_engine(HITLEngine::new(ai_agents_hitl::HITLConfig::default()))
18134 .approval_handler(Arc::new(handler))
18135 .hooks(hooks.clone())
18136 .build()
18137 .unwrap(),
18138 );
18139 let blocking_agent = Arc::clone(&agent);
18140 let blocker = tokio::spawn(async move {
18141 blocking_agent
18142 .invoke_tool(ToolExecutionRequest::new(
18143 "timeout-lock-blocker",
18144 "timeout_lock_blocker",
18145 serde_json::json!({"path": "./shared-timeout.txt"}),
18146 ToolCallSource::Manual,
18147 ))
18148 .await
18149 .unwrap()
18150 });
18151 blocker_gate.wait_until_entered().await;
18152
18153 let record = tokio::time::timeout(
18154 Duration::from_millis(500),
18155 agent.invoke_tool(ToolExecutionRequest::new(
18156 "invalid-modified-timeout",
18157 "approval_modified_timeout",
18158 serde_json::json!({
18159 "path": "./shared-timeout.txt",
18160 "invalid_timeout": false
18161 }),
18162 ToolCallSource::Manual,
18163 )),
18164 )
18165 .await
18166 .expect("final timeout validation must not wait for the held path lock")
18167 .unwrap();
18168
18169 blocker_gate.release();
18170 assert!(blocker.await.unwrap().success);
18171 assert!(!record.executed);
18172 assert!(!record.success);
18173 assert_eq!(record.policy.outcome, PermissionOutcome::Deny);
18174 assert!(record.output.contains(
18175 "effective tool timeout_ms must be no greater than 3153600000000000 milliseconds"
18176 ));
18177 assert_eq!(tool_calls.load(Ordering::SeqCst), 0);
18178 let invalid_request_events = hooks
18179 .events()
18180 .into_iter()
18181 .filter(|event| event.contains("approval_modified_timeout") || event == "error")
18182 .collect::<Vec<_>>();
18183 assert_eq!(
18184 invalid_request_events,
18185 vec![
18186 "start:approval_modified_timeout",
18187 "complete:approval_modified_timeout:false",
18188 "record:approval_modified_timeout:false",
18189 "error"
18190 ]
18191 );
18192 }
18193
18194 #[tokio::test]
18195 async fn side_effecting_tools_are_serialized_per_resource() {
18196 let mock = mock_with_response("hello");
18197 let active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
18198 let max_active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
18199 let agent = Arc::new(
18200 AgentBuilder::new()
18201 .system_prompt("You are helpful.")
18202 .llm(Arc::new(mock))
18203 .tool(Arc::new(LockedWriteTool {
18204 active: Arc::clone(&active),
18205 max_active: Arc::clone(&max_active),
18206 }))
18207 .build()
18208 .unwrap(),
18209 );
18210
18211 let left = {
18212 let agent = Arc::clone(&agent);
18213 tokio::spawn(async move {
18214 agent
18215 .invoke_tool(ToolExecutionRequest::new(
18216 "lock-1",
18217 "locked_write",
18218 serde_json::json!({"path": "./same.txt"}),
18219 ToolCallSource::Manual,
18220 ))
18221 .await
18222 .unwrap()
18223 })
18224 };
18225 let right = {
18226 let agent = Arc::clone(&agent);
18227 tokio::spawn(async move {
18228 agent
18229 .invoke_tool(ToolExecutionRequest::new(
18230 "lock-2",
18231 "locked_write",
18232 serde_json::json!({"path": "./same.txt"}),
18233 ToolCallSource::Manual,
18234 ))
18235 .await
18236 .unwrap()
18237 })
18238 };
18239
18240 let left = left.await.unwrap();
18241 let right = right.await.unwrap();
18242 assert!(left.success);
18243 assert!(right.success);
18244 assert_eq!(max_active.load(Ordering::SeqCst), 1);
18245 }
18246
18247 #[tokio::test]
18248 async fn path_resources_use_shared_global_lock_and_cleanup() {
18249 let active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
18250 let max_active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
18251 let bindings = ai_agents_core::ToolPolicyBindings {
18252 path_fields: vec![
18253 ai_agents_core::PathPolicyBinding::read_write("source_path"),
18254 ai_agents_core::PathPolicyBinding::write("destination_path"),
18255 ],
18256 ..Default::default()
18257 };
18258 let classification = ai_agents_core::ToolCallClassification::from_metadata(
18259 &MultiResourceWriteTool {
18260 active: Arc::clone(&active),
18261 max_active: Arc::clone(&max_active),
18262 }
18263 .safety_metadata(),
18264 );
18265 let left_args = serde_json::json!({
18266 "source_path": "./a/../first.txt",
18267 "destination_path": "./second.txt"
18268 });
18269 let right_args = serde_json::json!({
18270 "source_path": "./second.txt",
18271 "destination_path": "./first.txt"
18272 });
18273 let left_keys = tool_resource_lock_keys(
18274 "multi_resource_write",
18275 &left_args,
18276 &bindings,
18277 &classification,
18278 );
18279 let right_keys = tool_resource_lock_keys(
18280 "multi_resource_write",
18281 &right_args,
18282 &bindings,
18283 &classification,
18284 );
18285 assert_eq!(left_keys, right_keys);
18286 assert_eq!(left_keys, vec!["path-mutation:global".to_string()]);
18287
18288 let locks = new_tool_resource_locks();
18289 let build_agent = || {
18290 AgentBuilder::new()
18291 .system_prompt("Test shared resource locks.")
18292 .llm(Arc::new(mock_with_response("done")))
18293 .tool(Arc::new(MultiResourceWriteTool {
18294 active: Arc::clone(&active),
18295 max_active: Arc::clone(&max_active),
18296 }))
18297 .build()
18298 .unwrap()
18299 .with_shared_resource_locks(Arc::clone(&locks))
18300 };
18301 let left_agent = Arc::new(build_agent());
18302 let right_agent = Arc::new(build_agent());
18303 let left = tokio::spawn(async move {
18304 left_agent
18305 .invoke_tool(ToolExecutionRequest::new(
18306 "multi-left",
18307 "multi_resource_write",
18308 left_args,
18309 ToolCallSource::Manual,
18310 ))
18311 .await
18312 .unwrap()
18313 });
18314 let right = tokio::spawn(async move {
18315 right_agent
18316 .invoke_tool(ToolExecutionRequest::new(
18317 "multi-right",
18318 "multi_resource_write",
18319 right_args,
18320 ToolCallSource::Manual,
18321 ))
18322 .await
18323 .unwrap()
18324 });
18325 let (left, right) = tokio::time::timeout(std::time::Duration::from_secs(2), async {
18326 tokio::join!(left, right)
18327 })
18328 .await
18329 .expect("reversed resource acquisition must not deadlock");
18330
18331 assert!(left.unwrap().success);
18332 assert!(right.unwrap().success);
18333 assert_eq!(max_active.load(Ordering::SeqCst), 1);
18334 assert!(locks.read().is_empty());
18335 }
18336
18337 #[tokio::test]
18338 async fn global_path_lock_serializes_copy_destination_with_file_write() {
18339 assert_path_mutation_pair_serialized(
18340 "copy_path",
18341 CopyPathTool::new().policy_bindings().path_fields,
18342 serde_json::json!({
18343 "source_path": "./source.txt",
18344 "destination_path": "./shared.txt"
18345 }),
18346 "file_write",
18347 FileWriteTool::new().policy_bindings().path_fields,
18348 serde_json::json!({"path": "./shared.txt"}),
18349 )
18350 .await;
18351 }
18352
18353 #[tokio::test]
18354 async fn parent_and_spawned_runtime_share_global_path_lock() {
18355 let workspace = MutationTestWorkspace::new();
18356 let destination = workspace.root.join("spawned.txt");
18357 let parent_gate = PathMutationGate::new();
18358 let parent = Arc::new(
18359 AgentBuilder::from_yaml(
18360 r#"
18361name: LockParent
18362system_prompt: parent
18363llm:
18364 default: default
18365tools:
18366 - parent_path_write
18367spawner:
18368 shared_llms: true
18369"#,
18370 )
18371 .unwrap()
18372 .llm(Arc::new(mock_with_response("done")))
18373 .auto_configure_spawner()
18374 .await
18375 .unwrap()
18376 .tool(Arc::new(BlockingPathMutationTool {
18377 id: "parent_path_write",
18378 path_fields: vec![ai_agents_core::PathPolicyBinding::write("path")],
18379 gate: parent_gate.clone(),
18380 }))
18381 .build()
18382 .unwrap(),
18383 );
18384
18385 let mut child_spec = crate::spec::AgentSpec {
18386 name: "LockChild".to_string(),
18387 system_prompt: "child".to_string(),
18388 tools: Some(vec![crate::spec::ToolEntry::Simple(
18389 "file_write".to_string(),
18390 )]),
18391 ..Default::default()
18392 };
18393 child_spec.tool_security.enabled = true;
18394 child_spec.tool_security.fail_closed = true;
18395 let file_write_policy = ai_agents_tools::ToolPolicyConfig {
18396 write_paths: vec![workspace.root.to_string_lossy().into_owned()],
18397 allow_without_confirmation: true,
18398 ..Default::default()
18399 };
18400 child_spec
18401 .tool_security
18402 .tools
18403 .insert("file_write".to_string(), file_write_policy);
18404 let spawned = parent
18405 .spawner()
18406 .unwrap()
18407 .spawn_from_spec(child_spec)
18408 .await
18409 .unwrap();
18410 assert!(Arc::ptr_eq(
18411 &parent.resource_locks,
18412 &spawned.agent.resource_locks
18413 ));
18414 assert!(!Arc::ptr_eq(
18415 &parent.runtime_control,
18416 &spawned.agent.runtime_control
18417 ));
18418
18419 let parent_call = {
18420 let parent = Arc::clone(&parent);
18421 let destination = destination.clone();
18422 tokio::spawn(async move {
18423 parent
18424 .invoke_tool(ToolExecutionRequest::new(
18425 "parent-lock-holder",
18426 "parent_path_write",
18427 serde_json::json!({"path": destination}),
18428 ToolCallSource::Manual,
18429 ))
18430 .await
18431 .unwrap()
18432 })
18433 };
18434 parent_gate.wait_until_entered().await;
18435
18436 let child_call = {
18437 let child = Arc::clone(&spawned.agent);
18438 let destination = destination.clone();
18439 tokio::spawn(async move {
18440 child
18441 .invoke_tool(ToolExecutionRequest::new(
18442 "spawned-file-write",
18443 "file_write",
18444 serde_json::json!({
18445 "path": destination,
18446 "content": "spawned",
18447 "dry_run": false
18448 }),
18449 ToolCallSource::Manual,
18450 ))
18451 .await
18452 .unwrap()
18453 })
18454 };
18455 wait_for_resource_lock_strong_count(&parent.resource_locks, 2).await;
18456 assert!(!child_call.is_finished());
18457
18458 parent_gate.release();
18459 let (parent_record, child_record) =
18460 tokio::time::timeout(std::time::Duration::from_secs(2), async {
18461 tokio::join!(parent_call, child_call)
18462 })
18463 .await
18464 .expect("parent and spawned path mutations did not finish");
18465 assert!(parent_record.unwrap().success);
18466 assert!(child_record.unwrap().success);
18467 assert_eq!(std::fs::read_to_string(destination).unwrap(), "spawned");
18468 assert!(parent.resource_locks.read().is_empty());
18469 }
18470
18471 #[tokio::test]
18472 async fn cancelled_global_path_lock_waiter_does_not_retain_weak_entry() {
18473 let locks = new_tool_resource_locks();
18474 let holder_gate = PathMutationGate::new();
18475 let waiter_gate = PathMutationGate::new();
18476 waiter_gate.release();
18477 let holder = Arc::new(
18478 AgentBuilder::new()
18479 .system_prompt("Hold the global path lock.")
18480 .llm(Arc::new(mock_with_response("done")))
18481 .tool(Arc::new(BlockingPathMutationTool {
18482 id: "holder_write",
18483 path_fields: vec![ai_agents_core::PathPolicyBinding::write("path")],
18484 gate: holder_gate.clone(),
18485 }))
18486 .build()
18487 .unwrap()
18488 .with_shared_resource_locks(Arc::clone(&locks)),
18489 );
18490 let waiter = Arc::new(
18491 AgentBuilder::new()
18492 .system_prompt("Wait for the global path lock.")
18493 .llm(Arc::new(mock_with_response("done")))
18494 .tool(Arc::new(BlockingPathMutationTool {
18495 id: "waiter_write",
18496 path_fields: vec![ai_agents_core::PathPolicyBinding::write("path")],
18497 gate: waiter_gate.clone(),
18498 }))
18499 .build()
18500 .unwrap()
18501 .with_shared_resource_locks(Arc::clone(&locks)),
18502 );
18503
18504 let holder_call = {
18505 let holder = Arc::clone(&holder);
18506 tokio::spawn(async move {
18507 holder
18508 .invoke_tool(ToolExecutionRequest::new(
18509 "holder-call",
18510 "holder_write",
18511 serde_json::json!({"path": "./shared.txt"}),
18512 ToolCallSource::Manual,
18513 ))
18514 .await
18515 .unwrap()
18516 })
18517 };
18518 holder_gate.wait_until_entered().await;
18519
18520 let waiter_call = {
18521 let waiter = Arc::clone(&waiter);
18522 tokio::spawn(async move {
18523 waiter
18524 .invoke_tool(ToolExecutionRequest::new(
18525 "waiter-call",
18526 "waiter_write",
18527 serde_json::json!({"path": "./shared.txt"}),
18528 ToolCallSource::Manual,
18529 ))
18530 .await
18531 .unwrap()
18532 })
18533 };
18534 wait_for_resource_lock_strong_count(&locks, 2).await;
18535 waiter.runtime_control().cancel_all();
18536
18537 let waiter_record = tokio::time::timeout(std::time::Duration::from_secs(2), waiter_call)
18538 .await
18539 .expect("cancelled lock waiter did not finish")
18540 .unwrap();
18541 assert!(!waiter_record.success);
18542 assert!(!waiter_record.executed);
18543 assert!(waiter_record.cancelled);
18544 assert_eq!(
18545 waiter_record.cancellation_reason.as_deref(),
18546 Some("runtime control cancellation")
18547 );
18548 assert!(!waiter_gate.entered.load(Ordering::SeqCst));
18549 assert_eq!(
18550 locks
18551 .read()
18552 .get("path-mutation:global")
18553 .map_or(0, |lock| lock.strong_count()),
18554 1
18555 );
18556
18557 holder_gate.release();
18558 let holder_record = tokio::time::timeout(std::time::Duration::from_secs(2), holder_call)
18559 .await
18560 .expect("lock holder did not finish")
18561 .unwrap();
18562 assert!(holder_record.success);
18563 assert!(locks.read().is_empty());
18564 }
18565
18566 #[tokio::test]
18567 async fn path_mutation_policy_and_approval_denials_do_not_invoke_tools() {
18568 for denial in [MutationDenial::Policy, MutationDenial::Approval] {
18569 let tools: [Arc<dyn Tool>; 3] = [
18570 Arc::new(CopyPathTool::new()),
18571 Arc::new(MovePathTool::new()),
18572 Arc::new(DeletePathTool::new()),
18573 ];
18574 for tool in tools {
18575 assert_path_mutation_denied(tool, denial).await;
18576 }
18577 }
18578 }
18579
18580 #[tokio::test]
18581 async fn policy_denial_keeps_executor_hook_lifecycle_and_record_authority() {
18582 let workspace = MutationTestWorkspace::new();
18583 let target = workspace.root.join("denied.txt");
18584 let hooks = Arc::new(ToolLifecycleRecordingHooks::new());
18585 let agent = AgentBuilder::new()
18586 .system_prompt("Test denied tool hooks.")
18587 .llm(Arc::new(mock_with_response("done")))
18588 .tool(Arc::new(FileWriteTool::new()))
18589 .tool_security(ToolSecurityEngine::new(mutation_denial_security_config(
18590 "file_write",
18591 &workspace.root,
18592 MutationDenial::Policy,
18593 )))
18594 .hooks(hooks.clone())
18595 .build()
18596 .unwrap();
18597
18598 let record = agent
18599 .invoke_tool(ToolExecutionRequest::new(
18600 "denied-hook-call",
18601 "file_write",
18602 serde_json::json!({
18603 "path": target.to_string_lossy(),
18604 "content": "blocked"
18605 }),
18606 ToolCallSource::Manual,
18607 ))
18608 .await
18609 .unwrap();
18610
18611 assert!(!record.executed);
18612 assert!(!record.success);
18613 assert_eq!(record.policy.outcome, PermissionOutcome::Deny);
18614 assert_eq!(
18615 hooks.events(),
18616 vec![
18617 "start:file_write",
18618 "complete:file_write:false",
18619 "record:file_write:false",
18620 "error"
18621 ]
18622 );
18623 assert!(!target.exists());
18624 }
18625
18626 #[tokio::test]
18627 async fn approval_argument_changes_are_rechecked_against_final_scope() {
18628 let active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
18629 let max_active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
18630 let entered = Arc::new(tokio::sync::Barrier::new(2));
18631 let release = Arc::new(tokio::sync::Notify::new());
18632 let handler = Arc::new(BlockingApprovalHandler {
18633 entered: Arc::clone(&entered),
18634 release: Arc::clone(&release),
18635 result: ApprovalResult::Modified {
18636 changes: HashMap::from([(
18637 "path".to_string(),
18638 Value::String("./after-approval.txt".to_string()),
18639 )]),
18640 },
18641 });
18642 let agent = Arc::new(
18643 AgentBuilder::new()
18644 .system_prompt("Test final scope validation.")
18645 .llm(Arc::new(mock_with_response("done")))
18646 .tool(Arc::new(LockedWriteTool {
18647 active: Arc::clone(&active),
18648 max_active: Arc::clone(&max_active),
18649 }))
18650 .tool_security(ToolSecurityEngine::new(approval_security_config(true)))
18651 .hitl_engine(HITLEngine::new(ai_agents_hitl::HITLConfig::default()))
18652 .approval_handler(handler)
18653 .build()
18654 .unwrap(),
18655 );
18656 let control = agent.runtime_control();
18657 let running = Arc::clone(&agent);
18658 let call = tokio::spawn(async move {
18659 running
18660 .invoke_tool(ToolExecutionRequest::new(
18661 "approval-scope",
18662 "locked_write",
18663 serde_json::json!({"path": "./before-approval.txt"}),
18664 ToolCallSource::Manual,
18665 ))
18666 .await
18667 .unwrap()
18668 });
18669 entered.wait().await;
18670 let expected_version = control.set_tool_scope(Vec::new());
18671 release.notify_one();
18672 let record = call.await.unwrap();
18673
18674 assert!(!record.executed);
18675 assert!(!record.success);
18676 assert_eq!(record.runtime_config_version, expected_version);
18677 assert_eq!(record.executed_arguments["path"], "./after-approval.txt");
18678 assert_eq!(max_active.load(Ordering::SeqCst), 0);
18679 assert_eq!(
18680 record.metadata["runtime_scope_snapshot"],
18681 serde_json::json!([])
18682 );
18683 }
18684
18685 #[tokio::test]
18686 async fn approval_is_rechecked_against_final_policy_snapshot() {
18687 let active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
18688 let max_active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
18689 let entered = Arc::new(tokio::sync::Barrier::new(2));
18690 let release = Arc::new(tokio::sync::Notify::new());
18691 let handler = Arc::new(BlockingApprovalHandler {
18692 entered: Arc::clone(&entered),
18693 release: Arc::clone(&release),
18694 result: ApprovalResult::Approved,
18695 });
18696 let agent = Arc::new(
18697 AgentBuilder::new()
18698 .system_prompt("Test final policy validation.")
18699 .llm(Arc::new(mock_with_response("done")))
18700 .tool(Arc::new(LockedWriteTool {
18701 active: Arc::clone(&active),
18702 max_active: Arc::clone(&max_active),
18703 }))
18704 .tool_security(ToolSecurityEngine::new(approval_security_config(true)))
18705 .hitl_engine(HITLEngine::new(ai_agents_hitl::HITLConfig::default()))
18706 .approval_handler(handler)
18707 .build()
18708 .unwrap(),
18709 );
18710 let control = agent.runtime_control();
18711 let running = Arc::clone(&agent);
18712 let call = tokio::spawn(async move {
18713 running
18714 .invoke_tool(ToolExecutionRequest::new(
18715 "approval-policy",
18716 "locked_write",
18717 serde_json::json!({"path": "./policy.txt"}),
18718 ToolCallSource::Manual,
18719 ))
18720 .await
18721 .unwrap()
18722 });
18723 entered.wait().await;
18724 let expected_version = control.set_tool_security(approval_security_config(false));
18725 release.notify_one();
18726 let record = call.await.unwrap();
18727
18728 assert!(!record.executed);
18729 assert!(!record.success);
18730 assert_eq!(record.runtime_config_version, expected_version);
18731 assert_eq!(record.policy.outcome, PermissionOutcome::Unavailable);
18732 assert_eq!(max_active.load(Ordering::SeqCst), 0);
18733 assert!(record.metadata.contains_key("policy_snapshot"));
18734 }
18735
18736 #[test]
18737 fn invalid_live_policy_does_not_replace_snapshot_or_generation() {
18738 let agent = AgentBuilder::new()
18739 .system_prompt("Test runtime policy validation.")
18740 .llm(Arc::new(mock_with_response("done")))
18741 .build()
18742 .unwrap();
18743 let control = agent.runtime_control();
18744 let mut valid = ToolSecurityConfig::default();
18745 valid.tools.insert(
18746 "web_search".to_string(),
18747 ai_agents_tools::ToolPolicyConfig {
18748 max_results: Some(5),
18749 ..Default::default()
18750 },
18751 );
18752 let generation = control.try_set_tool_security(valid).unwrap();
18753
18754 let mut invalid = ToolSecurityConfig::default();
18755 invalid.tools.insert(
18756 "web_search".to_string(),
18757 ai_agents_tools::ToolPolicyConfig {
18758 max_results: Some(0),
18759 ..Default::default()
18760 },
18761 );
18762 let error = control.try_set_tool_security(invalid).unwrap_err();
18763
18764 assert!(
18765 error
18766 .to_string()
18767 .contains("max_results must be greater than 0")
18768 );
18769 assert_eq!(control.version(), generation);
18770 assert_eq!(
18771 control
18772 .state
18773 .tool_security_override
18774 .read()
18775 .as_ref()
18776 .unwrap()
18777 .config()
18778 .tools["web_search"]
18779 .max_results,
18780 Some(5)
18781 );
18782 }
18783
18784 #[test]
18786 fn invalid_timeout_config_stops_before_approval_or_tool_invocation() {
18787 let tool_calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
18788 let approval_calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
18789 let spec = crate::spec::AgentSpec {
18790 tool_security: ToolSecurityConfig {
18791 enabled: true,
18792 default_timeout_ms: u64::MAX,
18793 ..Default::default()
18794 },
18795 ..Default::default()
18796 };
18797
18798 let result = AgentBuilder::from_spec(spec)
18799 .llm(Arc::new(mock_with_response("done")))
18800 .tool(Arc::new(FlakyWriteTool {
18801 calls: Arc::clone(&tool_calls),
18802 }))
18803 .hitl_engine(HITLEngine::new(ai_agents_hitl::HITLConfig::default()))
18804 .approval_handler(Arc::new(CountingApprovalHandler {
18805 calls: Arc::clone(&approval_calls),
18806 }))
18807 .build();
18808
18809 assert!(result.is_err());
18810 assert_eq!(approval_calls.load(Ordering::SeqCst), 0);
18811 assert_eq!(tool_calls.load(Ordering::SeqCst), 0);
18812 }
18813
18814 #[test]
18816 fn invalid_recovery_timeout_config_stops_before_approval_or_tool_invocation() {
18817 use ai_agents_recovery::{ErrorRecoveryConfig, ToolRecoveryConfig, ToolRetryConfig};
18818
18819 let tool_calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
18820 let approval_calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
18821 let spec = crate::spec::AgentSpec {
18822 error_recovery: ErrorRecoveryConfig {
18823 tools: ToolRecoveryConfig {
18824 default: ToolRetryConfig {
18825 timeout_ms: Some(u64::MAX),
18826 ..Default::default()
18827 },
18828 ..Default::default()
18829 },
18830 ..Default::default()
18831 },
18832 ..Default::default()
18833 };
18834
18835 let result = AgentBuilder::from_spec(spec)
18836 .llm(Arc::new(mock_with_response("done")))
18837 .tool(Arc::new(FlakyWriteTool {
18838 calls: Arc::clone(&tool_calls),
18839 }))
18840 .hitl_engine(HITLEngine::new(ai_agents_hitl::HITLConfig::default()))
18841 .approval_handler(Arc::new(CountingApprovalHandler {
18842 calls: Arc::clone(&approval_calls),
18843 }))
18844 .build();
18845
18846 assert!(result.is_err());
18847 assert_eq!(approval_calls.load(Ordering::SeqCst), 0);
18848 assert_eq!(tool_calls.load(Ordering::SeqCst), 0);
18849 }
18850
18851 #[test]
18853 fn invalid_timeout_policy_does_not_replace_snapshot_or_generation() {
18854 let agent = AgentBuilder::new()
18855 .system_prompt("Test runtime timeout policy validation.")
18856 .llm(Arc::new(mock_with_response("done")))
18857 .build()
18858 .unwrap();
18859 let control = agent.runtime_control();
18860 let valid = ToolSecurityConfig {
18861 default_timeout_ms: 5_000,
18862 ..Default::default()
18863 };
18864 let generation = control.try_set_tool_security(valid).unwrap();
18865
18866 let invalid = ToolSecurityConfig {
18867 default_timeout_ms: MAX_TOOL_TIMEOUT_MS + 1,
18868 ..Default::default()
18869 };
18870 let error = control.try_set_tool_security(invalid).unwrap_err();
18871
18872 assert!(error.to_string().contains(&format!(
18873 "tool_security.default_timeout_ms must be no greater than {MAX_TOOL_TIMEOUT_MS} milliseconds"
18874 )));
18875 assert_eq!(control.version(), generation);
18876 assert_eq!(
18877 control
18878 .state
18879 .tool_security_override
18880 .read()
18881 .as_ref()
18882 .unwrap()
18883 .config()
18884 .default_timeout_ms,
18885 5_000
18886 );
18887 }
18888
18889 #[test]
18891 fn runtime_tool_timeout_conversion_enforces_the_stable_boundary() {
18892 let timeout = RuntimeAgent::validated_tool_timeout(MAX_TOOL_TIMEOUT_MS).unwrap();
18893 assert_eq!(timeout.timer, Duration::from_millis(MAX_TOOL_TIMEOUT_MS));
18894 assert_eq!(
18895 timeout.deadline_delta,
18896 chrono::Duration::milliseconds(MAX_TOOL_TIMEOUT_MS as i64)
18897 );
18898
18899 for timeout_ms in [MAX_TOOL_TIMEOUT_MS + 1, u64::MAX] {
18900 let error = RuntimeAgent::validated_tool_timeout(timeout_ms).unwrap_err();
18901 assert!(error.to_string().contains(&format!(
18902 "effective tool timeout_ms must be no greater than {MAX_TOOL_TIMEOUT_MS} milliseconds"
18903 )));
18904 }
18905 }
18906
18907 #[tokio::test]
18908 async fn persistent_override_preserves_rate_history_within_generation() {
18909 let calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
18910 let agent = AgentBuilder::new()
18911 .system_prompt("Test persistent policy overrides.")
18912 .llm(Arc::new(mock_with_response("done")))
18913 .tool(Arc::new(RecoveryTestTool {
18914 id: "limited_override".to_string(),
18915 succeeds: true,
18916 calls: Arc::clone(&calls),
18917 max_output_chars: None,
18918 }))
18919 .build()
18920 .unwrap();
18921 let mut security = ToolSecurityConfig {
18922 enabled: true,
18923 fail_closed: true,
18924 ..Default::default()
18925 };
18926 let policy = ai_agents_tools::ToolPolicyConfig {
18927 write_paths: vec![".".to_string()],
18928 rate_limit: Some(1),
18929 ..Default::default()
18930 };
18931 security
18932 .tools
18933 .insert("limited_override".to_string(), policy);
18934 let generation = agent.runtime_control().set_tool_security(security);
18935
18936 let first = agent
18937 .invoke_tool(ToolExecutionRequest::new(
18938 "limited-first",
18939 "limited_override",
18940 serde_json::json!({"path": "./limited.txt"}),
18941 ToolCallSource::Manual,
18942 ))
18943 .await
18944 .unwrap();
18945 let second = agent
18946 .invoke_tool(ToolExecutionRequest::new(
18947 "limited-second",
18948 "limited_override",
18949 serde_json::json!({"path": "./limited.txt"}),
18950 ToolCallSource::Manual,
18951 ))
18952 .await
18953 .unwrap();
18954
18955 assert!(first.success);
18956 assert_eq!(first.policy_version, generation);
18957 assert!(!second.executed);
18958 assert!(second.output.contains("Rate limit exceeded"));
18959 assert_eq!(second.policy_version, generation);
18960 assert_eq!(calls.load(Ordering::SeqCst), 1);
18961 }
18962
18963 #[tokio::test]
18964 async fn concurrent_rate_admission_consumes_capacity_atomically() {
18965 let calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
18966 let tool = Arc::new(RecoveryTestTool {
18967 id: "atomic_rate".to_string(),
18968 succeeds: true,
18969 calls: Arc::clone(&calls),
18970 max_output_chars: None,
18971 });
18972 let arguments = serde_json::json!({"path": "./atomic-rate.txt"});
18973 let bindings = tool.policy_bindings();
18974 let classification = tool.classify_call(&arguments);
18975 let resource_keys =
18976 tool_resource_lock_keys(tool.id(), &arguments, &bindings, &classification);
18977 let mut security = ToolSecurityConfig {
18978 enabled: true,
18979 fail_closed: true,
18980 ..Default::default()
18981 };
18982 let policy = ai_agents_tools::ToolPolicyConfig {
18983 write_paths: vec![".".to_string()],
18984 rate_limit: Some(1),
18985 ..Default::default()
18986 };
18987 security.tools.insert(tool.id().to_string(), policy);
18988 let agent = Arc::new(
18989 AgentBuilder::new()
18990 .system_prompt("Test atomic rate admission.")
18991 .llm(Arc::new(mock_with_response("done")))
18992 .tool(tool)
18993 .tool_security(ToolSecurityEngine::new(security))
18994 .build()
18995 .unwrap(),
18996 );
18997 let held = agent
18998 .acquire_tool_resource_locks(&resource_keys)
18999 .await
19000 .unwrap();
19001 let left = {
19002 let agent = Arc::clone(&agent);
19003 let arguments = arguments.clone();
19004 tokio::spawn(async move {
19005 agent
19006 .invoke_tool(ToolExecutionRequest::new(
19007 "atomic-rate-left",
19008 "atomic_rate",
19009 arguments,
19010 ToolCallSource::Manual,
19011 ))
19012 .await
19013 .unwrap()
19014 })
19015 };
19016 let right = {
19017 let agent = Arc::clone(&agent);
19018 tokio::spawn(async move {
19019 agent
19020 .invoke_tool(ToolExecutionRequest::new(
19021 "atomic-rate-right",
19022 "atomic_rate",
19023 arguments,
19024 ToolCallSource::Manual,
19025 ))
19026 .await
19027 .unwrap()
19028 })
19029 };
19030 tokio::time::sleep(std::time::Duration::from_millis(25)).await;
19031 drop(held);
19032 let (left, right) = tokio::join!(left, right);
19033 let records = [left.unwrap(), right.unwrap()];
19034
19035 assert_eq!(records.iter().filter(|record| record.success).count(), 1);
19036 assert_eq!(records.iter().filter(|record| record.executed).count(), 1);
19037 assert!(
19038 records.iter().any(|record| {
19039 !record.executed && record.output.contains("Rate limit exceeded")
19040 })
19041 );
19042 assert_eq!(calls.load(Ordering::SeqCst), 1);
19043 }
19044
19045 #[tokio::test]
19046 async fn changed_policy_generation_invalidates_pending_approval() {
19047 let active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19048 let max_active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19049 let entered = Arc::new(tokio::sync::Barrier::new(2));
19050 let release = Arc::new(tokio::sync::Notify::new());
19051 let handler = Arc::new(BlockingApprovalHandler {
19052 entered: Arc::clone(&entered),
19053 release: Arc::clone(&release),
19054 result: ApprovalResult::Approved,
19055 });
19056 let agent = Arc::new(
19057 AgentBuilder::new()
19058 .system_prompt("Test stale approval denial.")
19059 .llm(Arc::new(mock_with_response("done")))
19060 .tool(Arc::new(LockedWriteTool {
19061 active: Arc::clone(&active),
19062 max_active: Arc::clone(&max_active),
19063 }))
19064 .tool_security(ToolSecurityEngine::new(approval_security_config(true)))
19065 .hitl_engine(HITLEngine::new(ai_agents_hitl::HITLConfig::default()))
19066 .approval_handler(handler)
19067 .build()
19068 .unwrap(),
19069 );
19070 let running = Arc::clone(&agent);
19071 let call = tokio::spawn(async move {
19072 running
19073 .invoke_tool(ToolExecutionRequest::new(
19074 "stale-approval",
19075 "locked_write",
19076 serde_json::json!({"path": "./stale.txt"}),
19077 ToolCallSource::Manual,
19078 ))
19079 .await
19080 .unwrap()
19081 });
19082 entered.wait().await;
19083 let generation = agent
19084 .runtime_control()
19085 .set_tool_security(approval_security_config(true));
19086 release.notify_one();
19087 let record = call.await.unwrap();
19088
19089 assert!(!record.executed);
19090 assert!(record.output.contains("Approval became stale"));
19091 assert_eq!(record.policy_version, generation);
19092 assert_eq!(max_active.load(Ordering::SeqCst), 0);
19093 }
19094
19095 #[tokio::test]
19096 async fn final_policy_reapplies_argument_caps_after_approval_changes() {
19097 use ai_agents_hitl::CallbackHandler;
19098
19099 let mut security = ToolSecurityConfig {
19100 enabled: true,
19101 fail_closed: true,
19102 ..Default::default()
19103 };
19104 let policy = ai_agents_tools::ToolPolicyConfig {
19105 read_paths: vec![".".to_string()],
19106 max_results: Some(5),
19107 require_confirmation: true,
19108 ..Default::default()
19109 };
19110 security.tools.insert("context_echo".to_string(), policy);
19111 let handler = CallbackHandler::new(|_| ApprovalResult::Modified {
19112 changes: HashMap::from([("max_results".to_string(), serde_json::json!(99))]),
19113 });
19114 let agent = AgentBuilder::new()
19115 .system_prompt("Test final argument caps.")
19116 .llm(Arc::new(mock_with_response("done")))
19117 .tool(Arc::new(ContextEchoTool))
19118 .tool_security(ToolSecurityEngine::new(security))
19119 .hitl_engine(HITLEngine::new(ai_agents_hitl::HITLConfig::default()))
19120 .approval_handler(Arc::new(handler))
19121 .build()
19122 .unwrap();
19123
19124 let record = agent
19125 .invoke_tool(ToolExecutionRequest::new(
19126 "final-cap",
19127 "context_echo",
19128 serde_json::json!({"path": ".", "max_results": 1}),
19129 ToolCallSource::Manual,
19130 ))
19131 .await
19132 .unwrap();
19133
19134 assert!(record.success);
19135 assert_eq!(record.executed_arguments["max_results"], 5);
19136 assert_eq!(
19137 record.approval.unwrap().modified_arguments.unwrap()["max_results"],
19138 5
19139 );
19140 }
19141
19142 #[tokio::test]
19143 async fn no_binding_writes_use_canonical_fallback_lock() {
19144 let active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19145 let max_active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19146 let agent = Arc::new(
19147 AgentBuilder::new()
19148 .system_prompt("Test fallback resource locks.")
19149 .llm(Arc::new(mock_with_response("done")))
19150 .tool(Arc::new(NoBindingWriteTool {
19151 active: Arc::clone(&active),
19152 max_active: Arc::clone(&max_active),
19153 }))
19154 .build()
19155 .unwrap(),
19156 );
19157 let left = {
19158 let agent = Arc::clone(&agent);
19159 tokio::spawn(async move {
19160 agent
19161 .invoke_tool(ToolExecutionRequest::new(
19162 "no-binding-left",
19163 "no_binding_write",
19164 serde_json::json!({}),
19165 ToolCallSource::Manual,
19166 ))
19167 .await
19168 .unwrap()
19169 })
19170 };
19171 let right = {
19172 let agent = Arc::clone(&agent);
19173 tokio::spawn(async move {
19174 agent
19175 .invoke_tool(ToolExecutionRequest::new(
19176 "no-binding-right",
19177 "no_binding_write",
19178 serde_json::json!({}),
19179 ToolCallSource::Manual,
19180 ))
19181 .await
19182 .unwrap()
19183 })
19184 };
19185 let (left, right) = tokio::join!(left, right);
19186
19187 assert!(left.unwrap().success);
19188 assert!(right.unwrap().success);
19189 assert_eq!(max_active.load(Ordering::SeqCst), 1);
19190 }
19191
19192 #[tokio::test]
19193 async fn parent_and_child_paths_share_a_resource_lock() {
19194 let active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19195 let max_active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19196 let agent = Arc::new(
19197 AgentBuilder::new()
19198 .system_prompt("Test parent child resource locks.")
19199 .llm(Arc::new(mock_with_response("done")))
19200 .tool(Arc::new(LockedWriteTool {
19201 active: Arc::clone(&active),
19202 max_active: Arc::clone(&max_active),
19203 }))
19204 .build()
19205 .unwrap(),
19206 );
19207 let parent = format!("./lock-parent-{}", uuid::Uuid::new_v4());
19208 let child = format!("{}/child.txt", parent);
19209 let left = {
19210 let agent = Arc::clone(&agent);
19211 tokio::spawn(async move {
19212 agent
19213 .invoke_tool(ToolExecutionRequest::new(
19214 "parent-lock",
19215 "locked_write",
19216 serde_json::json!({"path": parent}),
19217 ToolCallSource::Manual,
19218 ))
19219 .await
19220 .unwrap()
19221 })
19222 };
19223 let right = {
19224 let agent = Arc::clone(&agent);
19225 tokio::spawn(async move {
19226 agent
19227 .invoke_tool(ToolExecutionRequest::new(
19228 "child-lock",
19229 "locked_write",
19230 serde_json::json!({"path": child}),
19231 ToolCallSource::Manual,
19232 ))
19233 .await
19234 .unwrap()
19235 })
19236 };
19237 let (left, right) = tokio::join!(left, right);
19238
19239 assert!(left.unwrap().success);
19240 assert!(right.unwrap().success);
19241 assert_eq!(max_active.load(Ordering::SeqCst), 1);
19242 }
19243
19244 #[tokio::test]
19245 async fn tool_hooks_can_reenter_after_resource_guards_are_dropped() {
19246 let calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19247 let hooks = Arc::new(ReentrantToolHooks {
19248 agent: parking_lot::Mutex::new(None),
19249 invoked: AtomicBool::new(false),
19250 nested_success: AtomicBool::new(false),
19251 });
19252 let agent = Arc::new(
19253 AgentBuilder::new()
19254 .system_prompt("Test hook reentrancy.")
19255 .llm(Arc::new(mock_with_response("done")))
19256 .tool(Arc::new(RecoveryTestTool {
19257 id: "reentrant_write".to_string(),
19258 succeeds: true,
19259 calls: Arc::clone(&calls),
19260 max_output_chars: None,
19261 }))
19262 .hooks(hooks.clone())
19263 .build()
19264 .unwrap(),
19265 );
19266 *hooks.agent.lock() = Some(Arc::downgrade(&agent));
19267 let record = tokio::time::timeout(
19268 std::time::Duration::from_secs(2),
19269 agent.invoke_tool(ToolExecutionRequest::new(
19270 "outer-hook-call",
19271 "reentrant_write",
19272 serde_json::json!({"path": "./hook.txt"}),
19273 ToolCallSource::Manual,
19274 )),
19275 )
19276 .await
19277 .expect("tool completion hook must not retain resource guards")
19278 .unwrap();
19279
19280 assert!(record.success);
19281 assert!(hooks.nested_success.load(Ordering::SeqCst));
19282 assert_eq!(calls.load(Ordering::SeqCst), 2);
19283 }
19284
19285 #[tokio::test]
19287 async fn fallback_finalizes_original_record_before_shared_execution() {
19288 let primary_calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19289 let fallback_calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19290 let hooks = Arc::new(ToolLifecycleRecordingHooks::new());
19291 let agent = AgentBuilder::new()
19292 .system_prompt("Test fallback execution.")
19293 .llm(Arc::new(mock_with_response("done")))
19294 .tool(Arc::new(RecoveryTestTool {
19295 id: "primary".to_string(),
19296 succeeds: false,
19297 calls: Arc::clone(&primary_calls),
19298 max_output_chars: None,
19299 }))
19300 .tool(Arc::new(RecoveryTestTool {
19301 id: "fallback".to_string(),
19302 succeeds: true,
19303 calls: Arc::clone(&fallback_calls),
19304 max_output_chars: None,
19305 }))
19306 .recovery_manager(recovery_manager_with_fallbacks([(
19307 "primary".to_string(),
19308 "fallback".to_string(),
19309 )]))
19310 .hooks(hooks.clone())
19311 .build()
19312 .unwrap();
19313 let record = tokio::time::timeout(
19314 std::time::Duration::from_secs(2),
19315 agent.invoke_tool(ToolExecutionRequest::new(
19316 "fallback-call",
19317 "primary",
19318 serde_json::json!({"path": "./shared.txt"}),
19319 ToolCallSource::Manual,
19320 )),
19321 )
19322 .await
19323 .expect("fallback must not retain the primary resource guard")
19324 .unwrap();
19325
19326 assert_eq!(
19327 hooks.events(),
19328 vec![
19329 "start:primary",
19330 "complete:primary:false",
19331 "record:primary:true",
19332 "error",
19333 "start:fallback",
19334 "complete:fallback:true",
19335 "record:fallback:true",
19336 ]
19337 );
19338 let records = hooks.records();
19339 assert_eq!(records.len(), 2);
19340 let original = &records[0];
19341 assert_eq!(original.canonical_id, "primary");
19342 assert!(matches!(original.source, ToolCallSource::Manual));
19343 assert!(original.executed);
19344 assert!(!original.success);
19345
19346 let fallback = &records[1];
19347 assert_eq!(fallback.canonical_id, "fallback");
19348 assert_eq!(fallback.call_id, "fallback-call");
19349 assert!(matches!(
19350 &fallback.source,
19351 ToolCallSource::Fallback { original_tool } if original_tool == "primary"
19352 ));
19353 assert!(fallback.executed);
19354 assert!(fallback.success);
19355 assert_eq!(record.canonical_id, fallback.canonical_id);
19356 assert_eq!(record.output, fallback.output);
19357
19358 let history = agent.tool_call_history();
19359 assert_eq!(
19360 history
19361 .iter()
19362 .map(|entry| entry.tool_id.as_str())
19363 .collect::<Vec<_>>(),
19364 vec!["primary", "fallback"]
19365 );
19366 assert_eq!(history[0].result.get("success"), Some(&Value::Bool(false)));
19367 assert_eq!(primary_calls.load(Ordering::SeqCst), 1);
19368 assert_eq!(fallback_calls.load(Ordering::SeqCst), 1);
19369 }
19370
19371 #[tokio::test]
19373 async fn self_fallback_cycle_is_denied_before_reinvocation() {
19374 let calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19375 let hooks = Arc::new(ToolLifecycleRecordingHooks::new());
19376 let agent = AgentBuilder::new()
19377 .system_prompt("Test self-fallback cycle admission.")
19378 .llm(Arc::new(mock_with_response("done")))
19379 .tool(Arc::new(RecoveryTestTool {
19380 id: "primary".to_string(),
19381 succeeds: false,
19382 calls: Arc::clone(&calls),
19383 max_output_chars: None,
19384 }))
19385 .recovery_manager(recovery_manager_with_fallbacks([(
19386 "primary".to_string(),
19387 "primary".to_string(),
19388 )]))
19389 .hooks(hooks.clone())
19390 .build()
19391 .unwrap();
19392
19393 let record = tokio::time::timeout(
19394 std::time::Duration::from_secs(2),
19395 agent.invoke_tool(ToolExecutionRequest::new(
19396 "self-fallback-call",
19397 "primary",
19398 serde_json::json!({"path": "./shared.txt"}),
19399 ToolCallSource::Manual,
19400 )),
19401 )
19402 .await
19403 .expect("self fallback must terminate without recursive execution")
19404 .unwrap();
19405
19406 assert_eq!(calls.load(Ordering::SeqCst), 1);
19407 assert_eq!(record.canonical_id, "primary");
19408 assert!(!record.executed);
19409 assert!(!record.success);
19410 assert_eq!(record.policy.outcome, PermissionOutcome::Deny);
19411 assert!(record.output.contains("fallback cycle"));
19412 assert!(matches!(
19413 record.source,
19414 ToolCallSource::Fallback { ref original_tool } if original_tool == "primary"
19415 ));
19416 assert_eq!(
19417 record.metadata.get("fallback_chain"),
19418 Some(&serde_json::json!(["primary"]))
19419 );
19420 assert_eq!(
19421 hooks.events(),
19422 vec![
19423 "start:primary",
19424 "complete:primary:false",
19425 "record:primary:true",
19426 "error",
19427 "complete:primary:false",
19428 "record:primary:false",
19429 "error",
19430 ]
19431 );
19432 assert_eq!(agent.tool_call_history().len(), 2);
19433 }
19434
19435 #[tokio::test]
19437 async fn alias_mediated_fallback_cycle_is_denied_canonically() {
19438 let primary_calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19439 let secondary_calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19440 let hooks = Arc::new(ToolLifecycleRecordingHooks::new());
19441 let agent = AgentBuilder::new()
19442 .system_prompt("Test canonical fallback cycle admission.")
19443 .llm(Arc::new(mock_with_response("done")))
19444 .tool(Arc::new(RecoveryTestTool {
19445 id: "primary".to_string(),
19446 succeeds: false,
19447 calls: Arc::clone(&primary_calls),
19448 max_output_chars: None,
19449 }))
19450 .tool(Arc::new(RecoveryTestTool {
19451 id: "secondary".to_string(),
19452 succeeds: false,
19453 calls: Arc::clone(&secondary_calls),
19454 max_output_chars: None,
19455 }))
19456 .recovery_manager(recovery_manager_with_fallbacks([
19457 ("primary".to_string(), "secondary".to_string()),
19458 ("secondary".to_string(), "primary alias".to_string()),
19459 ]))
19460 .hooks(hooks.clone())
19461 .build()
19462 .unwrap();
19463 agent.tools.set_tool_aliases(
19464 "primary",
19465 ToolAliases::new().with_name("en", "primary alias"),
19466 );
19467
19468 let record = tokio::time::timeout(
19469 std::time::Duration::from_secs(2),
19470 agent.invoke_tool(ToolExecutionRequest::new(
19471 "alias-fallback-call",
19472 "primary",
19473 serde_json::json!({"path": "./shared.txt"}),
19474 ToolCallSource::Manual,
19475 )),
19476 )
19477 .await
19478 .expect("alias-mediated fallback cycle must terminate")
19479 .unwrap();
19480
19481 assert_eq!(primary_calls.load(Ordering::SeqCst), 1);
19482 assert_eq!(secondary_calls.load(Ordering::SeqCst), 1);
19483 assert_eq!(record.requested_name, "primary alias");
19484 assert_eq!(record.canonical_id, "primary");
19485 assert!(!record.executed);
19486 assert!(record.output.contains("fallback cycle"));
19487 assert_eq!(
19488 record.metadata.get("fallback_chain"),
19489 Some(&serde_json::json!(["primary", "secondary"]))
19490 );
19491 assert_eq!(hooks.records().len(), 3);
19492 assert_eq!(agent.tool_call_history().len(), 3);
19493 }
19494
19495 #[tokio::test]
19497 async fn final_canonical_drift_cannot_bypass_fallback_ancestry() {
19498 let primary_calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19499 let secondary_calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19500 let provider = Arc::new(DriftingFallbackProvider {
19501 refreshed: AtomicBool::new(false),
19502 primary_calls: Arc::clone(&primary_calls),
19503 secondary_calls: Arc::clone(&secondary_calls),
19504 });
19505 let registry = ToolRegistry::new();
19506 registry.register_provider(provider).await.unwrap();
19507 let lifecycle = Arc::new(ToolLifecycleRecordingHooks::new());
19508 let hooks = Arc::new(RefreshFallbackProviderHooks {
19509 agent: parking_lot::Mutex::new(None),
19510 lifecycle: Arc::clone(&lifecycle),
19511 });
19512 let agent = Arc::new(
19513 AgentBuilder::new()
19514 .system_prompt("Test final canonical fallback admission.")
19515 .llm(Arc::new(mock_with_response("done")))
19516 .tools(registry)
19517 .recovery_manager(recovery_manager_with_fallbacks([(
19518 "primary".to_string(),
19519 "fallback alias".to_string(),
19520 )]))
19521 .hooks(hooks.clone())
19522 .build()
19523 .unwrap(),
19524 );
19525 *hooks.agent.lock() = Some(Arc::downgrade(&agent));
19526
19527 let record = agent
19528 .invoke_tool(ToolExecutionRequest::new(
19529 "drifting-fallback-call",
19530 "primary",
19531 serde_json::json!({"path": "./shared.txt"}),
19532 ToolCallSource::Manual,
19533 ))
19534 .await
19535 .unwrap();
19536
19537 assert_eq!(primary_calls.load(Ordering::SeqCst), 1);
19538 assert_eq!(secondary_calls.load(Ordering::SeqCst), 0);
19539 assert_eq!(record.canonical_id, "secondary");
19540 assert!(!record.executed);
19541 assert_eq!(record.policy.outcome, PermissionOutcome::Deny);
19542 assert!(record.output.contains("fallback cycle"));
19543 assert_eq!(
19544 record.metadata.get("fallback_chain"),
19545 Some(&serde_json::json!(["primary", "secondary"]))
19546 );
19547 assert_eq!(
19548 record.metadata.get("final_resolved_canonical_id"),
19549 Some(&serde_json::json!("primary"))
19550 );
19551 assert_eq!(
19552 lifecycle.events(),
19553 vec![
19554 "start:primary",
19555 "complete:primary:false",
19556 "record:primary:true",
19557 "error",
19558 "start:secondary",
19559 "complete:secondary:false",
19560 "record:secondary:false",
19561 "error",
19562 ]
19563 );
19564 let records = lifecycle.records();
19565 assert_eq!(records.len(), 2);
19566 assert_eq!(records[1].canonical_id, "secondary");
19567 assert_eq!(
19568 records[1].metadata.get("final_resolved_canonical_id"),
19569 Some(&serde_json::json!("primary"))
19570 );
19571 let history = agent.tool_call_history();
19572 assert_eq!(
19573 history
19574 .iter()
19575 .map(|entry| entry.tool_id.as_str())
19576 .collect::<Vec<_>>(),
19577 vec!["primary", "secondary"]
19578 );
19579 }
19580
19581 #[tokio::test]
19583 async fn acyclic_fallback_chain_is_denied_after_the_hop_limit() {
19584 let tool_count = MAX_TOOL_FALLBACK_HOPS + 2;
19585 let calls = (0..tool_count)
19586 .map(|_| Arc::new(std::sync::atomic::AtomicUsize::new(0)))
19587 .collect::<Vec<_>>();
19588 let mut builder = AgentBuilder::new()
19589 .system_prompt("Test bounded acyclic fallback admission.")
19590 .llm(Arc::new(mock_with_response("done")));
19591 for (index, counter) in calls.iter().enumerate() {
19592 builder = builder.tool(Arc::new(RecoveryTestTool {
19593 id: format!("fallback_{index}"),
19594 succeeds: false,
19595 calls: Arc::clone(counter),
19596 max_output_chars: None,
19597 }));
19598 }
19599 let fallbacks = (0..tool_count - 1).map(|index| {
19600 (
19601 format!("fallback_{index}"),
19602 format!("fallback_{}", index + 1),
19603 )
19604 });
19605 let agent = builder
19606 .recovery_manager(recovery_manager_with_fallbacks(fallbacks))
19607 .build()
19608 .unwrap();
19609
19610 let record = tokio::time::timeout(
19611 std::time::Duration::from_secs(2),
19612 agent.invoke_tool(ToolExecutionRequest::new(
19613 "bounded-fallback-call",
19614 "fallback_0",
19615 serde_json::json!({"path": "./shared.txt"}),
19616 ToolCallSource::Manual,
19617 )),
19618 )
19619 .await
19620 .expect("bounded fallback chain must terminate")
19621 .unwrap();
19622
19623 for counter in calls.iter().take(MAX_TOOL_FALLBACK_HOPS + 1) {
19624 assert_eq!(counter.load(Ordering::SeqCst), 1);
19625 }
19626 assert_eq!(calls[MAX_TOOL_FALLBACK_HOPS + 1].load(Ordering::SeqCst), 0);
19627 assert_eq!(
19628 record.canonical_id,
19629 format!("fallback_{}", MAX_TOOL_FALLBACK_HOPS + 1)
19630 );
19631 assert!(!record.executed);
19632 assert!(record.output.contains("maximum of 16 hops"));
19633 assert_eq!(agent.tool_call_history().len(), tool_count);
19634 }
19635
19636 #[tokio::test]
19637 async fn diagnostics_without_provider_records_unavailable_without_execution() {
19638 let mock = mock_with_response("hello");
19639 let yaml = r#"
19640name: DiagnosticsNoProviderAgent
19641system_prompt: "Review diagnostics."
19642tools: [diagnostics]
19643"#;
19644 let agent = AgentBuilder::from_yaml(yaml)
19645 .unwrap()
19646 .llm(Arc::new(mock))
19647 .auto_configure_features()
19648 .unwrap()
19649 .build()
19650 .unwrap();
19651
19652 let record = agent
19653 .invoke_tool(ToolExecutionRequest::new(
19654 "diagnostics-call",
19655 "diagnostics",
19656 serde_json::json!({}),
19657 ToolCallSource::Manual,
19658 ))
19659 .await
19660 .unwrap();
19661
19662 assert!(!record.executed);
19663 assert!(!record.success);
19664 assert_eq!(record.policy.outcome, PermissionOutcome::Unavailable);
19665 }
19666
19667 #[tokio::test]
19668 async fn web_search_without_provider_records_unavailable_without_execution() {
19669 let mock = mock_with_response("hello");
19670 let yaml = r#"
19671name: WebSearchNoProviderAgent
19672system_prompt: "You search the web."
19673tools: [web_search]
19674"#;
19675 let agent = AgentBuilder::from_yaml(yaml)
19676 .unwrap()
19677 .llm(Arc::new(mock))
19678 .auto_configure_features()
19679 .unwrap()
19680 .build()
19681 .unwrap();
19682
19683 let record = agent
19684 .invoke_tool(ToolExecutionRequest::new(
19685 "web-search-call",
19686 "web_search",
19687 serde_json::json!({"query": "rust async"}),
19688 ToolCallSource::Manual,
19689 ))
19690 .await
19691 .unwrap();
19692
19693 assert!(!record.executed);
19694 assert!(!record.success);
19695 assert_eq!(record.policy.outcome, PermissionOutcome::Unavailable);
19696 }
19697
19698 #[tokio::test]
19699 async fn unavailable_host_tool_does_not_request_approval() {
19700 let approvals = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19701 let handler = Arc::new(CountingApprovalHandler {
19702 calls: Arc::clone(&approvals),
19703 });
19704 let mut security = ToolSecurityConfig {
19705 enabled: true,
19706 fail_closed: true,
19707 ..Default::default()
19708 };
19709 security.tools.insert(
19710 "web_search".to_string(),
19711 ai_agents_tools::ToolPolicyConfig {
19712 enabled: true,
19713 require_confirmation: true,
19714 ..Default::default()
19715 },
19716 );
19717 let yaml = r#"
19718name: UnavailableApprovalAgent
19719system_prompt: "Search only with approval."
19720tools: [web_search]
19721"#;
19722 let agent = AgentBuilder::from_yaml(yaml)
19723 .unwrap()
19724 .llm(Arc::new(mock_with_response("done")))
19725 .auto_configure_features()
19726 .unwrap()
19727 .tool_security(ToolSecurityEngine::new(security))
19728 .hitl_engine(HITLEngine::new(ai_agents_hitl::HITLConfig::default()))
19729 .approval_handler(handler)
19730 .build()
19731 .unwrap();
19732
19733 let record = agent
19734 .invoke_tool(ToolExecutionRequest::new(
19735 "unavailable-before-approval",
19736 "web_search",
19737 serde_json::json!({"query": "rust async"}),
19738 ToolCallSource::Manual,
19739 ))
19740 .await
19741 .unwrap();
19742
19743 assert_eq!(approvals.load(Ordering::SeqCst), 0);
19744 assert!(!record.executed);
19745 assert!(!record.success);
19746 assert_eq!(record.policy.outcome, PermissionOutcome::Unavailable);
19747 assert!(
19748 record
19749 .approval
19750 .as_ref()
19751 .is_some_and(|approval| matches!(approval.status, ToolApprovalStatus::Unavailable))
19752 );
19753 }
19754
19755 #[tokio::test]
19756 async fn test_spawner_section_does_not_grant_core_tools_when_top_level_tools_omitted() {
19757 let mock = mock_with_response("hello");
19758 let yaml = r#"
19759name: SpawnerNoGrantAgent
19760system_prompt: "You manage agents."
19761spawner:
19762 max_agents: 2
19763"#;
19764 let agent = AgentBuilder::from_yaml(yaml)
19765 .unwrap()
19766 .llm(Arc::new(mock))
19767 .auto_configure_features()
19768 .unwrap()
19769 .auto_configure_spawner()
19770 .await
19771 .unwrap()
19772 .build()
19773 .unwrap();
19774
19775 let available = agent.get_available_tool_ids().await.unwrap();
19776 assert!(available.is_empty());
19777 }
19778
19779 #[tokio::test]
19780 async fn test_spawner_section_does_not_grant_core_tools_when_top_level_tools_empty() {
19781 let mock = mock_with_response("hello");
19782 let yaml = r#"
19783name: EmptySpawnerNoGrantAgent
19784system_prompt: "You manage agents."
19785tools: []
19786spawner:
19787 max_agents: 2
19788"#;
19789 let agent = AgentBuilder::from_yaml(yaml)
19790 .unwrap()
19791 .llm(Arc::new(mock))
19792 .auto_configure_features()
19793 .unwrap()
19794 .auto_configure_spawner()
19795 .await
19796 .unwrap()
19797 .build()
19798 .unwrap();
19799
19800 let available = agent.get_available_tool_ids().await.unwrap();
19801 assert!(available.is_empty());
19802 }
19803
19804 #[tokio::test]
19805 async fn test_management_tools_flag_grants_core_tools_when_top_level_tools_empty() {
19806 let mock = mock_with_response("hello");
19807 let yaml = r#"
19808name: ManagementGrantAgent
19809system_prompt: "You manage agents."
19810tools: []
19811spawner:
19812 management_tools: true
19813"#;
19814 let agent = AgentBuilder::from_yaml(yaml)
19815 .unwrap()
19816 .llm(Arc::new(mock))
19817 .auto_configure_features()
19818 .unwrap()
19819 .auto_configure_spawner()
19820 .await
19821 .unwrap()
19822 .build()
19823 .unwrap();
19824
19825 let available = agent.get_available_tool_ids().await.unwrap();
19826 assert_eq!(available.len(), 4);
19827 assert!(available.contains(&"spawn_agent".to_string()));
19828 assert!(available.contains(&"send_agent_message".to_string()));
19829 assert!(available.contains(&"list_agents".to_string()));
19830 assert!(available.contains(&"remove_agent".to_string()));
19831 }
19832
19833 #[tokio::test]
19834 async fn test_management_tools_flag_grants_core_tools_when_top_level_tools_omitted() {
19835 let mock = mock_with_response("hello");
19836 let yaml = r#"
19837name: ManagementOmittedToolsGrantAgent
19838system_prompt: "You manage agents."
19839spawner:
19840 management_tools: true
19841"#;
19842 let agent = AgentBuilder::from_yaml(yaml)
19843 .unwrap()
19844 .llm(Arc::new(mock))
19845 .auto_configure_features()
19846 .unwrap()
19847 .auto_configure_spawner()
19848 .await
19849 .unwrap()
19850 .build()
19851 .unwrap();
19852
19853 let available = agent.get_available_tool_ids().await.unwrap();
19854 assert_eq!(available.len(), 4);
19855 assert!(available.contains(&"spawn_agent".to_string()));
19856 assert!(available.contains(&"send_agent_message".to_string()));
19857 assert!(available.contains(&"list_agents".to_string()));
19858 assert!(available.contains(&"remove_agent".to_string()));
19859 }
19860
19861 #[tokio::test]
19862 async fn test_management_tools_selected_grants_only_selected_tools() {
19863 let mock = mock_with_response("hello");
19864 let yaml = r#"
19865name: ManagementSelectedGrantAgent
19866system_prompt: "You manage agents."
19867tools: []
19868spawner:
19869 management_tools:
19870 - spawn_agent
19871 - send_agent_message
19872 - list_agents
19873"#;
19874 let agent = AgentBuilder::from_yaml(yaml)
19875 .unwrap()
19876 .llm(Arc::new(mock))
19877 .auto_configure_features()
19878 .unwrap()
19879 .auto_configure_spawner()
19880 .await
19881 .unwrap()
19882 .build()
19883 .unwrap();
19884
19885 let available = agent.get_available_tool_ids().await.unwrap();
19886 assert_eq!(available.len(), 3);
19887 assert!(available.contains(&"spawn_agent".to_string()));
19888 assert!(available.contains(&"send_agent_message".to_string()));
19889 assert!(available.contains(&"list_agents".to_string()));
19890 assert!(!available.contains(&"remove_agent".to_string()));
19891 }
19892
19893 #[tokio::test]
19894 async fn test_orchestration_tools_flag_grants_tools_when_top_level_tools_empty() {
19895 let mock = mock_with_response("hello");
19896 let yaml = r#"
19897name: OrchestrationGrantAgent
19898system_prompt: "You coordinate agents."
19899llms:
19900 default:
19901 provider: openai
19902 model: gpt-4
19903 router:
19904 provider: openai
19905 model: gpt-4
19906llm:
19907 default: default
19908 router: router
19909tools: []
19910spawner:
19911 orchestration_tools: true
19912"#;
19913 let agent = AgentBuilder::from_yaml(yaml)
19914 .unwrap()
19915 .llm(Arc::new(mock))
19916 .auto_configure_features()
19917 .unwrap()
19918 .auto_configure_spawner()
19919 .await
19920 .unwrap()
19921 .build()
19922 .unwrap();
19923
19924 let available = agent.get_available_tool_ids().await.unwrap();
19925 assert_eq!(available.len(), 5);
19926 assert!(available.contains(&"route_to_agent".to_string()));
19927 assert!(available.contains(&"pipeline_process".to_string()));
19928 assert!(available.contains(&"concurrent_ask".to_string()));
19929 assert!(available.contains(&"group_discussion".to_string()));
19930 assert!(available.contains(&"handoff_conversation".to_string()));
19931 }
19932
19933 #[tokio::test]
19934 async fn test_persona_evolve_flag_grants_tool_when_top_level_tools_empty() {
19935 let mock = mock_with_response("hello");
19936 let yaml = r#"
19937name: PersonaGrantAgent
19938system_prompt: "You can evolve persona."
19939llm:
19940 provider: openai
19941 model: gpt-4
19942tools: []
19943persona:
19944 identity:
19945 name: "Guide"
19946 role: "Helper"
19947 evolution:
19948 enabled: true
19949 allow_llm_evolve: true
19950 mutable_fields:
19951 - traits.personality
19952"#;
19953 let agent = AgentBuilder::from_yaml(yaml)
19954 .unwrap()
19955 .llm(Arc::new(mock))
19956 .build()
19957 .unwrap();
19958
19959 let available = agent.get_available_tool_ids().await.unwrap();
19960 assert_eq!(available, vec!["persona_evolve".to_string()]);
19961 }
19962
19963 #[tokio::test]
19964 async fn test_persona_evolve_flag_grants_tool_when_top_level_tools_omitted() {
19965 let mock = mock_with_response("hello");
19966 let yaml = r#"
19967name: PersonaOmittedToolsGrantAgent
19968system_prompt: "You can evolve persona."
19969llm:
19970 provider: openai
19971 model: gpt-4
19972persona:
19973 identity:
19974 name: "Guide"
19975 role: "Helper"
19976 evolution:
19977 enabled: true
19978 allow_llm_evolve: true
19979 mutable_fields:
19980 - traits.personality
19981"#;
19982 let agent = AgentBuilder::from_yaml(yaml)
19983 .unwrap()
19984 .llm(Arc::new(mock))
19985 .build()
19986 .unwrap();
19987
19988 let available = agent.get_available_tool_ids().await.unwrap();
19989 assert_eq!(available, vec!["persona_evolve".to_string()]);
19990 }
19991
19992 #[tokio::test]
19993 async fn test_omitted_yaml_tools_exposes_no_tools() {
19994 let mock = mock_with_response("hello");
19995 let yaml = r#"
19996name: NoToolsAgent
19997system_prompt: "You are helpful."
19998"#;
19999 let agent = AgentBuilder::from_yaml(yaml)
20000 .unwrap()
20001 .llm(Arc::new(mock))
20002 .auto_configure_features()
20003 .unwrap()
20004 .build()
20005 .unwrap();
20006
20007 let available = agent.get_available_tool_ids().await.unwrap();
20008 assert!(available.is_empty());
20009 }
20010
20011 #[tokio::test]
20012 async fn runtime_scope_cannot_widen_omitted_or_empty_yaml_grants() {
20013 for tools in ["", "tools: []"] {
20014 let yaml = format!(
20015 r#"
20016name: RuntimeScopeNoGrantAgent
20017system_prompt: "No ordinary tools are granted."
20018{tools}
20019"#
20020 );
20021 let agent = AgentBuilder::from_yaml(&yaml)
20022 .unwrap()
20023 .llm(Arc::new(mock_with_response("done")))
20024 .auto_configure_features()
20025 .unwrap()
20026 .build()
20027 .unwrap();
20028
20029 agent
20030 .runtime_control()
20031 .set_tool_scope(vec!["calculator".to_string()]);
20032
20033 assert!(agent.get_available_tool_ids().await.unwrap().is_empty());
20034 }
20035 }
20036
20037 #[tokio::test]
20038 async fn runtime_scope_widening_attempt_keeps_only_declared_tools() {
20039 let yaml = r#"
20040name: RuntimeScopeWideningAgent
20041system_prompt: "Runtime scope cannot add authority."
20042tools: [calculator]
20043"#;
20044 let agent = AgentBuilder::from_yaml(yaml)
20045 .unwrap()
20046 .llm(Arc::new(mock_with_response("done")))
20047 .auto_configure_features()
20048 .unwrap()
20049 .build()
20050 .unwrap();
20051
20052 agent
20053 .runtime_control()
20054 .set_tool_scope(vec!["calculator".to_string(), "datetime".to_string()]);
20055
20056 assert_eq!(
20057 agent.get_available_tool_ids().await.unwrap(),
20058 vec!["calculator".to_string()]
20059 );
20060 }
20061
20062 #[tokio::test]
20063 async fn runtime_scope_is_canonical_unique_ordered_and_clear_restores_declared_grant() {
20064 let yaml = r#"
20065name: RuntimeScopeIntersectionAgent
20066system_prompt: "Use only declared tools."
20067tools: [calculator, datetime]
20068"#;
20069 let agent = AgentBuilder::from_yaml(yaml)
20070 .unwrap()
20071 .llm(Arc::new(mock_with_response("done")))
20072 .auto_configure_features()
20073 .unwrap()
20074 .build()
20075 .unwrap();
20076 let mut aliases = ai_agents_tools::ToolAliases::default();
20077 aliases
20078 .names
20079 .insert("en".to_string(), "calculate_alias".to_string());
20080 agent.tools.set_tool_aliases("calculator", aliases);
20081 let control = agent.runtime_control();
20082
20083 control.set_tool_scope(vec![
20084 "datetime".to_string(),
20085 "calculate_alias".to_string(),
20086 "calculator".to_string(),
20087 "unknown".to_string(),
20088 "datetime".to_string(),
20089 ]);
20090 assert_eq!(
20091 agent.get_available_tool_ids().await.unwrap(),
20092 vec!["calculator".to_string(), "datetime".to_string()]
20093 );
20094
20095 control.set_tool_scope(vec!["datetime".to_string()]);
20096 assert_eq!(
20097 agent.get_available_tool_ids().await.unwrap(),
20098 vec!["datetime".to_string()]
20099 );
20100
20101 control.clear_tool_scope_override();
20102 assert_eq!(
20103 agent.get_available_tool_ids().await.unwrap(),
20104 vec!["calculator".to_string(), "datetime".to_string()]
20105 );
20106 }
20107
20108 #[tokio::test]
20109 async fn runtime_scope_preserves_programmatic_registration_as_declared_grant() {
20110 let agent = AgentBuilder::new()
20111 .system_prompt("Use registered tools.")
20112 .llm(Arc::new(mock_with_response("done")))
20113 .tool(Arc::new(ContextEchoTool))
20114 .tool(Arc::new(SlowTool))
20115 .build()
20116 .unwrap();
20117
20118 agent.runtime_control().set_tool_scope(vec![
20119 "Context Echo".to_string(),
20120 "context_echo".to_string(),
20121 "unknown".to_string(),
20122 ]);
20123
20124 assert_eq!(
20125 agent.get_available_tool_ids().await.unwrap(),
20126 vec!["context_echo".to_string()]
20127 );
20128 }
20129
20130 #[tokio::test]
20131 async fn nested_state_scopes_intersect_every_ancestor_with_aliases() {
20132 let yaml = r#"
20133name: NestedStateScopeAgent
20134system_prompt: "Honor every state scope."
20135tools: [calculator, datetime, echo]
20136states:
20137 initial: root
20138 states:
20139 root:
20140 tools: [calculate_alias, datetime]
20141 initial: middle
20142 states:
20143 middle:
20144 initial: leaf
20145 states:
20146 leaf:
20147 tools: [datetime_alias, echo]
20148"#;
20149 let agent = AgentBuilder::from_yaml(yaml)
20150 .unwrap()
20151 .llm(Arc::new(mock_with_response("done")))
20152 .auto_configure_features()
20153 .unwrap()
20154 .build()
20155 .unwrap();
20156 let mut calculator_aliases = ai_agents_tools::ToolAliases::default();
20157 calculator_aliases
20158 .names
20159 .insert("en".to_string(), "calculate_alias".to_string());
20160 agent
20161 .tools
20162 .set_tool_aliases("calculator", calculator_aliases);
20163 let mut datetime_aliases = ai_agents_tools::ToolAliases::default();
20164 datetime_aliases
20165 .names
20166 .insert("en".to_string(), "datetime_alias".to_string());
20167 agent.tools.set_tool_aliases("datetime", datetime_aliases);
20168 agent.runtime_control().set_tool_scope(vec![
20169 "unknown".to_string(),
20170 "datetime_alias".to_string(),
20171 "calculate_alias".to_string(),
20172 "datetime".to_string(),
20173 ]);
20174
20175 assert_eq!(agent.current_state().as_deref(), Some("root.middle.leaf"));
20176 assert_eq!(
20177 agent.get_available_tool_ids().await.unwrap(),
20178 vec!["datetime".to_string()]
20179 );
20180 }
20181
20182 #[tokio::test]
20183 async fn ancestor_empty_state_scope_denies_omitted_descendants() {
20184 let yaml = r#"
20185name: NestedEmptyStateScopeAgent
20186system_prompt: "An empty ancestor scope denies all tools."
20187tools: [calculator]
20188states:
20189 initial: root
20190 states:
20191 root:
20192 tools: []
20193 initial: middle
20194 states:
20195 middle:
20196 initial: leaf
20197 states:
20198 leaf: {}
20199"#;
20200 let agent = AgentBuilder::from_yaml(yaml)
20201 .unwrap()
20202 .llm(Arc::new(mock_with_response("done")))
20203 .auto_configure_features()
20204 .unwrap()
20205 .build()
20206 .unwrap();
20207
20208 assert!(agent.get_available_tool_ids().await.unwrap().is_empty());
20209 }
20210
20211 #[tokio::test]
20212 async fn state_change_during_approval_invalidates_the_reviewed_authority() {
20213 let active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
20214 let max_active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
20215 let entered = Arc::new(tokio::sync::Barrier::new(2));
20216 let release = Arc::new(tokio::sync::Notify::new());
20217 let handler = Arc::new(BlockingApprovalHandler {
20218 entered: Arc::clone(&entered),
20219 release: Arc::clone(&release),
20220 result: ApprovalResult::Approved,
20221 });
20222 let yaml = r#"
20223name: ApprovalStateGenerationAgent
20224system_prompt: "State authority may change during approval."
20225tools: [locked_write]
20226states:
20227 initial: first
20228 states:
20229 first:
20230 tools: [locked_write]
20231 second:
20232 tools: [locked_write]
20233"#;
20234 let agent = Arc::new(
20235 AgentBuilder::from_yaml(yaml)
20236 .unwrap()
20237 .llm(Arc::new(mock_with_response("done")))
20238 .tool(Arc::new(LockedWriteTool {
20239 active: Arc::clone(&active),
20240 max_active: Arc::clone(&max_active),
20241 }))
20242 .tool_security(ToolSecurityEngine::new(approval_security_config(true)))
20243 .hitl_engine(HITLEngine::new(ai_agents_hitl::HITLConfig::default()))
20244 .approval_handler(handler)
20245 .build()
20246 .unwrap(),
20247 );
20248 let running = Arc::clone(&agent);
20249 let call = tokio::spawn(async move {
20250 running
20251 .invoke_tool(ToolExecutionRequest::new(
20252 "approval-state-generation",
20253 "locked_write",
20254 serde_json::json!({"path": "./state-generation.txt"}),
20255 ToolCallSource::Manual,
20256 ))
20257 .await
20258 .unwrap()
20259 });
20260
20261 entered.wait().await;
20262 agent.transition_to("second").await.unwrap();
20263 release.notify_one();
20264 let record = call.await.unwrap();
20265
20266 assert!(!record.executed);
20267 assert!(record.output.contains("Approval became stale"));
20268 assert_eq!(max_active.load(Ordering::SeqCst), 0);
20269 }
20270
20271 #[tokio::test]
20272 async fn state_change_while_waiting_for_resource_lock_fails_final_admission() {
20273 let holder_gate = PathMutationGate::new();
20274 let waiter_gate = PathMutationGate::new();
20275 let yaml = r#"
20276name: LockedStateGenerationAgent
20277system_prompt: "State authority must remain stable through admission."
20278tools: [state_lock_holder, state_lock_waiter]
20279states:
20280 initial: first
20281 states:
20282 first:
20283 tools: [state_lock_holder, state_lock_waiter]
20284 second:
20285 tools: [state_lock_holder, state_lock_waiter]
20286"#;
20287 let agent = Arc::new(
20288 AgentBuilder::from_yaml(yaml)
20289 .unwrap()
20290 .llm(Arc::new(mock_with_response("done")))
20291 .tool(Arc::new(BlockingPathMutationTool {
20292 id: "state_lock_holder",
20293 path_fields: vec![ai_agents_core::PathPolicyBinding::write("path")],
20294 gate: holder_gate.clone(),
20295 }))
20296 .tool(Arc::new(BlockingPathMutationTool {
20297 id: "state_lock_waiter",
20298 path_fields: vec![ai_agents_core::PathPolicyBinding::write("path")],
20299 gate: waiter_gate.clone(),
20300 }))
20301 .build()
20302 .unwrap(),
20303 );
20304 let holder_call = {
20305 let agent = Arc::clone(&agent);
20306 tokio::spawn(async move {
20307 agent
20308 .invoke_tool(ToolExecutionRequest::new(
20309 "state-lock-holder",
20310 "state_lock_holder",
20311 serde_json::json!({"path": "./shared-state-path.txt"}),
20312 ToolCallSource::Manual,
20313 ))
20314 .await
20315 .unwrap()
20316 })
20317 };
20318 holder_gate.wait_until_entered().await;
20319 let waiter_call = {
20320 let agent = Arc::clone(&agent);
20321 tokio::spawn(async move {
20322 agent
20323 .invoke_tool(ToolExecutionRequest::new(
20324 "state-lock-waiter",
20325 "state_lock_waiter",
20326 serde_json::json!({"path": "./shared-state-path.txt"}),
20327 ToolCallSource::Manual,
20328 ))
20329 .await
20330 .unwrap()
20331 })
20332 };
20333
20334 wait_for_resource_lock_strong_count(&agent.resource_locks, 2).await;
20335 agent.transition_to("second").await.unwrap();
20336 holder_gate.release();
20337 let holder_record = holder_call.await.unwrap();
20338 let waiter_record = waiter_call.await.unwrap();
20339
20340 assert!(holder_record.success);
20341 assert!(!waiter_record.executed);
20342 assert!(
20343 waiter_record
20344 .output
20345 .contains("state scope changed before admission")
20346 );
20347 assert!(!waiter_gate.entered.load(Ordering::SeqCst));
20348 }
20349
20350 #[tokio::test]
20351 async fn test_state_tools_cannot_widen_top_level_grant() {
20352 let mock = mock_with_response("hello");
20353 let yaml = r#"
20354name: NarrowToolsAgent
20355system_prompt: "You are helpful."
20356tools:
20357 - calculator
20358states:
20359 initial: current
20360 states:
20361 current:
20362 tools: [datetime]
20363"#;
20364 let agent = AgentBuilder::from_yaml(yaml)
20365 .unwrap()
20366 .llm(Arc::new(mock))
20367 .auto_configure_features()
20368 .unwrap()
20369 .build()
20370 .unwrap();
20371
20372 let available = agent.get_available_tool_ids().await.unwrap();
20373 assert!(available.is_empty());
20374 }
20375
20376 #[tokio::test]
20378 async fn test_integration_tool_execution() {
20379 let mock = mock_with_responses(vec![
20381 r#"I'll calculate that for you.
20383{"tool": "calculator", "arguments": {"expression": "2+2"}}"#,
20384 "The answer is 4.",
20386 ]);
20387 let observed = mock.clone();
20388 let mut tools = ai_agents_tools::ToolRegistry::new();
20389 tools
20390 .register(Arc::new(ai_agents_tools::CalculatorTool))
20391 .unwrap();
20392
20393 let agent = AgentBuilder::new()
20394 .system_prompt("You are a calculator assistant.")
20395 .llm(Arc::new(mock))
20396 .tools(tools)
20397 .build()
20398 .unwrap();
20399
20400 let response = agent.chat("What is 2+2?").await.unwrap();
20401
20402 assert_eq!(response.content, "The answer is 4.");
20403 assert_eq!(response.tool_calls.as_ref().map(Vec::len), Some(1));
20404 assert_eq!(
20405 observed.call_count(),
20406 2,
20407 "tool result must trigger a second LLM call"
20408 );
20409 let history = agent.tool_call_history();
20410 assert_eq!(history.len(), 1);
20411 assert_eq!(history[0].tool_id, "calculator");
20412 assert_eq!(
20413 history[0].result.get("result"),
20414 Some(&serde_json::json!(4.0)),
20415 "{:?}",
20416 history[0].result
20417 );
20418 }
20419
20420 #[test]
20423 fn legacy_tool_call_marker_is_plain_text() {
20424 let agent = AgentBuilder::new()
20425 .system_prompt("x")
20426 .llm(Arc::new(mock_with_response("x")))
20427 .build()
20428 .unwrap();
20429 let parsed = agent
20430 .parse_tool_calls(
20431 r#"[TOOL_CALL: {"name": "calculator", "arguments": {"expression": "2+2"}}]"#,
20432 )
20433 .unwrap();
20434 assert!(parsed.is_none());
20435 }
20436
20437 #[tokio::test]
20438 async fn test_tool_hitl_rejection_finalizes_blocking_turn() {
20439 let responses = Arc::new(std::sync::atomic::AtomicUsize::new(0));
20440 let hooks = Arc::new(ResponseCountingHooks {
20441 responses: Arc::clone(&responses),
20442 });
20443 let mock = mock_with_response(r#"{"tool":"echo","arguments":{"message":"hello"}}"#);
20444 let yaml = r#"
20445name: ToolRejectAgent
20446system_prompt: "You use tools when requested."
20447tools:
20448 - echo
20449hitl:
20450 tools:
20451 echo:
20452 require_approval: true
20453 approval_message: "Approve echo?"
20454"#;
20455 let agent = AgentBuilder::from_yaml(yaml)
20456 .unwrap()
20457 .llm(Arc::new(mock))
20458 .auto_configure_features()
20459 .unwrap()
20460 .hooks(hooks)
20461 .build()
20462 .unwrap();
20463
20464 let response = agent.chat("echo hello").await.unwrap();
20465
20466 assert!(
20467 response.content.contains("Operation cancelled"),
20468 "unexpected response: {}",
20469 response.content
20470 );
20471 assert_eq!(responses.load(Ordering::SeqCst), 1);
20472 let messages = agent.memory.get_messages(None).await.unwrap();
20473 assert_eq!(messages.len(), 3);
20474 assert_eq!(messages[0].content, "echo hello");
20475 assert!(messages[1].content.contains("\"tool\":\"echo\""));
20476 assert!(messages[2].content.contains("rejected by the approver"));
20477 }
20478
20479 #[tokio::test]
20480 async fn test_tool_hitl_rejection_finalizes_streaming_turn() {
20481 use futures::StreamExt;
20482
20483 let responses = Arc::new(std::sync::atomic::AtomicUsize::new(0));
20484 let hooks = Arc::new(ResponseCountingHooks {
20485 responses: Arc::clone(&responses),
20486 });
20487 let mock = mock_with_response(r#"{"tool":"echo","arguments":{"message":"hello"}}"#);
20488 let yaml = r#"
20489name: ToolRejectStreamingAgent
20490system_prompt: "You use tools when requested."
20491tools:
20492 - echo
20493streaming:
20494 enabled: true
20495hitl:
20496 tools:
20497 echo:
20498 require_approval: true
20499 approval_message: "Approve echo?"
20500"#;
20501 let agent = AgentBuilder::from_yaml(yaml)
20502 .unwrap()
20503 .llm(Arc::new(mock))
20504 .auto_configure_features()
20505 .unwrap()
20506 .hooks(hooks)
20507 .build()
20508 .unwrap();
20509
20510 let mut stream = agent.chat_stream("echo hello").await.unwrap();
20511 let mut terminal_error = String::new();
20512 let mut done = false;
20513 while let Some(chunk) = stream.next().await {
20514 match chunk {
20515 StreamChunk::Error { message } => terminal_error = message,
20516 StreamChunk::Done {} => {
20517 done = true;
20518 break;
20519 }
20520 _ => {}
20521 }
20522 }
20523
20524 assert!(done);
20525 assert!(
20526 terminal_error.contains("Operation cancelled"),
20527 "unexpected terminal error: {}",
20528 terminal_error
20529 );
20530 assert_eq!(responses.load(Ordering::SeqCst), 1);
20531 let messages = agent.memory.get_messages(None).await.unwrap();
20532 assert_eq!(messages.len(), 3);
20533 assert_eq!(messages[0].content, "echo hello");
20534 assert!(messages[1].content.contains("\"tool\":\"echo\""));
20535 assert!(messages[2].content.contains("rejected by the approver"));
20536 }
20537
20538 #[tokio::test]
20539 async fn tool_hitl_rejection_preserves_legacy_error_but_finalizes_event_stream() {
20540 let mock = mock_with_response(r#"{"tool":"echo","arguments":{"message":"hello"}}"#);
20541 let yaml = r#"
20542name: ToolRejectEventAgent
20543system_prompt: "You use tools when requested."
20544tools:
20545 - echo
20546streaming:
20547 enabled: true
20548hitl:
20549 tools:
20550 echo:
20551 require_approval: true
20552 approval_message: "Approve echo?"
20553"#;
20554 let agent = AgentBuilder::from_yaml(yaml)
20555 .unwrap()
20556 .llm(Arc::new(mock))
20557 .auto_configure_features()
20558 .unwrap()
20559 .build()
20560 .unwrap();
20561
20562 let mut stream = agent.chat_stream_events("echo hello").await.unwrap();
20563 let mut error_seen = false;
20564 let mut final_response = None;
20565 while let Some(event) = stream.next().await {
20566 match event {
20567 AgentStreamEvent::Chunk(StreamChunk::Error { .. }) => error_seen = true,
20568 AgentStreamEvent::Final(response) => final_response = Some(response),
20569 AgentStreamEvent::Chunk(_) => {}
20570 }
20571 }
20572
20573 assert!(!error_seen);
20574 assert!(
20575 final_response
20576 .is_some_and(|response| { response.content.contains("Operation cancelled") })
20577 );
20578 }
20579
20580 #[tokio::test]
20581 async fn test_pre_response_guard_transition_skips_old_state_llm() {
20582 let mock = mock_with_response("Billing state response");
20583 let call_counter = mock.clone();
20584 let yaml = r#"
20585name: OptimizedStateAgent
20586system_prompt: "You route before answering."
20587runtime:
20588 optimization:
20589 enabled: true
20590 pre_response_deterministic_transitions: true
20591states:
20592 initial: greeting
20593 states:
20594 greeting:
20595 prompt: "Old state prompt that should be skipped."
20596 transitions:
20597 - to: billing
20598 guard:
20599 context:
20600 topic:
20601 eq: billing
20602 timing: pre_response
20603 billing:
20604 prompt: "Answer from the billing state."
20605"#;
20606 let agent = AgentBuilder::from_yaml(yaml)
20607 .unwrap()
20608 .llm(Arc::new(mock))
20609 .build()
20610 .unwrap();
20611 agent
20612 .set_context("topic", serde_json::json!("billing"))
20613 .unwrap();
20614
20615 let response = agent.chat("I need billing help").await.unwrap();
20616
20617 assert_eq!(agent.current_state().as_deref(), Some("billing"));
20618 assert_eq!(response.content, "Billing state response");
20619 assert_eq!(call_counter.call_count(), 1);
20620 assert_eq!(agent.actor_facts().len(), 0);
20621 }
20622
20623 #[tokio::test]
20624 async fn test_set_context_supports_dotted_paths_for_pre_response_guards() {
20625 let mock = mock_with_response("Billing state response");
20626 let call_counter = mock.clone();
20627 let yaml = r#"
20628name: OptimizedStateAgent
20629system_prompt: "You route before answering."
20630runtime:
20631 optimization:
20632 enabled: true
20633 pre_response_deterministic_transitions: true
20634context:
20635 request:
20636 type: runtime
20637 default:
20638 topic: general
20639states:
20640 initial: greeting
20641 states:
20642 greeting:
20643 prompt: "Old state prompt that should be skipped."
20644 transitions:
20645 - to: billing
20646 guard:
20647 context:
20648 request.topic:
20649 eq: billing
20650 timing: pre_response
20651 billing:
20652 prompt: "Answer from the billing state."
20653"#;
20654 let agent = AgentBuilder::from_yaml(yaml)
20655 .unwrap()
20656 .llm(Arc::new(mock))
20657 .build()
20658 .unwrap();
20659 agent
20660 .set_context("request.topic", serde_json::json!("billing"))
20661 .unwrap();
20662
20663 let response = agent.chat("I need billing help").await.unwrap();
20664
20665 assert_eq!(agent.current_state().as_deref(), Some("billing"));
20666 assert_eq!(response.content, "Billing state response");
20667 assert_eq!(call_counter.call_count(), 1);
20668 assert_eq!(
20669 agent.get_context().get("request"),
20670 Some(&serde_json::json!({"topic": "billing"}))
20671 );
20672 }
20673
20674 #[tokio::test]
20675 async fn test_pre_response_rejection_does_not_commit_staged_context_or_user() {
20676 let mock = mock_with_response("billing");
20677 let yaml = r#"
20678name: OptimizedStateAgent
20679system_prompt: "You route before answering."
20680runtime:
20681 optimization:
20682 enabled: true
20683 pre_response_deterministic_transitions: true
20684hitl:
20685 states:
20686 billing:
20687 on_enter: require_approval
20688 approval_message: "Approve billing route?"
20689states:
20690 initial: greeting
20691 states:
20692 greeting:
20693 prompt: "Old state prompt."
20694 extract:
20695 - key: topic
20696 description: "Support topic"
20697 transitions:
20698 - to: billing
20699 guard:
20700 context:
20701 topic:
20702 eq: billing
20703 timing: pre_response
20704 run_extractors: true
20705 billing:
20706 prompt: "Billing state."
20707"#;
20708 let agent = AgentBuilder::from_yaml(yaml)
20709 .unwrap()
20710 .llm(Arc::new(mock))
20711 .build()
20712 .unwrap();
20713
20714 let response = agent
20715 .try_pre_response_transition("billing please")
20716 .await
20717 .unwrap();
20718
20719 assert!(response.is_none());
20720 assert_eq!(agent.current_state().as_deref(), Some("greeting"));
20721 assert!(!agent.get_context().contains_key("topic"));
20722 assert_eq!(agent.memory.get_messages(None).await.unwrap().len(), 0);
20723 }
20724
20725 #[tokio::test]
20726 async fn test_pre_response_extractor_commits_context_on_winning_path() {
20727 let mock = mock_with_responses(vec!["billing", "Billing response"]);
20728 let yaml = r#"
20729name: OptimizedStateAgent
20730system_prompt: "You route before answering."
20731runtime:
20732 optimization:
20733 enabled: true
20734 pre_response_deterministic_transitions: true
20735states:
20736 initial: greeting
20737 states:
20738 greeting:
20739 prompt: "Old state prompt."
20740 extract:
20741 - key: topic
20742 description: "Support topic"
20743 transitions:
20744 - to: billing
20745 guard:
20746 context:
20747 topic:
20748 eq: billing
20749 timing: pre_response
20750 run_extractors: true
20751 billing:
20752 prompt: "Billing state."
20753"#;
20754 let agent = AgentBuilder::from_yaml(yaml)
20755 .unwrap()
20756 .llm(Arc::new(mock))
20757 .build()
20758 .unwrap();
20759
20760 let response = agent.chat("billing please").await.unwrap();
20761
20762 assert_eq!(agent.current_state().as_deref(), Some("billing"));
20763 assert_eq!(response.content, "Billing response");
20764 assert_eq!(
20765 agent.get_context().get("topic"),
20766 Some(&serde_json::json!("billing"))
20767 );
20768 }
20769
20770 #[tokio::test]
20771 async fn test_pre_response_extractor_miss_does_not_mutate_context() {
20772 let mock = mock_with_response("__NONE__");
20773 let yaml = r#"
20774name: OptimizedStateAgent
20775system_prompt: "You route before answering."
20776runtime:
20777 optimization:
20778 enabled: true
20779 pre_response_deterministic_transitions: true
20780states:
20781 initial: greeting
20782 states:
20783 greeting:
20784 prompt: "Old state prompt."
20785 extract:
20786 - key: topic
20787 description: "Support topic"
20788 transitions:
20789 - to: billing
20790 guard:
20791 context:
20792 topic:
20793 eq: billing
20794 timing: pre_response
20795 run_extractors: true
20796 billing:
20797 prompt: "Billing state."
20798"#;
20799 let agent = AgentBuilder::from_yaml(yaml)
20800 .unwrap()
20801 .llm(Arc::new(mock))
20802 .build()
20803 .unwrap();
20804
20805 let response = agent.try_pre_response_transition("hello").await.unwrap();
20806
20807 assert!(response.is_none());
20808 assert_eq!(agent.current_state().as_deref(), Some("greeting"));
20809 assert!(!agent.get_context().contains_key("topic"));
20810 }
20811
20812 #[tokio::test]
20813 async fn test_default_guard_transition_stays_post_response() {
20814 let mock = mock_with_responses(vec!["Greeting response", "Billing response"]);
20815 let call_counter = mock.clone();
20816 let yaml = r#"
20817name: TimingAgent
20818system_prompt: "You route carefully."
20819runtime:
20820 optimization:
20821 enabled: true
20822 pre_response_deterministic_transitions: true
20823states:
20824 initial: greeting
20825 states:
20826 greeting:
20827 prompt: "Old state prompt."
20828 transitions:
20829 - to: billing
20830 guard:
20831 context:
20832 topic:
20833 eq: billing
20834 billing:
20835 prompt: "Billing state."
20836"#;
20837 let agent = AgentBuilder::from_yaml(yaml)
20838 .unwrap()
20839 .llm(Arc::new(mock))
20840 .build()
20841 .unwrap();
20842 agent
20843 .set_context("topic", serde_json::json!("billing"))
20844 .unwrap();
20845
20846 let response = agent.chat("billing please").await.unwrap();
20847
20848 assert_eq!(agent.current_state().as_deref(), Some("billing"));
20849 assert_eq!(response.content, "Billing response");
20850 assert_eq!(call_counter.call_count(), 2);
20851 }
20852
20853 #[tokio::test]
20854 async fn test_explicit_post_response_guard_transition_stays_post_response() {
20855 let mock = mock_with_responses(vec!["Greeting response", "Billing response"]);
20856 let call_counter = mock.clone();
20857 let yaml = r#"
20858name: TimingAgent
20859system_prompt: "You route carefully."
20860runtime:
20861 optimization:
20862 enabled: true
20863 pre_response_deterministic_transitions: true
20864states:
20865 initial: greeting
20866 states:
20867 greeting:
20868 prompt: "Old state prompt."
20869 transitions:
20870 - to: billing
20871 guard:
20872 context:
20873 topic:
20874 eq: billing
20875 timing: post_response
20876 billing:
20877 prompt: "Billing state."
20878"#;
20879 let agent = AgentBuilder::from_yaml(yaml)
20880 .unwrap()
20881 .llm(Arc::new(mock))
20882 .build()
20883 .unwrap();
20884 agent
20885 .set_context("topic", serde_json::json!("billing"))
20886 .unwrap();
20887
20888 let response = agent.chat("billing please").await.unwrap();
20889
20890 assert_eq!(agent.current_state().as_deref(), Some("billing"));
20891 assert_eq!(response.content, "Billing response");
20892 assert_eq!(call_counter.call_count(), 2);
20893 }
20894
20895 #[tokio::test]
20896 async fn test_pre_response_extractors_are_transition_scoped() {
20897 let mock = mock_with_responses(vec!["billing", "Billing response"]);
20898 let yaml = r#"
20899name: ScopedExtractorAgent
20900system_prompt: "You route carefully."
20901runtime:
20902 optimization:
20903 enabled: true
20904 pre_response_deterministic_transitions: true
20905states:
20906 initial: greeting
20907 states:
20908 greeting:
20909 prompt: "Old state prompt."
20910 extract:
20911 - key: topic
20912 description: "Support topic"
20913 transitions:
20914 - to: wrong
20915 guard:
20916 context:
20917 topic:
20918 eq: billing
20919 timing: pre_response
20920 - to: billing
20921 guard:
20922 context:
20923 topic:
20924 eq: billing
20925 timing: pre_response
20926 run_extractors: true
20927 wrong:
20928 prompt: "Wrong state."
20929 billing:
20930 prompt: "Billing state."
20931"#;
20932 let agent = AgentBuilder::from_yaml(yaml)
20933 .unwrap()
20934 .llm(Arc::new(mock))
20935 .build()
20936 .unwrap();
20937
20938 let response = agent.chat("billing please").await.unwrap();
20939
20940 assert_eq!(agent.current_state().as_deref(), Some("billing"));
20941 assert_eq!(response.content, "Billing response");
20942 }
20943
20944 #[tokio::test]
20945 async fn test_pre_response_resolved_intent_routes_early() {
20946 let mock = mock_with_response("Billing response");
20947 let yaml = r#"
20948name: IntentAgent
20949system_prompt: "You route carefully."
20950runtime:
20951 optimization:
20952 enabled: true
20953 pre_response_deterministic_transitions: true
20954states:
20955 initial: greeting
20956 states:
20957 greeting:
20958 prompt: "Old state prompt."
20959 transitions:
20960 - to: billing
20961 intent: billing
20962 timing: pre_response
20963 billing:
20964 prompt: "Billing state."
20965"#;
20966 let agent = AgentBuilder::from_yaml(yaml)
20967 .unwrap()
20968 .llm(Arc::new(mock))
20969 .build()
20970 .unwrap();
20971 agent
20972 .set_context("resolved_intent", serde_json::json!("billing"))
20973 .unwrap();
20974
20975 let response = agent
20976 .try_pre_response_transition("I need billing help")
20977 .await
20978 .unwrap()
20979 .unwrap();
20980
20981 assert_eq!(agent.current_state().as_deref(), Some("billing"));
20982 assert_eq!(response.content, "Billing response");
20983 }
20984
20985 #[tokio::test]
20986 async fn test_background_overflow_error_surfaces() {
20987 let mut config = RuntimeConfig::default();
20988 config.optimization.enabled = true;
20989 config.optimization.post_turn.max_background_tasks = 1;
20990 config.optimization.post_turn.on_background_overflow = BackgroundOverflowPolicy::Error;
20991 let policy = crate::optimization::MaintenanceTaskPolicy {
20992 mode: MaintenanceMode::Background,
20993 await_before_next_turn: AwaitBeforeNextTurn::Always,
20994 };
20995 let agent = AgentBuilder::new()
20996 .system_prompt("You are helpful.")
20997 .llm(Arc::new(mock_with_response("ok")))
20998 .build()
20999 .unwrap()
21000 .with_runtime_config(config);
21001 agent
21002 .background_maintenance
21003 .spawn(None, async { std::future::pending::<Result<()>>().await })
21004 .unwrap();
21005
21006 let result = agent
21007 .spawn_or_handle_background(None, async { Ok(()) }, "facts", &policy)
21008 .await;
21009
21010 assert!(result.is_err());
21011 }
21012
21013 #[tokio::test]
21014 async fn test_speculative_reasoning_low_cap_uses_serial_reasoning() {
21015 let default_mock = mock_with_response("Plain draft response");
21016 let router_mock = mock_with_response("cot");
21017 let router_counter = router_mock.clone();
21018 let yaml = r#"
21019name: ReasoningReservationAgent
21020system_prompt: "You answer plainly unless reasoning wins."
21021llm:
21022 default: default
21023 router: router
21024observability:
21025 enabled: true
21026 export:
21027 write_raw_events: true
21028reasoning:
21029 mode: auto
21030 judge_llm: router
21031runtime:
21032 optimization:
21033 enabled: true
21034 max_speculative_llm_calls_per_turn: 1
21035 speculative_reasoning_auto: true
21036 max_parallel_runtime_tasks: 2
21037"#;
21038 let agent = AgentBuilder::from_yaml(yaml)
21039 .unwrap()
21040 .llm_alias("default", Arc::new(default_mock))
21041 .llm_alias("router", Arc::new(router_mock))
21042 .build()
21043 .unwrap();
21044
21045 let response = agent.chat("hello").await.unwrap();
21046
21047 assert_eq!(response.content, "Plain draft response");
21048 assert_eq!(router_counter.call_count(), 1);
21049 let events = agent.observability().unwrap().raw_events();
21050 assert!(!events.iter().any(|event| {
21051 event.dimensions.get("commit_behavior") == Some(&"reasoning_decision".to_string())
21052 }));
21053 }
21054
21055 #[tokio::test]
21056 async fn test_forced_reasoning_skips_plain_speculative_draft() {
21057 let mock = mock_with_response("Reasoned response");
21058 let yaml = r#"
21059name: ForcedReasoningAgent
21060system_prompt: "You reason before answering."
21061observability:
21062 enabled: true
21063 export:
21064 write_raw_events: true
21065reasoning:
21066 mode: cot
21067runtime:
21068 optimization:
21069 enabled: true
21070 max_speculative_llm_calls_per_turn: 2
21071 speculative_state_transitions: true
21072 max_parallel_runtime_tasks: 2
21073states:
21074 initial: triage
21075 states:
21076 triage:
21077 prompt: "Answer from triage."
21078 transitions:
21079 - to: billing
21080 guard:
21081 context:
21082 route:
21083 eq: billing
21084 timing: parallel
21085 billing:
21086 prompt: "Billing state."
21087"#;
21088 let agent = AgentBuilder::from_yaml(yaml)
21089 .unwrap()
21090 .llm(Arc::new(mock))
21091 .build()
21092 .unwrap();
21093
21094 let response = agent.chat("hello").await.unwrap();
21095
21096 assert_eq!(response.content, "Reasoned response");
21097 let events = agent.observability().unwrap().raw_events();
21098 assert!(
21099 !events
21100 .iter()
21101 .any(|event| event.dimensions.contains_key("branch_status"))
21102 );
21103 }
21104
21105 #[tokio::test]
21106 async fn test_speculative_skill_low_cap_uses_serial_skill_route() {
21107 let default_mock = mock_with_response("Skill committed response");
21108 let router_mock = mock_with_response("helper");
21109 let router_counter = router_mock.clone();
21110 let yaml = r#"
21111name: SkillReservationAgent
21112system_prompt: "Use skills when they match."
21113llm:
21114 default: default
21115 router: router
21116observability:
21117 enabled: true
21118 export:
21119 write_raw_events: true
21120runtime:
21121 optimization:
21122 enabled: true
21123 max_speculative_llm_calls_per_turn: 1
21124 speculative_skill_routing: true
21125 max_parallel_runtime_tasks: 2
21126skills:
21127 - id: helper
21128 description: "Answer helper requests"
21129 trigger: "User asks for helper"
21130 steps:
21131 - prompt: "Answer the helper request: {{ user_input }}"
21132"#;
21133 let agent = AgentBuilder::from_yaml(yaml)
21134 .unwrap()
21135 .llm_alias("default", Arc::new(default_mock))
21136 .llm_alias("router", Arc::new(router_mock))
21137 .build()
21138 .unwrap();
21139
21140 let response = agent.chat("please use helper").await.unwrap();
21141
21142 assert_eq!(response.content, "Skill committed response");
21143 assert_eq!(router_counter.call_count(), 1);
21144 let events = agent.observability().unwrap().raw_events();
21145 assert!(
21146 !events
21147 .iter()
21148 .any(|event| event.dimensions.contains_key("branch_status"))
21149 );
21150 }
21151
21152 #[tokio::test]
21153 async fn test_parallel_transition_low_cap_allows_deterministic_route() {
21154 let mock = mock_with_response("unused");
21155 let call_counter = mock.clone();
21156 let yaml = r#"
21157name: ParallelTransitionLowCapAgent
21158system_prompt: "Route before stale responses when safe."
21159runtime:
21160 optimization:
21161 enabled: true
21162 max_speculative_llm_calls_per_turn: 1
21163 speculative_state_transitions: true
21164 max_parallel_runtime_tasks: 2
21165states:
21166 initial: triage
21167 states:
21168 triage:
21169 prompt: "Triage state."
21170 transitions:
21171 - to: billing
21172 guard:
21173 context:
21174 route:
21175 eq: billing
21176 timing: parallel
21177 billing:
21178 prompt: "Billing state."
21179"#;
21180 let agent = AgentBuilder::from_yaml(yaml)
21181 .unwrap()
21182 .llm(Arc::new(mock))
21183 .build()
21184 .unwrap();
21185 agent
21186 .set_context("route", serde_json::json!("billing"))
21187 .unwrap();
21188 agent.update_active_turn_context("billing help", HashMap::new());
21189 assert!(
21190 agent.reserve_active_speculative_llm_call(
21191 RuntimeOptimizationKind::ParallelStateTransition
21192 )
21193 );
21194
21195 let selection = agent
21196 .select_parallel_transition_candidate("billing help")
21197 .await
21198 .unwrap();
21199 agent.end_root_turn();
21200
21201 match selection {
21202 ParallelTransitionSelection::Candidate(candidate) => {
21203 assert_eq!(candidate.target(), "billing");
21204 }
21205 ParallelTransitionSelection::NoMatch => panic!("deterministic route did not match"),
21206 ParallelTransitionSelection::ReservationExhausted => {
21207 panic!("deterministic route consumed LLM budget")
21208 }
21209 }
21210 assert_eq!(call_counter.call_count(), 0);
21211 }
21212
21213 #[tokio::test]
21214 async fn speculative_transition_drops_loser_before_state_actions() {
21215 let lock = Arc::new(tokio::sync::Mutex::new(()));
21216 let first_started = Arc::new(tokio::sync::Notify::new());
21217 let first_dropped = Arc::new(AtomicBool::new(false));
21218 let committed_after_drop = Arc::new(AtomicBool::new(false));
21219 let default = Arc::new(FirstCallLockingProvider {
21220 lock,
21221 first_started: Arc::clone(&first_started),
21222 first_dropped: Arc::clone(&first_dropped),
21223 committed_after_drop: Arc::clone(&committed_after_drop),
21224 calls: AtomicU64::new(0),
21225 });
21226 let router = Arc::new(RoutingAfterProviderStart {
21227 provider_started: first_started,
21228 });
21229 let yaml = r#"
21230name: SpeculativeCancellationAgent
21231system_prompt: "Route before committed work."
21232llm:
21233 default: default
21234 router: router
21235runtime:
21236 optimization:
21237 enabled: true
21238 max_speculative_llm_calls_per_turn: 2
21239 speculative_state_transitions: true
21240 max_parallel_runtime_tasks: 2
21241states:
21242 initial: triage
21243 states:
21244 triage:
21245 prompt: "Triage state."
21246 transitions:
21247 - to: technical
21248 when: "The request needs technical support"
21249 timing: parallel
21250 technical:
21251 prompt: "Technical state."
21252 on_enter:
21253 - prompt: "Prepare technical context."
21254 llm: default
21255 store_as: preparation
21256"#;
21257 let agent = AgentBuilder::from_yaml(yaml)
21258 .unwrap()
21259 .llm_alias("default", default)
21260 .llm_alias("router", router)
21261 .build()
21262 .unwrap();
21263
21264 let response = tokio::time::timeout(
21265 std::time::Duration::from_secs(2),
21266 agent.chat("I cannot log in because of AUTH-17."),
21267 )
21268 .await
21269 .expect("committed work must not wait on the losing provider future")
21270 .unwrap();
21271
21272 assert_eq!(response.content, "Committed technical response.");
21273 assert_eq!(agent.current_state().as_deref(), Some("technical"));
21274 assert!(first_dropped.load(Ordering::SeqCst));
21275 assert!(committed_after_drop.load(Ordering::SeqCst));
21276 }
21277
21278 #[tokio::test]
21279 async fn buffered_transition_drops_stale_stream_before_redispatch() {
21280 use futures::StreamExt;
21281
21282 let lock = Arc::new(tokio::sync::Mutex::new(()));
21283 let stream_started = Arc::new(tokio::sync::Notify::new());
21284 let stream_dropped = Arc::new(AtomicBool::new(false));
21285 let committed_after_drop = Arc::new(AtomicBool::new(false));
21286 let default = Arc::new(BufferedLockingProvider {
21287 lock,
21288 stream_started: Arc::clone(&stream_started),
21289 stream_dropped: Arc::clone(&stream_dropped),
21290 committed_after_drop: Arc::clone(&committed_after_drop),
21291 });
21292 let router = Arc::new(RoutingAfterProviderStart {
21293 provider_started: stream_started,
21294 });
21295 let yaml = r#"
21296name: BufferedCancellationAgent
21297system_prompt: "Hide stale streamed output."
21298llm:
21299 default: default
21300 router: router
21301streaming:
21302 enabled: true
21303 buffer_size: 8
21304runtime:
21305 optimization:
21306 enabled: true
21307 max_speculative_llm_calls_per_turn: 2
21308 speculative_state_transitions: true
21309 streaming_policy: buffer_until_routing_done
21310 max_parallel_runtime_tasks: 2
21311states:
21312 initial: triage
21313 states:
21314 triage:
21315 prompt: "Triage state."
21316 transitions:
21317 - to: technical
21318 when: "The request needs technical support"
21319 timing: parallel
21320 technical:
21321 prompt: "Technical state."
21322"#;
21323 let agent = AgentBuilder::from_yaml(yaml)
21324 .unwrap()
21325 .llm_alias("default", default)
21326 .llm_alias("router", router)
21327 .build()
21328 .unwrap();
21329
21330 let content = tokio::time::timeout(std::time::Duration::from_secs(2), async {
21331 let mut stream = agent
21332 .chat_stream("AUTH-17 needs technical help.")
21333 .await
21334 .unwrap();
21335 let mut content = String::new();
21336 while let Some(chunk) = stream.next().await {
21337 match chunk {
21338 StreamChunk::Content { text } => content.push_str(&text),
21339 StreamChunk::Done {} => break,
21340 StreamChunk::Error { message } => panic!("unexpected stream error: {message}"),
21341 _ => {}
21342 }
21343 }
21344 content
21345 })
21346 .await
21347 .expect("redispatch must not wait on the stale streaming future");
21348
21349 assert_eq!(content, "Committed technical response.");
21350 assert_eq!(agent.current_state().as_deref(), Some("technical"));
21351 assert!(stream_dropped.load(Ordering::SeqCst));
21352 assert!(committed_after_drop.load(Ordering::SeqCst));
21353 }
21354
21355 #[tokio::test]
21356 async fn buffered_transition_drops_established_stream_before_redispatch() {
21357 use futures::StreamExt;
21358
21359 let stream_started = Arc::new(tokio::sync::Notify::new());
21360 let stream_dropped = Arc::new(AtomicBool::new(false));
21361 let stream_dropped_notify = Arc::new(tokio::sync::Notify::new());
21362 let committed_after_drop = Arc::new(AtomicBool::new(false));
21363 let default = Arc::new(EstablishedStreamProvider {
21364 stream_started: Arc::clone(&stream_started),
21365 stream_dropped: Arc::clone(&stream_dropped),
21366 stream_dropped_notify,
21367 committed_after_drop: Arc::clone(&committed_after_drop),
21368 });
21369 let router = Arc::new(RoutingAfterProviderStart {
21370 provider_started: stream_started,
21371 });
21372 let yaml = r#"
21373name: EstablishedStreamCancellationAgent
21374system_prompt: "Hide stale streamed output."
21375llm:
21376 default: default
21377 router: router
21378streaming:
21379 enabled: true
21380 buffer_size: 8
21381runtime:
21382 optimization:
21383 enabled: true
21384 max_speculative_llm_calls_per_turn: 2
21385 speculative_state_transitions: true
21386 streaming_policy: buffer_until_routing_done
21387 max_parallel_runtime_tasks: 2
21388states:
21389 initial: triage
21390 states:
21391 triage:
21392 prompt: "Triage state."
21393 transitions:
21394 - to: technical
21395 when: "The request needs technical support"
21396 timing: parallel
21397 technical:
21398 prompt: "Technical state."
21399"#;
21400 let agent = AgentBuilder::from_yaml(yaml)
21401 .unwrap()
21402 .llm_alias("default", default)
21403 .llm_alias("router", router)
21404 .build()
21405 .unwrap();
21406
21407 let content = tokio::time::timeout(std::time::Duration::from_secs(2), async {
21408 let mut stream = agent
21409 .chat_stream("AUTH-17 needs technical help.")
21410 .await
21411 .unwrap();
21412 let mut content = String::new();
21413 while let Some(chunk) = stream.next().await {
21414 match chunk {
21415 StreamChunk::Content { text } => content.push_str(&text),
21416 StreamChunk::Done {} => break,
21417 StreamChunk::Error { message } => panic!("unexpected stream error: {message}"),
21418 _ => {}
21419 }
21420 }
21421 content
21422 })
21423 .await
21424 .expect("redispatch must wait for the established stale stream to be dropped");
21425
21426 assert_eq!(content, "Committed technical response.");
21427 assert_eq!(agent.current_state().as_deref(), Some("technical"));
21428 assert!(stream_dropped.load(Ordering::SeqCst));
21429 assert!(committed_after_drop.load(Ordering::SeqCst));
21430 }
21431
21432 #[tokio::test]
21433 async fn test_buffered_streaming_transition_reservation_falls_back() {
21434 use futures::StreamExt;
21435
21436 let mock = mock_with_responses(vec![
21437 "Serial streaming response",
21438 "Serial streaming response",
21439 ]);
21440 let router_mock = mock_with_response("1");
21441 let router_counter = router_mock.clone();
21442 let yaml = r#"
21443name: BufferedReservationFallbackAgent
21444system_prompt: "Stream normally if speculative routing cannot be evaluated."
21445llm:
21446 default: default
21447 router: router
21448observability:
21449 enabled: true
21450 export:
21451 write_raw_events: true
21452streaming:
21453 enabled: true
21454 buffer_size: 8
21455runtime:
21456 optimization:
21457 enabled: true
21458 max_speculative_llm_calls_per_turn: 1
21459 speculative_state_transitions: true
21460 streaming_policy: buffer_until_routing_done
21461 max_parallel_runtime_tasks: 2
21462states:
21463 initial: triage
21464 states:
21465 triage:
21466 prompt: "Triage state."
21467 transitions:
21468 - to: billing
21469 guard:
21470 context:
21471 route:
21472 eq: billing
21473 when: "User asks about billing"
21474 timing: parallel
21475 billing:
21476 prompt: "Billing state."
21477"#;
21478 let agent = AgentBuilder::from_yaml(yaml)
21479 .unwrap()
21480 .llm_alias("default", Arc::new(mock))
21481 .llm_alias("router", Arc::new(router_mock))
21482 .build()
21483 .unwrap();
21484
21485 let mut stream = agent.chat_stream("hello").await.unwrap();
21486 let mut content = String::new();
21487 let mut error = None;
21488 while let Some(chunk) = stream.next().await {
21489 match chunk {
21490 StreamChunk::Content { text } => content.push_str(&text),
21491 StreamChunk::Error { message } => error = Some(message),
21492 StreamChunk::Done {} => break,
21493 _ => {}
21494 }
21495 }
21496
21497 assert_eq!(error, None);
21498 assert_eq!(content, "Serial streaming response");
21499 assert_eq!(router_counter.call_count(), 0);
21500 let events = agent.observability().unwrap().raw_events();
21501 assert!(events.iter().any(|event| {
21502 event.dimensions.get("branch_status") == Some(&"cancelled".to_string())
21503 && event.dimensions.get("commit_behavior")
21504 == Some(&"transition_decision".to_string())
21505 }));
21506 }
21507
21508 #[tokio::test]
21509 async fn test_blocking_error_cleanup_resets_root_turn_for_next_chat() {
21510 let mut mock = mock_with_response("Recovered response");
21511 mock.set_error("boom");
21512 let mut handle = mock.clone();
21513 let agent = AgentBuilder::new()
21514 .system_prompt("You are helpful.")
21515 .llm(Arc::new(mock))
21516 .build()
21517 .unwrap();
21518
21519 assert!(agent.chat("first").await.is_err());
21520 handle.clear_error();
21521 let response = agent.chat("second").await.unwrap();
21522
21523 assert_eq!(response.content, "Recovered response");
21524 let messages = agent.memory.get_messages(None).await.unwrap();
21525 let user_count = messages
21526 .iter()
21527 .filter(|message| message.role == ai_agents_core::Role::User)
21528 .count();
21529 assert_eq!(user_count, 2);
21530 }
21531
21532 #[tokio::test]
21533 async fn test_streaming_error_cleanup_resets_root_turn_for_next_chat() {
21534 use futures::StreamExt;
21535
21536 let mut mock = mock_with_response("Recovered response");
21537 mock.set_error("stream boom");
21538 let mut handle = mock.clone();
21539 let agent = AgentBuilder::new()
21540 .system_prompt("You are helpful.")
21541 .llm(Arc::new(mock))
21542 .build()
21543 .unwrap();
21544
21545 let mut stream = agent.chat_stream("first").await.unwrap();
21546 let mut saw_error = false;
21547 while let Some(chunk) = stream.next().await {
21548 if matches!(chunk, StreamChunk::Error { .. }) {
21549 saw_error = true;
21550 }
21551 }
21552 assert!(saw_error);
21553
21554 handle.clear_error();
21555 let response = agent.chat("second").await.unwrap();
21556
21557 assert_eq!(response.content, "Recovered response");
21558 let messages = agent.memory.get_messages(None).await.unwrap();
21559 let user_count = messages
21560 .iter()
21561 .filter(|message| message.role == ai_agents_core::Role::User)
21562 .count();
21563 assert_eq!(user_count, 2);
21564 }
21565
21566 #[tokio::test]
21567 async fn test_buffered_streaming_route_miss_releases_buffer_limit() {
21568 use futures::StreamExt;
21569
21570 let mut mock = mock_with_response("one two three");
21571 mock.set_latency(10);
21572 let yaml = r#"
21573name: BufferedMissAgent
21574system_prompt: "You stream safely."
21575llm:
21576 default: default
21577streaming:
21578 enabled: true
21579 buffer_size: 1
21580runtime:
21581 optimization:
21582 enabled: true
21583 max_speculative_llm_calls_per_turn: 2
21584 speculative_state_transitions: true
21585 streaming_policy: buffer_until_routing_done
21586 max_parallel_runtime_tasks: 2
21587states:
21588 initial: triage
21589 states:
21590 triage:
21591 prompt: "Answer from triage."
21592 transitions:
21593 - to: billing
21594 guard:
21595 context:
21596 route:
21597 eq: billing
21598 timing: parallel
21599 billing:
21600 prompt: "Billing state."
21601"#;
21602 let agent = AgentBuilder::from_yaml(yaml)
21603 .unwrap()
21604 .llm_alias("default", Arc::new(mock))
21605 .build()
21606 .unwrap();
21607
21608 let mut stream = agent.chat_stream("hello").await.unwrap();
21609 let mut content = String::new();
21610 let mut error = None;
21611 while let Some(chunk) = stream.next().await {
21612 match chunk {
21613 StreamChunk::Content { text } => content.push_str(&text),
21614 StreamChunk::Error { message } => error = Some(message),
21615 StreamChunk::Done {} => break,
21616 _ => {}
21617 }
21618 }
21619
21620 assert_eq!(error, None);
21621 assert_eq!(content, "one two three");
21622 }
21623
21624 #[tokio::test]
21625 async fn test_buffered_streaming_main_failure_finalizes_branch() {
21626 use futures::StreamExt;
21627
21628 let mock = mock_with_response("one two");
21629 let mut router_mock = mock_with_response("0");
21630 router_mock.set_latency(50);
21631 let yaml = r#"
21632name: BufferedFailureAgent
21633system_prompt: "You stream safely."
21634llm:
21635 default: default
21636 router: router
21637observability:
21638 enabled: true
21639 export:
21640 write_raw_events: true
21641streaming:
21642 enabled: true
21643 buffer_size: 1
21644runtime:
21645 optimization:
21646 enabled: true
21647 max_speculative_llm_calls_per_turn: 2
21648 speculative_state_transitions: true
21649 streaming_policy: buffer_until_routing_done
21650 max_parallel_runtime_tasks: 2
21651states:
21652 initial: triage
21653 states:
21654 triage:
21655 prompt: "Ask for the category."
21656 transitions:
21657 - to: billing
21658 when: "User asks about billing"
21659 timing: parallel
21660 billing:
21661 prompt: "Billing state."
21662"#;
21663 let agent = AgentBuilder::from_yaml(yaml)
21664 .unwrap()
21665 .llm_alias("default", Arc::new(mock))
21666 .llm_alias("router", Arc::new(router_mock))
21667 .build()
21668 .unwrap();
21669
21670 let mut stream = agent.chat_stream("hello").await.unwrap();
21671 let mut error = String::new();
21672 while let Some(chunk) = stream.next().await {
21673 if let StreamChunk::Error { message } = chunk {
21674 error = message;
21675 }
21676 }
21677
21678 assert!(
21679 error.contains("stream buffer filled"),
21680 "unexpected stream error: {}",
21681 error
21682 );
21683 let events = agent.observability().unwrap().raw_events();
21684 assert!(events.iter().any(|event| {
21685 event.dimensions.get("branch_status") == Some(&"failed".to_string())
21686 && event.dimensions.get("commit_behavior") == Some(&"final_response".to_string())
21687 && event.dimensions.get("optimization")
21688 == Some(&"buffered_streaming_routing".to_string())
21689 }));
21690 }
21691
21692 #[tokio::test]
21693 async fn test_streaming_preflight_does_not_emit_old_state_content() {
21694 use futures::StreamExt;
21695
21696 let mock = mock_with_response("Billing streamed response");
21697 let yaml = r#"
21698name: StreamingOptimizedAgent
21699system_prompt: "You route before streaming."
21700runtime:
21701 optimization:
21702 enabled: true
21703 pre_response_deterministic_transitions: true
21704streaming:
21705 enabled: true
21706states:
21707 initial: greeting
21708 states:
21709 greeting:
21710 prompt: "OLD_STATE_SENTINEL"
21711 transitions:
21712 - to: billing
21713 guard:
21714 context:
21715 topic:
21716 eq: billing
21717 timing: pre_response
21718 billing:
21719 prompt: "Billing state."
21720"#;
21721 let agent = AgentBuilder::from_yaml(yaml)
21722 .unwrap()
21723 .llm(Arc::new(mock))
21724 .build()
21725 .unwrap();
21726 agent
21727 .set_context("topic", serde_json::json!("billing"))
21728 .unwrap();
21729
21730 let mut stream = agent.chat_stream("billing please").await.unwrap();
21731 let mut content = String::new();
21732 while let Some(chunk) = stream.next().await {
21733 match chunk {
21734 StreamChunk::Content { text } => content.push_str(&text),
21735 StreamChunk::Error { message } => panic!("stream error: {}", message),
21736 StreamChunk::Done {} => break,
21737 _ => {}
21738 }
21739 }
21740
21741 assert_eq!(agent.current_state().as_deref(), Some("billing"));
21742 assert!(content.contains("Billing streamed response"));
21743 assert!(!content.contains("OLD_STATE_SENTINEL"));
21744 }
21745
21746 #[tokio::test]
21748 async fn test_integration_state_machine_basic() {
21749 let yaml = r#"
21750name: StateAgent
21751system_prompt: "You are a support agent."
21752states:
21753 initial: greeting
21754 states:
21755 greeting:
21756 prompt: "Welcome the user warmly."
21757 transitions:
21758 - to: support
21759 when: "User needs help"
21760 auto: true
21761 support:
21762 prompt: "Help solve the user's problem."
21763"#;
21764 let mock = mock_with_responses(vec![
21765 "Welcome! How can I help?", "1", "I'll help you with that.", ]);
21769 let builder = AgentBuilder::from_yaml(yaml).unwrap();
21770 let agent = builder.llm(Arc::new(mock)).build().unwrap();
21771
21772 assert_eq!(agent.current_state(), Some("greeting".to_string()));
21773 let _ = agent.chat("I need help").await.unwrap();
21774 }
21777
21778 #[tokio::test]
21780 async fn test_integration_state_on_enter_set_context() {
21781 let yaml = r#"
21782name: ActionAgent
21783system_prompt: "You are helpful."
21784states:
21785 initial: step1
21786 states:
21787 step1:
21788 prompt: "Step 1"
21789 on_exit:
21790 - set_context:
21791 step1_exited: true
21792 transitions:
21793 - to: step2
21794 when: "always"
21795 auto: true
21796 step2:
21797 prompt: "Step 2"
21798 on_enter:
21799 - set_context:
21800 step2_entered: true
21801"#;
21802 let mock = mock_with_responses(vec![
21804 "Processing step 1.",
21805 "0", ]);
21807 let builder = AgentBuilder::from_yaml(yaml).unwrap();
21808 let agent = builder.llm(Arc::new(mock)).build().unwrap();
21809
21810 assert_eq!(agent.current_state(), Some("step1".to_string()));
21811
21812 agent.transition_to("step2").await.unwrap();
21814
21815 assert_eq!(agent.current_state(), Some("step2".to_string()));
21816
21817 let ctx = agent.get_context();
21819 assert_eq!(ctx.get("step1_exited"), Some(&serde_json::json!(true)));
21820 assert_eq!(ctx.get("step2_entered"), Some(&serde_json::json!(true)));
21821 }
21822
21823 #[tokio::test]
21824 async fn state_action_tool_preserves_source_in_stored_record() {
21825 let yaml = r#"
21826name: StateActionToolAgent
21827system_prompt: "You are helpful."
21828tools:
21829 - context_echo
21830states:
21831 initial: idle
21832 states:
21833 idle:
21834 prompt: "Idle"
21835 active:
21836 prompt: "Active"
21837 on_enter:
21838 - set_context:
21839 action_started: true
21840 - tool: context_echo
21841 args: {}
21842"#;
21843 let agent = AgentBuilder::from_yaml(yaml)
21844 .unwrap()
21845 .llm(Arc::new(mock_with_response("unused")))
21846 .tool(Arc::new(ContextEchoTool))
21847 .build()
21848 .unwrap();
21849
21850 agent.transition_to("active").await.unwrap();
21851
21852 let record: ToolExecutionRecord = serde_json::from_value(
21853 agent
21854 .get_context()
21855 .get("last_tool_record")
21856 .cloned()
21857 .expect("successful state action must store its execution record"),
21858 )
21859 .unwrap();
21860 assert!(record.executed);
21861 assert!(record.success);
21862 assert_eq!(record.canonical_id, "context_echo");
21863 assert!(matches!(
21864 &record.source,
21865 ToolCallSource::StateAction {
21866 state: Some(state),
21867 action_index: 1,
21868 } if state == "active"
21869 ));
21870 }
21871
21872 #[tokio::test]
21873 async fn test_ordinary_transition_uses_on_enter_then_on_reenter() {
21874 let yaml = r#"
21875name: OrdinaryLifecycleAgent
21876system_prompt: "You are helpful."
21877states:
21878 initial: intake
21879 regenerate_on_transition: false
21880 states:
21881 intake:
21882 prompt: "Intake"
21883 transitions:
21884 - to: drafting
21885 guard:
21886 context:
21887 route:
21888 eq: drafting
21889 drafting:
21890 prompt: "Drafting"
21891 on_enter:
21892 - set_context:
21893 draft_version: 1
21894 on_reenter:
21895 - set_context:
21896 draft_version: 2
21897 transitions:
21898 - to: review
21899 guard:
21900 context:
21901 route:
21902 eq: review
21903 review:
21904 prompt: "Review"
21905 on_enter:
21906 - set_context:
21907 review_entry: first
21908 transitions:
21909 - to: drafting
21910 guard:
21911 context:
21912 route:
21913 eq: drafting
21914"#;
21915 let agent = AgentBuilder::from_yaml(yaml)
21916 .unwrap()
21917 .llm(Arc::new(mock_with_responses(vec![
21918 "Intake response",
21919 "Draft response",
21920 "Review response",
21921 ])))
21922 .build()
21923 .unwrap();
21924
21925 agent
21926 .set_context("route", serde_json::json!("drafting"))
21927 .unwrap();
21928 agent.chat("Start a draft").await.unwrap();
21929 assert_eq!(agent.current_state().as_deref(), Some("drafting"));
21930 assert_eq!(
21931 agent.get_context().get("draft_version"),
21932 Some(&serde_json::json!(1))
21933 );
21934
21935 agent
21936 .set_context("route", serde_json::json!("review"))
21937 .unwrap();
21938 agent.chat("Review this").await.unwrap();
21939 assert_eq!(agent.current_state().as_deref(), Some("review"));
21940 assert_eq!(
21941 agent.get_context().get("review_entry"),
21942 Some(&serde_json::json!("first"))
21943 );
21944
21945 agent
21946 .set_context("route", serde_json::json!("drafting"))
21947 .unwrap();
21948 agent.chat("Revise this").await.unwrap();
21949 assert_eq!(agent.current_state().as_deref(), Some("drafting"));
21950 assert_eq!(
21951 agent.get_context().get("draft_version"),
21952 Some(&serde_json::json!(2))
21953 );
21954 }
21955
21956 #[tokio::test]
21957 async fn test_manual_transition_uses_on_enter_then_on_reenter() {
21958 let yaml = r#"
21959name: ManualLifecycleAgent
21960system_prompt: "You are helpful."
21961states:
21962 initial: intake
21963 states:
21964 intake:
21965 prompt: "Intake"
21966 drafting:
21967 prompt: "Drafting"
21968 on_enter:
21969 - set_context:
21970 draft_version: 1
21971 on_reenter:
21972 - set_context:
21973 draft_version: 2
21974 review:
21975 prompt: "Review"
21976"#;
21977 let agent = AgentBuilder::from_yaml(yaml)
21978 .unwrap()
21979 .llm(Arc::new(mock_with_response("unused")))
21980 .build()
21981 .unwrap();
21982
21983 assert!(!agent.get_context().contains_key("draft_version"));
21984 agent.transition_to("drafting").await.unwrap();
21985 assert_eq!(agent.current_state().as_deref(), Some("drafting"));
21986 assert_eq!(
21987 agent.get_context().get("draft_version"),
21988 Some(&serde_json::json!(1))
21989 );
21990
21991 agent.transition_to("review").await.unwrap();
21992 agent.transition_to("drafting").await.unwrap();
21993 assert_eq!(agent.current_state().as_deref(), Some("drafting"));
21994 assert_eq!(
21995 agent.get_context().get("draft_version"),
21996 Some(&serde_json::json!(2))
21997 );
21998 }
21999
22000 #[tokio::test]
22001 async fn test_timeout_transition_uses_on_enter_then_on_reenter() {
22002 let yaml = r#"
22003name: TimeoutLifecycleAgent
22004system_prompt: "You are helpful."
22005states:
22006 initial: intake
22007 regenerate_on_transition: false
22008 states:
22009 intake:
22010 prompt: "Intake"
22011 max_turns: 1
22012 timeout_to: drafting
22013 drafting:
22014 prompt: "Drafting"
22015 max_turns: 1
22016 timeout_to: review
22017 on_enter:
22018 - set_context:
22019 draft_version: 1
22020 on_reenter:
22021 - set_context:
22022 draft_version: 2
22023 review:
22024 prompt: "Review"
22025 max_turns: 1
22026 timeout_to: drafting
22027 on_enter:
22028 - set_context:
22029 review_entry: first
22030"#;
22031 let agent = AgentBuilder::from_yaml(yaml)
22032 .unwrap()
22033 .llm(Arc::new(mock_with_responses(vec![
22034 "Intake",
22035 "First draft",
22036 "Review",
22037 "Revised draft",
22038 ])))
22039 .build()
22040 .unwrap();
22041
22042 agent.chat("First turn").await.unwrap();
22043 assert_eq!(agent.current_state().as_deref(), Some("intake"));
22044 assert!(!agent.get_context().contains_key("draft_version"));
22045
22046 agent.chat("Second turn").await.unwrap();
22047 assert_eq!(agent.current_state().as_deref(), Some("drafting"));
22048 assert_eq!(
22049 agent.get_context().get("draft_version"),
22050 Some(&serde_json::json!(1))
22051 );
22052
22053 agent.chat("Third turn").await.unwrap();
22054 assert_eq!(agent.current_state().as_deref(), Some("review"));
22055 assert_eq!(
22056 agent.get_context().get("review_entry"),
22057 Some(&serde_json::json!("first"))
22058 );
22059
22060 agent.chat("Fourth turn").await.unwrap();
22061 assert_eq!(agent.current_state().as_deref(), Some("drafting"));
22062 assert_eq!(
22063 agent.get_context().get("draft_version"),
22064 Some(&serde_json::json!(2))
22065 );
22066 }
22067
22068 #[tokio::test]
22070 async fn test_integration_process_normalize() {
22071 let yaml = r#"
22072name: ProcessAgent
22073system_prompt: "You are helpful."
22074process:
22075 input:
22076 - type: normalize
22077 config:
22078 trim: true
22079 collapse_whitespace: true
22080"#;
22081 let mock = mock_with_response("Got your message.");
22082 let builder = AgentBuilder::from_yaml(yaml).unwrap();
22083 let agent = builder.llm(Arc::new(mock.clone())).build().unwrap();
22084
22085 let _ = agent.chat(" hello world ").await.unwrap();
22086
22087 let history = mock.call_history();
22089 assert!(!history.is_empty());
22090 let last_call = history.last().unwrap();
22092 let user_msg = last_call
22093 .messages
22094 .iter()
22095 .find(|m| m.role == ai_agents_core::Role::User)
22096 .unwrap();
22097 assert_eq!(user_msg.content, "hello world");
22098 }
22099
22100 #[tokio::test]
22104 async fn test_integration_memory_compression() {
22105 let yaml = r#"
22106name: MemoryAgent
22107system_prompt: "You are helpful."
22108memory:
22109 type: compacting
22110 max_messages: 100
22111 compress_threshold: 5
22112 max_recent_messages: 3
22113 summarize_batch_size: 2
22114"#;
22115 let responses: Vec<&str> = (0..8).map(|_| "Response from assistant.").collect();
22117 let mock = mock_with_responses(responses);
22118 let builder = AgentBuilder::from_yaml(yaml).unwrap();
22119 let agent = builder.llm(Arc::new(mock)).build().unwrap();
22120
22121 for i in 0..6 {
22123 let _ = agent.chat(&format!("Message {}", i)).await.unwrap();
22124 }
22125
22126 let messages = agent.memory.get_messages(None).await.unwrap();
22129 assert!(messages.len() <= 12); }
22133
22134 #[tokio::test]
22136 async fn test_integration_multi_llm_registry() {
22137 let mut mock_default = MockLLMProvider::new("default");
22138 mock_default.set_response("Default LLM response.");
22139 let mut mock_router = MockLLMProvider::new("router");
22140 mock_router.set_response("Router response.");
22141
22142 let agent = AgentBuilder::new()
22143 .system_prompt("You are helpful.")
22144 .llm_alias("default", Arc::new(mock_default))
22145 .llm_alias("router", Arc::new(mock_router))
22146 .build()
22147 .unwrap();
22148
22149 let response = agent.chat("Hello").await.unwrap();
22150 assert_eq!(response.content, "Default LLM response.");
22151 }
22152
22153 #[tokio::test]
22155 async fn test_integration_agent_reset() {
22156 let mock = mock_with_responses(vec!["Hello!", "Hello again!"]);
22157 let agent = AgentBuilder::new()
22158 .system_prompt("You are helpful.")
22159 .llm(Arc::new(mock))
22160 .build()
22161 .unwrap();
22162
22163 let _ = agent.chat("Hi").await.unwrap();
22164 let messages = agent.memory.get_messages(None).await.unwrap();
22165 assert_eq!(messages.len(), 2); agent.reset().await.unwrap();
22168 let messages = agent.memory.get_messages(None).await.unwrap();
22169 assert_eq!(messages.len(), 0);
22170 }
22171
22172 #[tokio::test]
22174 async fn test_integration_process_validate_reject() {
22175 use ai_agents_process::{ProcessConfig, ProcessProcessor};
22176
22177 let validate_config = ai_agents_process::ValidateStage {
22178 id: Some("length_check".to_string()),
22179 condition: None,
22180 config: ai_agents_process::ValidateConfig {
22181 rules: vec![ai_agents_process::ValidationRule::MinLength {
22182 min_length: 10,
22183 on_fail: ai_agents_process::ValidationAction {
22184 action: ai_agents_process::ValidationActionType::Reject,
22185 message: None,
22186 },
22187 }],
22188 ..Default::default()
22189 },
22190 };
22191 let process_config = ProcessConfig {
22192 input: vec![ai_agents_process::ProcessStage::Validate(validate_config)],
22193 ..Default::default()
22194 };
22195 let processor = ProcessProcessor::new(process_config);
22196
22197 let mock = mock_with_response("Should not reach here.");
22198 let agent = AgentBuilder::new()
22199 .system_prompt("You are helpful.")
22200 .llm(Arc::new(mock))
22201 .process_processor(processor)
22202 .build()
22203 .unwrap();
22204
22205 let response = agent.chat("Hi").await.unwrap();
22206 assert!(
22208 response.content.contains("rejected")
22209 || response.content.contains("Input rejected")
22210 || response.content.contains("too short")
22211 || response.content.contains("Too short")
22212 || response.content.len() < 50, "Expected rejection response, got: {}",
22214 response.content
22215 );
22216 }
22217
22218 #[tokio::test]
22220 async fn test_llm_fallback_on_failure() {
22221 use ai_agents_recovery::{ErrorRecoveryConfig, LLMFailureAction, LLMRecoveryConfig};
22222
22223 let mut primary = MockLLMProvider::new("primary");
22224 primary.set_error("Primary LLM is unavailable");
22225
22226 let mut fallback = MockLLMProvider::new("fallback");
22227 fallback.set_response("Fallback response works!");
22228
22229 let agent = AgentBuilder::new()
22230 .system_prompt("You are helpful.")
22231 .llm_alias("default", Arc::new(primary))
22232 .llm_alias("backup", Arc::new(fallback))
22233 .recovery_manager(RecoveryManager::new(ErrorRecoveryConfig {
22234 llm: LLMRecoveryConfig {
22235 on_failure: LLMFailureAction::FallbackLlm {
22236 fallback_llm: "backup".to_string(),
22237 },
22238 ..Default::default()
22239 },
22240 ..Default::default()
22241 }))
22242 .build()
22243 .unwrap();
22244
22245 let response = agent.chat("Hello").await.unwrap();
22246 assert!(
22247 response.content.contains("Fallback response"),
22248 "Expected fallback response, got: {}",
22249 response.content
22250 );
22251 }
22252
22253 #[tokio::test]
22255 async fn test_llm_fallback_response_static_message() {
22256 use ai_agents_recovery::{ErrorRecoveryConfig, LLMFailureAction, LLMRecoveryConfig};
22257
22258 let mut primary = MockLLMProvider::new("primary");
22259 primary.set_error("Primary LLM is unavailable");
22260
22261 let agent = AgentBuilder::new()
22262 .system_prompt("You are helpful.")
22263 .llm(Arc::new(primary))
22264 .recovery_manager(RecoveryManager::new(ErrorRecoveryConfig {
22265 llm: LLMRecoveryConfig {
22266 on_failure: LLMFailureAction::FallbackResponse {
22267 message: "I am temporarily unavailable. Please try again later."
22268 .to_string(),
22269 },
22270 ..Default::default()
22271 },
22272 ..Default::default()
22273 }))
22274 .build()
22275 .unwrap();
22276
22277 let response = agent.chat("Hello").await.unwrap();
22278 assert!(
22279 response.content.contains("temporarily unavailable"),
22280 "Expected static fallback message, got: {}",
22281 response.content
22282 );
22283 }
22284
22285 #[tokio::test]
22288 async fn test_tool_failure_skip() {
22289 use ai_agents_recovery::{
22290 ErrorRecoveryConfig, ToolFailureAction, ToolRecoveryConfig, ToolRetryConfig,
22291 };
22292
22293 let mock = mock_with_responses(vec![
22294 r#"{"tool": "calculator", "arguments": {"expression": "not a number +"}}"#,
22295 "The calculation was skipped, but I can still help you.",
22296 ]);
22297 let observed = mock.clone();
22298 let mut tools = ai_agents_tools::ToolRegistry::new();
22299 tools
22300 .register(Arc::new(ai_agents_tools::CalculatorTool))
22301 .unwrap();
22302
22303 let agent = AgentBuilder::new()
22304 .system_prompt("You are helpful.")
22305 .llm(Arc::new(mock))
22306 .tools(tools)
22307 .recovery_manager(RecoveryManager::new(ErrorRecoveryConfig {
22308 tools: ToolRecoveryConfig {
22309 default: ToolRetryConfig {
22310 max_retries: 0,
22311 timeout_ms: None,
22312 on_failure: ToolFailureAction::Skip,
22313 },
22314 ..Default::default()
22315 },
22316 ..Default::default()
22317 }))
22318 .build()
22319 .unwrap();
22320
22321 let response = agent.chat("Compute this").await.unwrap();
22322
22323 assert_eq!(
22324 response.content,
22325 "The calculation was skipped, but I can still help you."
22326 );
22327 assert_eq!(observed.call_count(), 2);
22328 let history = agent.tool_call_history();
22330 assert_eq!(history.len(), 1);
22331 assert_eq!(history[0].tool_id, "calculator");
22332 assert_eq!(
22333 history[0].result.get("skipped"),
22334 Some(&serde_json::json!(true)),
22335 "{:?}",
22336 history[0].result
22337 );
22338 }
22339
22340 #[tokio::test]
22342 async fn test_unregistered_tool_call_records_unavailable_and_continues() {
22343 let mock = mock_with_responses(vec![
22344 r#"{"tool": "nonexistent_tool", "arguments": {}}"#,
22345 "The tool was unavailable, but I can still help you.",
22346 ]);
22347 let observed = mock.clone();
22348
22349 let agent = AgentBuilder::new()
22350 .system_prompt("You are helpful.")
22351 .llm(Arc::new(mock))
22352 .build()
22353 .unwrap();
22354
22355 let response = agent.chat("Use the nonexistent tool").await.unwrap();
22356
22357 assert_eq!(
22358 response.content,
22359 "The tool was unavailable, but I can still help you."
22360 );
22361 assert_eq!(observed.call_count(), 2);
22362 let history = agent.tool_call_history();
22363 assert_eq!(history.len(), 1);
22364 assert_eq!(history[0].tool_id, "nonexistent_tool");
22365 assert_eq!(
22366 history[0].result.pointer("/error/kind"),
22367 Some(&serde_json::json!("tool_unavailable")),
22368 "{:?}",
22369 history[0].result
22370 );
22371 }
22372
22373 fn fallback_llm_recovery(fallback_llm: &str) -> RecoveryManager {
22378 use ai_agents_recovery::{ErrorRecoveryConfig, LLMFailureAction, LLMRecoveryConfig};
22379 RecoveryManager::new(ErrorRecoveryConfig {
22380 llm: LLMRecoveryConfig {
22381 on_failure: LLMFailureAction::FallbackLlm {
22382 fallback_llm: fallback_llm.to_string(),
22383 },
22384 ..Default::default()
22385 },
22386 ..Default::default()
22387 })
22388 }
22389
22390 #[tokio::test]
22391 async fn test_stream_llm_fallback_on_open_failure() {
22392 let mut primary = MockLLMProvider::new("primary");
22393 primary.set_error("Primary LLM is unavailable");
22394 let mut fallback = MockLLMProvider::new("fallback");
22395 fallback.set_response("Fallback response works!");
22396
22397 let agent = AgentBuilder::new()
22398 .system_prompt("You are helpful.")
22399 .llm_alias("default", Arc::new(primary))
22400 .llm_alias("backup", Arc::new(fallback))
22401 .recovery_manager(fallback_llm_recovery("backup"))
22402 .build()
22403 .unwrap();
22404
22405 let (content, chunks, final_response) = collect_stream_events(&agent, "Hello").await;
22406 assert!(
22407 !chunks.iter().any(StreamChunk::is_error),
22408 "fallback must not surface as a stream error: {chunks:?}"
22409 );
22410 let final_response = final_response.expect("Final must be emitted after fallback");
22411 assert!(content.contains("Fallback response"));
22412 assert!(final_response.content.contains("Fallback response"));
22413 }
22414
22415 #[tokio::test]
22416 async fn test_stream_llm_fallback_response_static_message() {
22417 use ai_agents_recovery::{ErrorRecoveryConfig, LLMFailureAction, LLMRecoveryConfig};
22418
22419 let mut primary = MockLLMProvider::new("primary");
22420 primary.set_error("Primary LLM is unavailable");
22421
22422 let agent = AgentBuilder::new()
22423 .system_prompt("You are helpful.")
22424 .llm(Arc::new(primary))
22425 .recovery_manager(RecoveryManager::new(ErrorRecoveryConfig {
22426 llm: LLMRecoveryConfig {
22427 on_failure: LLMFailureAction::FallbackResponse {
22428 message: "Service is temporarily unavailable.".to_string(),
22429 },
22430 ..Default::default()
22431 },
22432 ..Default::default()
22433 }))
22434 .build()
22435 .unwrap();
22436
22437 let (content, chunks, final_response) = collect_stream_events(&agent, "Hello").await;
22438 assert!(!chunks.iter().any(StreamChunk::is_error));
22439 let content_chunks = chunks.iter().filter(|c| c.is_content()).count();
22440 assert_eq!(content_chunks, 1, "static fallback is one content chunk");
22441 assert_eq!(content, "Service is temporarily unavailable.");
22442 assert_eq!(
22443 final_response.expect("Final").content,
22444 "Service is temporarily unavailable."
22445 );
22446 }
22447
22448 struct FailOnceStreamProvider {
22450 remaining_failures: Arc<std::sync::atomic::AtomicUsize>,
22451 open_attempts: Arc<std::sync::atomic::AtomicUsize>,
22452 }
22453
22454 #[async_trait]
22455 impl LLMProvider for FailOnceStreamProvider {
22456 async fn complete(
22457 &self,
22458 _messages: &[ChatMessage],
22459 _config: Option<&LLMConfig>,
22460 ) -> std::result::Result<LLMResponse, LLMError> {
22461 Ok(LLMResponse::new("blocking path", FinishReason::Stop))
22462 }
22463
22464 async fn complete_stream(
22465 &self,
22466 _messages: &[ChatMessage],
22467 _config: Option<&LLMConfig>,
22468 ) -> std::result::Result<
22469 Box<dyn Stream<Item = std::result::Result<LLMChunk, LLMError>> + Unpin + Send>,
22470 LLMError,
22471 > {
22472 self.open_attempts.fetch_add(1, Ordering::SeqCst);
22473 if self
22474 .remaining_failures
22475 .fetch_update(Ordering::SeqCst, Ordering::SeqCst, |n| n.checked_sub(1))
22476 .is_ok()
22477 {
22478 return Err(LLMError::Network("connection reset".to_string()));
22479 }
22480 Ok(Box::new(futures::stream::iter(vec![Ok(LLMChunk::new(
22481 "Recovered after retry",
22482 true,
22483 ))])))
22484 }
22485
22486 fn provider_name(&self) -> &str {
22487 "fail-once-stream"
22488 }
22489
22490 fn supports(&self, feature: LLMFeature) -> bool {
22491 matches!(feature, LLMFeature::Streaming)
22492 }
22493 }
22494
22495 #[tokio::test]
22496 async fn test_stream_llm_retry_then_success() {
22497 use ai_agents_recovery::{BackoffConfig, ErrorRecoveryConfig, RetryConfig};
22498
22499 let open_attempts = Arc::new(std::sync::atomic::AtomicUsize::new(0));
22500 let provider = FailOnceStreamProvider {
22501 remaining_failures: Arc::new(std::sync::atomic::AtomicUsize::new(1)),
22502 open_attempts: Arc::clone(&open_attempts),
22503 };
22504
22505 let agent = AgentBuilder::new()
22506 .system_prompt("You are helpful.")
22507 .llm(Arc::new(provider))
22508 .recovery_manager(RecoveryManager::new(ErrorRecoveryConfig {
22509 default: RetryConfig {
22510 max_retries: 1,
22511 backoff: BackoffConfig {
22512 initial_ms: 1,
22513 max_ms: 1,
22514 ..Default::default()
22515 },
22516 ..Default::default()
22517 },
22518 ..Default::default()
22519 }))
22520 .build()
22521 .unwrap();
22522
22523 let (content, chunks, final_response) = collect_stream_events(&agent, "Hello").await;
22524 assert!(!chunks.iter().any(StreamChunk::is_error), "{chunks:?}");
22525 assert_eq!(open_attempts.load(Ordering::SeqCst), 2);
22526 assert_eq!(content, "Recovered after retry");
22527 assert_eq!(
22528 final_response.expect("Final").content,
22529 "Recovered after retry"
22530 );
22531 }
22532
22533 #[tokio::test]
22534 async fn test_stream_llm_error_action_error_emits_terminal_error() {
22535 let mut primary = MockLLMProvider::new("primary");
22536 primary.set_error("Primary LLM is unavailable");
22537
22538 let agent = AgentBuilder::new()
22539 .system_prompt("You are helpful.")
22540 .llm(Arc::new(primary))
22541 .build()
22542 .unwrap();
22543
22544 let (_, chunks, final_response) = collect_stream_events(&agent, "Hello").await;
22545 assert!(
22546 final_response.is_none(),
22547 "default Error action must not produce Final"
22548 );
22549 assert!(
22550 chunks.iter().any(StreamChunk::is_error),
22551 "default Error action must surface a stream error"
22552 );
22553 }
22554
22555 struct MidStreamFailureProvider;
22557
22558 #[async_trait]
22559 impl LLMProvider for MidStreamFailureProvider {
22560 async fn complete(
22561 &self,
22562 _messages: &[ChatMessage],
22563 _config: Option<&LLMConfig>,
22564 ) -> std::result::Result<LLMResponse, LLMError> {
22565 Ok(LLMResponse::new("blocking path", FinishReason::Stop))
22566 }
22567
22568 async fn complete_stream(
22569 &self,
22570 _messages: &[ChatMessage],
22571 _config: Option<&LLMConfig>,
22572 ) -> std::result::Result<
22573 Box<dyn Stream<Item = std::result::Result<LLMChunk, LLMError>> + Unpin + Send>,
22574 LLMError,
22575 > {
22576 Ok(Box::new(futures::stream::iter(vec![
22577 Ok(LLMChunk::new("Partial ", false)),
22578 Err(LLMError::Network("connection dropped".to_string())),
22579 ])))
22580 }
22581
22582 fn provider_name(&self) -> &str {
22583 "mid-stream-failure"
22584 }
22585
22586 fn supports(&self, feature: LLMFeature) -> bool {
22587 matches!(feature, LLMFeature::Streaming)
22588 }
22589 }
22590
22591 #[tokio::test]
22592 async fn test_stream_mid_stream_failure_is_terminal() {
22593 let mut fallback = MockLLMProvider::new("fallback");
22594 fallback.set_response("Fallback must not run");
22595 let fallback_calls = fallback.clone();
22596
22597 let agent = AgentBuilder::new()
22598 .system_prompt("You are helpful.")
22599 .llm_alias("default", Arc::new(MidStreamFailureProvider))
22600 .llm_alias("backup", Arc::new(fallback))
22601 .recovery_manager(fallback_llm_recovery("backup"))
22602 .build()
22603 .unwrap();
22604
22605 let (content, chunks, final_response) = collect_stream_events(&agent, "Hello").await;
22606 assert_eq!(content, "Partial ");
22607 assert!(chunks.iter().any(StreamChunk::is_error));
22608 assert!(final_response.is_none());
22609 assert_eq!(
22610 fallback_calls.call_count(),
22611 0,
22612 "fallback must not run after a visible delta"
22613 );
22614 }
22615
22616 #[tokio::test]
22617 async fn test_buffered_streaming_draft_uses_fallback_llm() {
22618 use futures::StreamExt;
22619
22620 let mut primary = MockLLMProvider::new("primary");
22621 primary.set_error("Primary LLM is unavailable");
22622 let fallback = mock_with_response("fallback one two");
22623 let yaml = r#"
22624name: BufferedFallbackAgent
22625system_prompt: "You stream safely."
22626llm:
22627 default: default
22628streaming:
22629 enabled: true
22630 buffer_size: 8
22631runtime:
22632 optimization:
22633 enabled: true
22634 max_speculative_llm_calls_per_turn: 2
22635 speculative_state_transitions: true
22636 streaming_policy: buffer_until_routing_done
22637 max_parallel_runtime_tasks: 2
22638states:
22639 initial: triage
22640 states:
22641 triage:
22642 prompt: "Answer from triage."
22643 transitions:
22644 - to: billing
22645 guard:
22646 context:
22647 route:
22648 eq: billing
22649 timing: parallel
22650 billing:
22651 prompt: "Billing state."
22652"#;
22653 let agent = AgentBuilder::from_yaml(yaml)
22654 .unwrap()
22655 .llm_alias("default", Arc::new(primary))
22656 .llm_alias("backup", Arc::new(fallback))
22657 .recovery_manager(fallback_llm_recovery("backup"))
22658 .build()
22659 .unwrap();
22660
22661 let mut stream = agent.chat_stream("hello").await.unwrap();
22662 let mut content = String::new();
22663 let mut error = None;
22664 while let Some(chunk) = stream.next().await {
22665 match chunk {
22666 StreamChunk::Content { text } => content.push_str(&text),
22667 StreamChunk::Error { message } => error = Some(message),
22668 StreamChunk::Done {} => break,
22669 _ => {}
22670 }
22671 }
22672
22673 assert_eq!(error, None);
22674 assert_eq!(content, "fallback one two");
22675 }
22676
22677 #[tokio::test]
22678 async fn parity_llm_fallback_llm() {
22679 let build = || {
22680 let mut primary = MockLLMProvider::new("primary");
22681 primary.set_error("Primary LLM is unavailable");
22682 let mut fallback = MockLLMProvider::new("fallback");
22683 fallback.set_response("Fallback response works!");
22684 AgentBuilder::new()
22685 .system_prompt("You are helpful.")
22686 .llm_alias("default", Arc::new(primary))
22687 .llm_alias("backup", Arc::new(fallback))
22688 .recovery_manager(fallback_llm_recovery("backup"))
22689 .build()
22690 .unwrap()
22691 };
22692 assert_blocking_streaming_parity(build, "Hello").await;
22693 }
22694
22695 #[tokio::test]
22696 async fn parity_llm_fallback_response() {
22697 use ai_agents_recovery::{ErrorRecoveryConfig, LLMFailureAction, LLMRecoveryConfig};
22698 let build = || {
22699 let mut primary = MockLLMProvider::new("primary");
22700 primary.set_error("Primary LLM is unavailable");
22701 AgentBuilder::new()
22702 .system_prompt("You are helpful.")
22703 .llm(Arc::new(primary))
22704 .recovery_manager(RecoveryManager::new(ErrorRecoveryConfig {
22705 llm: LLMRecoveryConfig {
22706 on_failure: LLMFailureAction::FallbackResponse {
22707 message: "Service is temporarily unavailable.".to_string(),
22708 },
22709 ..Default::default()
22710 },
22711 ..Default::default()
22712 }))
22713 .build()
22714 .unwrap()
22715 };
22716 assert_blocking_streaming_parity(build, "Hello").await;
22717 }
22718
22719 #[tokio::test]
22720 async fn parity_basic_chat() {
22721 let build = || {
22722 AgentBuilder::new()
22723 .system_prompt("You are helpful.")
22724 .llm(Arc::new(mock_with_response("Plain answer")))
22725 .build()
22726 .unwrap()
22727 };
22728 assert_blocking_streaming_parity(build, "Hello").await;
22729 }
22730
22731 fn skills_with_parallel_transition_yaml(extra_optimization: &str, streaming: &str) -> String {
22737 format!(
22738 r#"
22739name: SkillsBesideTransitionAgent
22740system_prompt: "Use skills when they match."
22741llm:
22742 default: default
22743 router: router
22744observability:
22745 enabled: true
22746 export:
22747 write_raw_events: true
22748{streaming}
22749runtime:
22750 optimization:
22751 enabled: true
22752 speculative_state_transitions: true
22753{extra_optimization}
22754states:
22755 initial: triage
22756 states:
22757 triage:
22758 prompt: "Triage state."
22759 transitions:
22760 - to: billing
22761 guard:
22762 context:
22763 route:
22764 eq: billing
22765 timing: parallel
22766 billing:
22767 prompt: "Billing state."
22768skills:
22769 - id: helper
22770 description: "Answer helper requests"
22771 trigger: "User asks for helper"
22772 steps:
22773 - prompt: "Answer the helper request: {{{{ user_input }}}}"
22774 llm: skill
22775"#
22776 )
22777 }
22778
22779 struct RoleMocks {
22783 main: MockLLMProvider,
22784 router: MockLLMProvider,
22785 skill: MockLLMProvider,
22786 }
22787
22788 fn role_mocks(main: MockLLMProvider, router: MockLLMProvider) -> RoleMocks {
22789 RoleMocks {
22790 main,
22791 router,
22792 skill: mock_with_response("Skill step response"),
22793 }
22794 }
22795
22796 fn build_skills_beside_transition_agent(yaml: &str, mocks: RoleMocks) -> RuntimeAgent {
22797 AgentBuilder::from_yaml(yaml)
22798 .unwrap()
22799 .llm_alias("default", Arc::new(mocks.main))
22800 .llm_alias("router", Arc::new(mocks.router))
22801 .llm_alias("skill", Arc::new(mocks.skill))
22802 .build()
22803 .unwrap()
22804 }
22805
22806 fn branch_events_with_commit_behavior(agent: &RuntimeAgent, behavior: &str) -> usize {
22807 agent
22808 .observability()
22809 .unwrap()
22810 .raw_events()
22811 .iter()
22812 .filter(|event| event.dimensions.get("commit_behavior") == Some(&behavior.to_string()))
22813 .count()
22814 }
22815
22816 #[tokio::test]
22817 async fn test_speculative_transition_with_skills_and_no_skill_branch_routes_skill_serially() {
22818 let default_mock = mock_with_response("Draft response");
22819 let router_mock = mock_with_response("helper");
22820 let router_counter = router_mock.clone();
22821 let yaml = skills_with_parallel_transition_yaml(
22822 " max_speculative_llm_calls_per_turn: 2\n max_parallel_runtime_tasks: 2",
22823 "",
22824 );
22825 let agent =
22826 build_skills_beside_transition_agent(&yaml, role_mocks(default_mock, router_mock));
22827
22828 let response = agent.chat("please use helper").await.unwrap();
22829
22830 assert_eq!(
22831 response.metadata.as_ref().and_then(|m| m.get("skill_id")),
22832 Some(&serde_json::json!("helper")),
22833 "skill must route even without a skill branch: {response:?}"
22834 );
22835 assert_eq!(router_counter.call_count(), 1);
22836 assert!(branch_events_with_commit_behavior(&agent, "transition_decision") > 0);
22838 assert_eq!(
22839 branch_events_with_commit_behavior(&agent, "skill_selection"),
22840 0
22841 );
22842 }
22843
22844 #[tokio::test]
22845 async fn test_speculative_transition_with_skills_no_match_commits_draft() {
22846 let default_mock = mock_with_response("Draft response");
22847 let router_mock = mock_with_response("none");
22848 let router_counter = router_mock.clone();
22849 let yaml = skills_with_parallel_transition_yaml(
22850 " max_speculative_llm_calls_per_turn: 2\n max_parallel_runtime_tasks: 2",
22851 "",
22852 );
22853 let agent =
22854 build_skills_beside_transition_agent(&yaml, role_mocks(default_mock, router_mock));
22855
22856 let response = agent.chat("just chat").await.unwrap();
22857
22858 assert_eq!(response.content, "Draft response");
22859 assert!(
22860 response
22861 .metadata
22862 .as_ref()
22863 .is_none_or(|m| !m.contains_key("skill_id"))
22864 );
22865 assert_eq!(router_counter.call_count(), 1);
22866 assert!(branch_events_with_commit_behavior(&agent, "final_response") > 0);
22867 }
22868
22869 #[tokio::test]
22870 async fn test_speculative_transition_win_skips_serial_skill_selection() {
22871 let default_mock = mock_with_response("Billing answer");
22872 let router_mock = mock_with_response("none");
22873 let router_counter = router_mock.clone();
22874 let yaml = skills_with_parallel_transition_yaml(
22875 " max_speculative_llm_calls_per_turn: 2\n max_parallel_runtime_tasks: 2",
22876 "",
22877 );
22878 let agent =
22879 build_skills_beside_transition_agent(&yaml, role_mocks(default_mock, router_mock));
22880 agent
22881 .set_context("route", serde_json::json!("billing"))
22882 .unwrap();
22883
22884 let response = agent.chat("billing please").await.unwrap();
22885
22886 assert_eq!(agent.current_state().as_deref(), Some("billing"));
22887 assert_eq!(response.content, "Billing answer");
22888 assert_eq!(router_counter.call_count(), 1);
22890 }
22891
22892 #[tokio::test]
22893 async fn test_speculative_skill_capacity_exhausted_still_routes_skill_serially() {
22894 let default_mock = mock_with_response("Draft response");
22895 let router_mock = mock_with_response("helper");
22896 let router_counter = router_mock.clone();
22897 let yaml = skills_with_parallel_transition_yaml(
22899 " speculative_skill_routing: true\n max_speculative_llm_calls_per_turn: 2\n max_parallel_runtime_tasks: 2",
22900 "",
22901 );
22902 let agent =
22903 build_skills_beside_transition_agent(&yaml, role_mocks(default_mock, router_mock));
22904
22905 let response = agent.chat("please use helper").await.unwrap();
22906
22907 assert_eq!(
22908 response.metadata.as_ref().and_then(|m| m.get("skill_id")),
22909 Some(&serde_json::json!("helper"))
22910 );
22911 assert_eq!(router_counter.call_count(), 1);
22912 assert_eq!(
22913 branch_events_with_commit_behavior(&agent, "skill_selection"),
22914 0
22915 );
22916 }
22917
22918 #[tokio::test]
22919 async fn test_speculative_transition_and_skill_both_enabled_unchanged() {
22920 let default_mock = mock_with_response("Draft response");
22921 let router_mock = mock_with_response("helper");
22922 let router_counter = router_mock.clone();
22923 let yaml = skills_with_parallel_transition_yaml(
22924 " speculative_skill_routing: true\n max_speculative_llm_calls_per_turn: 3\n max_parallel_runtime_tasks: 3",
22925 "",
22926 );
22927 let agent =
22928 build_skills_beside_transition_agent(&yaml, role_mocks(default_mock, router_mock));
22929
22930 let response = agent.chat("please use helper").await.unwrap();
22931
22932 assert_eq!(
22933 response.metadata.as_ref().and_then(|m| m.get("skill_id")),
22934 Some(&serde_json::json!("helper"))
22935 );
22936 assert_eq!(router_counter.call_count(), 1);
22937 assert!(branch_events_with_commit_behavior(&agent, "skill_selection") > 0);
22939 }
22940
22941 const BUFFERED_STREAMING_YAML_FRAGMENT: &str = "streaming:\n enabled: true\n buffer_size: 16";
22942 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";
22943
22944 #[tokio::test]
22945 async fn test_buffered_streaming_skill_wins_after_transition_miss() {
22946 let mut default_mock = mock_with_response("draft one two");
22947 default_mock.set_latency(10);
22948 let router_mock = mock_with_response("helper");
22949 let router_counter = router_mock.clone();
22950 let yaml = skills_with_parallel_transition_yaml(
22951 BUFFERED_OPTIMIZATION_FRAGMENT,
22952 BUFFERED_STREAMING_YAML_FRAGMENT,
22953 );
22954 let agent =
22955 build_skills_beside_transition_agent(&yaml, role_mocks(default_mock, router_mock));
22956
22957 let (content, chunks, final_response) =
22958 collect_stream_events(&agent, "please use helper").await;
22959
22960 assert!(!chunks.iter().any(StreamChunk::is_error), "{chunks:?}");
22961 assert!(
22962 !content.contains("draft"),
22963 "buffered draft must be discarded when a skill wins: {content:?}"
22964 );
22965 let final_response = final_response.expect("Final");
22966 assert_eq!(
22967 final_response
22968 .metadata
22969 .as_ref()
22970 .and_then(|m| m.get("skill_id")),
22971 Some(&serde_json::json!("helper"))
22972 );
22973 assert_eq!(content, final_response.content);
22974 assert_eq!(router_counter.call_count(), 1);
22975 }
22976
22977 #[tokio::test]
22978 async fn test_buffered_streaming_skill_miss_releases_buffer_and_commits_draft() {
22979 let mut default_mock = mock_with_response("draft one two");
22980 default_mock.set_latency(10);
22981 let router_mock = mock_with_response("none");
22982 let router_counter = router_mock.clone();
22983 let yaml = skills_with_parallel_transition_yaml(
22984 BUFFERED_OPTIMIZATION_FRAGMENT,
22985 BUFFERED_STREAMING_YAML_FRAGMENT,
22986 );
22987 let agent =
22988 build_skills_beside_transition_agent(&yaml, role_mocks(default_mock, router_mock));
22989
22990 let (content, chunks, final_response) = collect_stream_events(&agent, "just chat").await;
22991
22992 assert!(!chunks.iter().any(StreamChunk::is_error), "{chunks:?}");
22993 assert_eq!(content, "draft one two");
22994 assert_eq!(final_response.expect("Final").content, "draft one two");
22995 assert_eq!(router_counter.call_count(), 1);
22996 }
22997
22998 #[tokio::test]
22999 async fn test_buffered_streaming_transition_win_skips_skill_selection() {
23000 let default_mock = mock_with_response("Billing answer");
23001 let router_mock = mock_with_response("none");
23002 let router_counter = router_mock.clone();
23003 let yaml = skills_with_parallel_transition_yaml(
23004 BUFFERED_OPTIMIZATION_FRAGMENT,
23005 BUFFERED_STREAMING_YAML_FRAGMENT,
23006 );
23007 let agent =
23008 build_skills_beside_transition_agent(&yaml, role_mocks(default_mock, router_mock));
23009 agent
23010 .set_context("route", serde_json::json!("billing"))
23011 .unwrap();
23012
23013 let (content, chunks, final_response) =
23014 collect_stream_events(&agent, "billing please").await;
23015
23016 assert!(!chunks.iter().any(StreamChunk::is_error), "{chunks:?}");
23017 assert_eq!(agent.current_state().as_deref(), Some("billing"));
23018 assert_eq!(content, "Billing answer");
23019 assert_eq!(final_response.expect("Final").content, "Billing answer");
23020 assert_eq!(router_counter.call_count(), 1);
23022 }
23023
23024 #[tokio::test]
23025 async fn parity_buffered_policy_with_skills() {
23026 let yaml = skills_with_parallel_transition_yaml(
23027 BUFFERED_OPTIMIZATION_FRAGMENT,
23028 BUFFERED_STREAMING_YAML_FRAGMENT,
23029 );
23030 let build = || {
23031 build_skills_beside_transition_agent(
23032 &yaml,
23033 role_mocks(
23034 mock_with_response("draft one two"),
23035 mock_with_response("helper"),
23036 ),
23037 )
23038 };
23039 let (blocking, _, _) = assert_blocking_streaming_parity(build, "please use helper").await;
23040 assert_eq!(
23041 blocking.metadata.as_ref().and_then(|m| m.get("skill_id")),
23042 Some(&serde_json::json!("helper"))
23043 );
23044 }
23045
23046 #[tokio::test]
23047 async fn parity_buffered_policy_with_cot() {
23048 let yaml = format!(
23049 r#"
23050name: BufferedCotAgent
23051system_prompt: "Think first."
23052llm:
23053 default: default
23054streaming:
23055 enabled: true
23056 buffer_size: 16
23057reasoning:
23058 mode: cot
23059runtime:
23060 optimization:
23061 enabled: true
23062 speculative_state_transitions: true
23063{BUFFERED_OPTIMIZATION_FRAGMENT}
23064states:
23065 initial: triage
23066 states:
23067 triage:
23068 prompt: "Triage state."
23069 transitions:
23070 - to: billing
23071 guard:
23072 context:
23073 route:
23074 eq: billing
23075 timing: parallel
23076 billing:
23077 prompt: "Billing state."
23078"#
23079 );
23080 let build = || {
23081 AgentBuilder::from_yaml(&yaml)
23082 .unwrap()
23083 .llm_alias(
23084 "default",
23085 Arc::new(mock_with_response(
23086 "<thinking>step by step</thinking>Reasoned answer",
23087 )),
23088 )
23089 .build()
23090 .unwrap()
23091 };
23092 let (blocking, streamed, _) = assert_blocking_streaming_parity(build, "hello").await;
23093 assert_eq!(blocking.content, "Reasoned answer");
23094 let mode = streamed
23095 .metadata
23096 .as_ref()
23097 .and_then(|m| m.get("reasoning"))
23098 .and_then(|r| r.get("mode_used"))
23099 .cloned();
23100 assert_eq!(
23102 mode,
23103 Some(serde_json::to_value(ReasoningMode::CoT).unwrap())
23104 );
23105 }
23106
23107 fn post_response_transition_yaml(states_extra: &str, billing_extra: &str) -> String {
23114 format!(
23115 r#"
23116name: PostResponseTransitionAgent
23117system_prompt: "You are helpful."
23118streaming:
23119 enabled: true
23120states:
23121 initial: intake
23122{states_extra}
23123 states:
23124 intake:
23125 prompt: "Intake"
23126 transitions:
23127 - to: billing
23128 guard:
23129 context:
23130 route:
23131 eq: billing
23132 billing:
23133 prompt: "Billing"
23134{billing_extra}
23135"#
23136 )
23137 }
23138
23139 fn build_post_response_transition_agent(yaml: &str, mock: MockLLMProvider) -> RuntimeAgent {
23140 let agent = AgentBuilder::from_yaml(yaml)
23141 .unwrap()
23142 .llm(Arc::new(mock))
23143 .build()
23144 .unwrap();
23145 agent
23146 .set_context("route", serde_json::json!("billing"))
23147 .unwrap();
23148 agent
23149 }
23150
23151 fn count_occurrences(haystack: &str, needle: &str) -> usize {
23152 haystack.matches(needle).count()
23153 }
23154
23155 #[tokio::test]
23156 async fn test_stream_transition_without_regeneration_emits_content_once() {
23157 let yaml = post_response_transition_yaml(" regenerate_on_transition: false", "");
23158 let agent =
23159 build_post_response_transition_agent(&yaml, mock_with_response("Intake answer"));
23160
23161 let (content, chunks, final_response) = collect_stream_events(&agent, "hello").await;
23162
23163 assert!(!chunks.iter().any(StreamChunk::is_error), "{chunks:?}");
23164 assert_eq!(agent.current_state().as_deref(), Some("billing"));
23165 assert_eq!(
23166 count_occurrences(&content, "Intake answer"),
23167 1,
23168 "committed content must not be emitted twice: {content:?}"
23169 );
23170 assert_eq!(final_response.expect("Final").content, content);
23171 assert!(
23172 chunks
23173 .iter()
23174 .any(|c| matches!(c, StreamChunk::StateTransition { .. }))
23175 );
23176 }
23177
23178 #[tokio::test]
23179 async fn test_stream_transition_without_regeneration_buffered_emits_content_once() {
23180 let yaml = post_response_transition_yaml(" regenerate_on_transition: false", "");
23181 let mut mock = mock_with_response("Intake answer");
23182 mock.set_tool_choice(Some(ToolChoice::Auto));
23184 let agent = build_post_response_transition_agent(&yaml, mock);
23185
23186 let (content, chunks, final_response) = collect_stream_events(&agent, "hello").await;
23187
23188 assert!(!chunks.iter().any(StreamChunk::is_error), "{chunks:?}");
23189 assert_eq!(agent.current_state().as_deref(), Some("billing"));
23190 assert_eq!(
23191 count_occurrences(&content, "Intake answer"),
23192 1,
23193 "{content:?}"
23194 );
23195 assert_eq!(final_response.expect("Final").content, content);
23196 }
23197
23198 #[tokio::test]
23199 async fn test_stream_state_regenerate_on_enter_false_emits_content_once() {
23200 let yaml = post_response_transition_yaml("", " regenerate_on_enter: false");
23201 let agent =
23202 build_post_response_transition_agent(&yaml, mock_with_response("Intake answer"));
23203
23204 let (content, chunks, final_response) = collect_stream_events(&agent, "hello").await;
23205
23206 assert!(!chunks.iter().any(StreamChunk::is_error), "{chunks:?}");
23207 assert_eq!(agent.current_state().as_deref(), Some("billing"));
23208 assert_eq!(
23209 count_occurrences(&content, "Intake answer"),
23210 1,
23211 "{content:?}"
23212 );
23213 assert_eq!(final_response.expect("Final").content, content);
23214 }
23215
23216 #[tokio::test]
23217 async fn test_stream_transition_with_regeneration_emits_replacement() {
23218 let yaml = post_response_transition_yaml("", "");
23219 let agent = build_post_response_transition_agent(
23220 &yaml,
23221 mock_with_responses(vec!["Intake answer", "Billing answer"]),
23222 );
23223
23224 let (content, chunks, final_response) = collect_stream_events(&agent, "hello").await;
23225
23226 assert!(!chunks.iter().any(StreamChunk::is_error), "{chunks:?}");
23227 assert_eq!(agent.current_state().as_deref(), Some("billing"));
23228 assert_eq!(count_occurrences(&content, "Intake answer"), 1);
23230 assert_eq!(count_occurrences(&content, "Billing answer"), 1);
23231 assert_eq!(final_response.expect("Final").content, "Billing answer");
23232 }
23233
23234 #[tokio::test]
23235 async fn test_blocking_transition_without_regeneration_unchanged() {
23236 let yaml = post_response_transition_yaml(" regenerate_on_transition: false", "");
23237 let agent =
23238 build_post_response_transition_agent(&yaml, mock_with_response("Intake answer"));
23239
23240 let response = agent.chat("hello").await.unwrap();
23241
23242 assert_eq!(response.content, "Intake answer");
23243 assert_eq!(agent.current_state().as_deref(), Some("billing"));
23244 }
23245
23246 #[tokio::test]
23247 async fn parity_transition_regenerate_off() {
23248 let yaml = post_response_transition_yaml(" regenerate_on_transition: false", "");
23249 let build =
23250 || build_post_response_transition_agent(&yaml, mock_with_response("Intake answer"));
23251 assert_blocking_streaming_parity(build, "hello").await;
23252 }
23253
23254 #[tokio::test]
23255 async fn parity_transition_regenerate_on() {
23256 let yaml = post_response_transition_yaml("", "");
23257 let build = || {
23258 build_post_response_transition_agent(
23259 &yaml,
23260 mock_with_responses(vec!["Intake answer", "Billing answer"]),
23261 )
23262 };
23263 let (blocking, _, _) = assert_blocking_streaming_parity(build, "hello").await;
23264 assert_eq!(blocking.content, "Billing answer");
23265 }
23266
23267 fn rejecting_process_processor() -> ProcessProcessor {
23272 use ai_agents_process::ProcessConfig;
23273 let validate_config = ai_agents_process::ValidateStage {
23274 id: Some("length_check".to_string()),
23275 condition: None,
23276 config: ai_agents_process::ValidateConfig {
23277 rules: vec![ai_agents_process::ValidationRule::MinLength {
23278 min_length: 10,
23279 on_fail: ai_agents_process::ValidationAction {
23280 action: ai_agents_process::ValidationActionType::Reject,
23281 message: None,
23282 },
23283 }],
23284 ..Default::default()
23285 },
23286 };
23287 ProcessProcessor::new(ProcessConfig {
23288 input: vec![ai_agents_process::ProcessStage::Validate(validate_config)],
23289 ..Default::default()
23290 })
23291 }
23292
23293 fn looks_like_rejection(content: &str) -> bool {
23295 content.contains("rejected")
23296 || content.contains("Input rejected")
23297 || content.contains("too short")
23298 || content.contains("Too short")
23299 || content.len() < 50
23300 }
23301
23302 #[tokio::test]
23303 async fn test_stream_input_rejection_is_final_response() {
23304 let responses = Arc::new(std::sync::atomic::AtomicUsize::new(0));
23305 let hooks = Arc::new(ResponseCountingHooks {
23306 responses: Arc::clone(&responses),
23307 });
23308 let mock = mock_with_response("Should not reach here.");
23309 let llm_calls = mock.clone();
23310 let agent = AgentBuilder::new()
23311 .system_prompt("You are helpful.")
23312 .llm(Arc::new(mock))
23313 .process_processor(rejecting_process_processor())
23314 .hooks(hooks.clone())
23315 .build()
23316 .unwrap();
23317
23318 let (content, chunks, final_response) = collect_stream_events(&agent, "Hi").await;
23319
23320 assert!(
23321 !chunks.iter().any(StreamChunk::is_error),
23322 "rejection is a response, not a stream error: {chunks:?}"
23323 );
23324 let final_response = final_response.expect("rejection must finalize as Final");
23325 assert!(
23326 looks_like_rejection(&final_response.content),
23327 "Expected rejection response, got: {}",
23328 final_response.content
23329 );
23330 assert_eq!(content, final_response.content);
23331 assert_eq!(
23332 llm_calls.call_count(),
23333 0,
23334 "rejected input must not reach the LLM"
23335 );
23336 assert_eq!(responses.load(Ordering::SeqCst), 1, "on_response must fire");
23337 }
23338
23339 #[tokio::test]
23340 async fn parity_input_rejection() {
23341 let build = || {
23342 AgentBuilder::new()
23343 .system_prompt("You are helpful.")
23344 .llm(Arc::new(mock_with_response("Should not reach here.")))
23345 .process_processor(rejecting_process_processor())
23346 .build()
23347 .unwrap()
23348 };
23349 let (blocking, _, _) = assert_blocking_streaming_parity(build, "Hi").await;
23350 assert!(
23351 looks_like_rejection(&blocking.content),
23352 "{}",
23353 blocking.content
23354 );
23355 }
23356
23357 fn pre_response_transition_yaml(streaming_policy: &str) -> String {
23358 format!(
23359 r#"
23360name: StreamingPreflightAgent
23361system_prompt: "You route before streaming."
23362runtime:
23363 optimization:
23364 enabled: true
23365 pre_response_deterministic_transitions: true
23366 streaming_policy: {streaming_policy}
23367streaming:
23368 enabled: true
23369 buffer_size: 16
23370states:
23371 initial: greeting
23372 states:
23373 greeting:
23374 prompt: "OLD_STATE_SENTINEL"
23375 transitions:
23376 - to: billing
23377 guard:
23378 context:
23379 topic:
23380 eq: billing
23381 timing: pre_response
23382 billing:
23383 prompt: "Billing state."
23384"#
23385 )
23386 }
23387
23388 fn build_pre_response_transition_agent(yaml: &str) -> RuntimeAgent {
23389 let agent = AgentBuilder::from_yaml(yaml)
23390 .unwrap()
23391 .llm(Arc::new(mock_with_response("Billing streamed response")))
23392 .build()
23393 .unwrap();
23394 agent
23395 .set_context("topic", serde_json::json!("billing"))
23396 .unwrap();
23397 agent
23398 }
23399
23400 #[tokio::test]
23401 async fn test_stream_buffered_policy_runs_pre_response_deterministic_transition() {
23402 let yaml = pre_response_transition_yaml("buffer_until_routing_done");
23403 let agent = build_pre_response_transition_agent(&yaml);
23404
23405 let (content, chunks, final_response) =
23406 collect_stream_events(&agent, "billing please").await;
23407
23408 assert!(!chunks.iter().any(StreamChunk::is_error), "{chunks:?}");
23409 assert_eq!(agent.current_state().as_deref(), Some("billing"));
23410 assert!(content.contains("Billing streamed response"));
23411 assert!(!content.contains("OLD_STATE_SENTINEL"));
23412 assert_eq!(final_response.expect("Final").content, content);
23413 }
23414
23415 #[tokio::test]
23416 async fn test_stream_disabled_policy_skips_preflight() {
23417 let yaml = pre_response_transition_yaml("disabled");
23418 let agent = build_pre_response_transition_agent(&yaml);
23419
23420 let (_, chunks, final_response) = collect_stream_events(&agent, "billing please").await;
23421
23422 assert!(!chunks.iter().any(StreamChunk::is_error), "{chunks:?}");
23423 assert!(final_response.is_some());
23424 assert_eq!(agent.current_state().as_deref(), Some("greeting"));
23428 }
23429
23430 #[tokio::test]
23431 async fn parity_pre_response_transition_buffered_policy() {
23432 let yaml = pre_response_transition_yaml("buffer_until_routing_done");
23433 let build = || build_pre_response_transition_agent(&yaml);
23434 assert_blocking_streaming_parity(build, "billing please").await;
23435 }
23436
23437 fn calculator_agent_with(mock: MockLLMProvider) -> RuntimeAgent {
23442 let mut tools = ai_agents_tools::ToolRegistry::new();
23443 tools
23444 .register(Arc::new(ai_agents_tools::CalculatorTool))
23445 .unwrap();
23446 AgentBuilder::new()
23447 .system_prompt("You are a calculator assistant.")
23448 .llm(Arc::new(mock))
23449 .tools(tools)
23450 .build()
23451 .unwrap()
23452 }
23453
23454 #[tokio::test]
23455 async fn test_stream_tool_start_events_precede_results_for_batch() {
23456 let mock = mock_with_responses(vec![
23457 r#"[{"tool": "calculator", "arguments": {"expression": "1+1"}}, {"tool": "calculator", "arguments": {"expression": "2+2"}}]"#,
23458 "Both answers are ready.",
23459 ]);
23460 let agent = calculator_agent_with(mock);
23461
23462 let (_, chunks, final_response) = collect_stream_events(&agent, "compute both").await;
23463
23464 assert!(!chunks.iter().any(StreamChunk::is_error), "{chunks:?}");
23465 let final_response = final_response.expect("Final");
23466 assert_eq!(final_response.tool_calls.as_ref().map(Vec::len), Some(2));
23467
23468 let tool_events: Vec<&StreamChunk> = chunks
23469 .iter()
23470 .filter(|c| {
23471 matches!(
23472 c,
23473 StreamChunk::ToolCallStart { .. }
23474 | StreamChunk::ToolResult { .. }
23475 | StreamChunk::ToolCallEnd { .. }
23476 )
23477 })
23478 .collect();
23479 assert_eq!(tool_events.len(), 6, "{tool_events:?}");
23480 assert!(matches!(tool_events[0], StreamChunk::ToolCallStart { .. }));
23482 assert!(matches!(tool_events[1], StreamChunk::ToolCallStart { .. }));
23483 assert!(matches!(
23484 tool_events[2],
23485 StreamChunk::ToolResult { success: true, .. }
23486 ));
23487 assert!(matches!(tool_events[3], StreamChunk::ToolCallEnd { .. }));
23488 assert!(matches!(
23489 tool_events[4],
23490 StreamChunk::ToolResult { success: true, .. }
23491 ));
23492 assert!(matches!(tool_events[5], StreamChunk::ToolCallEnd { .. }));
23493 }
23494
23495 #[tokio::test]
23496 async fn test_stream_clarification_final_carries_options_and_detection() {
23497 let responses = || {
23498 vec![
23499 r#"{"is_ambiguous":true,"confidence":0.2,"ambiguity_type":"missing_target","reasoning":"target missing","what_is_unclear":["target"],"detected_language":"en"}"#,
23500 r#"{"question":"What should I send?","options":["report","invoice"]}"#,
23501 ]
23502 };
23503 let (blocking_agent, _) = state_disambiguation_agent(responses(), true, None, true);
23504 let (streaming_agent, _) = state_disambiguation_agent(responses(), true, None, true);
23505
23506 let blocking = blocking_agent.chat("Send it").await.unwrap();
23507 let (_, chunks, streamed) = collect_stream_events(&streaming_agent, "Send it").await;
23508 assert!(!chunks.iter().any(StreamChunk::is_error), "{chunks:?}");
23509 let streamed = streamed.expect("clarification must finalize as Final");
23510
23511 assert_eq!(streamed.content, "What should I send?");
23512 let streamed_meta = streamed
23513 .metadata
23514 .as_ref()
23515 .and_then(|m| m.get("disambiguation"))
23516 .cloned()
23517 .expect("disambiguation metadata");
23518 for key in ["status", "options", "clarifying", "detection"] {
23519 assert!(
23520 streamed_meta.get(key).is_some(),
23521 "missing {key}: {streamed_meta}"
23522 );
23523 }
23524 assert_eq!(
23525 streamed_meta.get("detection").and_then(|d| d.get("type")),
23526 Some(&serde_json::json!("missing_target"))
23527 );
23528 assert_eq!(
23529 blocking
23530 .metadata
23531 .as_ref()
23532 .and_then(|m| m.get("disambiguation")),
23533 Some(&streamed_meta),
23534 "blocking and streaming clarification metadata must be identical"
23535 );
23536 }
23537
23538 struct FailingMemory {
23540 messages: parking_lot::RwLock<Vec<ChatMessage>>,
23541 fail_on_add: usize,
23542 adds: std::sync::atomic::AtomicUsize,
23543 }
23544
23545 #[async_trait]
23546 impl ai_agents_core::Memory for FailingMemory {
23547 async fn add_message(&self, message: ChatMessage) -> Result<()> {
23548 let n = self.adds.fetch_add(1, Ordering::SeqCst) + 1;
23549 if n == self.fail_on_add {
23550 return Err(AgentError::Other(format!(
23551 "simulated memory failure on add #{n}"
23552 )));
23553 }
23554 self.messages.write().push(message);
23555 Ok(())
23556 }
23557
23558 async fn get_messages(&self, limit: Option<usize>) -> Result<Vec<ChatMessage>> {
23559 let messages = self.messages.read();
23560 Ok(match limit {
23561 Some(n) if n < messages.len() => messages[messages.len() - n..].to_vec(),
23562 _ => messages.clone(),
23563 })
23564 }
23565
23566 async fn clear(&self) -> Result<()> {
23567 self.messages.write().clear();
23568 Ok(())
23569 }
23570
23571 fn len(&self) -> usize {
23572 self.messages.read().len()
23573 }
23574
23575 async fn restore(&self, snapshot: ai_agents_core::MemorySnapshot) -> Result<()> {
23576 *self.messages.write() = snapshot.messages;
23577 Ok(())
23578 }
23579 }
23580
23581 impl ai_agents_memory::Memory for FailingMemory {}
23582
23583 #[tokio::test]
23584 async fn test_stream_memory_write_failure_surfaces_as_error() {
23585 let yaml = r#"
23588name: TransitionOnToolCallAgent
23589system_prompt: "You are helpful."
23590streaming:
23591 enabled: true
23592states:
23593 initial: intake
23594 states:
23595 intake:
23596 prompt: "Intake"
23597 transitions:
23598 - to: billing
23599 guard:
23600 context:
23601 route:
23602 eq: billing
23603 billing:
23604 prompt: "Billing"
23605"#;
23606 let build = |fail_on_add: usize| {
23607 let mut tools = ai_agents_tools::ToolRegistry::new();
23608 tools
23609 .register(Arc::new(ai_agents_tools::CalculatorTool))
23610 .unwrap();
23611 let agent = AgentBuilder::from_yaml(yaml)
23612 .unwrap()
23613 .llm(Arc::new(mock_with_responses(vec![
23614 r#"{"tool": "calculator", "arguments": {"expression": "1+1"}}"#,
23615 "Billing answer",
23616 ])))
23617 .tools(tools)
23618 .memory(Arc::new(FailingMemory {
23619 messages: parking_lot::RwLock::new(Vec::new()),
23620 fail_on_add,
23621 adds: std::sync::atomic::AtomicUsize::new(0),
23622 }))
23623 .build()
23624 .unwrap();
23625 agent
23626 .set_context("route", serde_json::json!("billing"))
23627 .unwrap();
23628 agent
23629 };
23630
23631 let blocking = build(2).chat("compute").await;
23632 assert!(
23633 blocking.is_err(),
23634 "blocking must surface the memory failure"
23635 );
23636
23637 let (_, chunks, final_response) = collect_stream_events(&build(2), "compute").await;
23638 assert!(
23639 final_response.is_none(),
23640 "streaming must not finalize after a memory failure"
23641 );
23642 assert!(
23643 chunks.iter().any(|c| matches!(c, StreamChunk::Error { message } if message.contains("simulated memory failure"))),
23644 "streaming must surface the memory failure: {chunks:?}"
23645 );
23646
23647 assert!(build(usize::MAX).chat("compute").await.is_ok());
23649 }
23650
23651 #[tokio::test]
23652 async fn parity_tool_execution() {
23653 let build = || {
23654 calculator_agent_with(mock_with_responses(vec![
23655 r#"{"tool": "calculator", "arguments": {"expression": "2+2"}}"#,
23656 "The answer is 4.",
23657 ]))
23658 };
23659 let (blocking, _, chunks) = assert_blocking_streaming_parity(build, "What is 2+2?").await;
23660 assert_eq!(blocking.content, "The answer is 4.");
23661 assert!(
23662 chunks
23663 .iter()
23664 .any(|c| matches!(c, StreamChunk::ToolResult { .. }))
23665 );
23666 }
23667
23668 #[tokio::test]
23669 async fn parity_disambiguation_clarification() {
23670 let build = || {
23671 state_disambiguation_agent(
23672 vec![
23673 r#"{"is_ambiguous":true,"confidence":0.2,"ambiguity_type":"missing_target","reasoning":"target missing","what_is_unclear":["target"],"detected_language":"en"}"#,
23674 r#"{"question":"What should I send?","options":null}"#,
23675 ],
23676 true,
23677 None,
23678 true,
23679 )
23680 .0
23681 };
23682 let (blocking, _, _) = assert_blocking_streaming_parity(build, "Send it").await;
23683 assert_eq!(blocking.content, "What should I send?");
23684 }
23685
23686 #[tokio::test]
23687 async fn parity_reflection_enabled() {
23688 let yaml = r#"
23689name: ReflectionAgent
23690system_prompt: "You are careful."
23691reflection:
23692 enabled: true
23693 criteria:
23694 - "Is the answer helpful?"
23695"#;
23696 let build = || {
23697 AgentBuilder::from_yaml(yaml)
23698 .unwrap()
23699 .llm(Arc::new(mock_with_responses(vec![
23700 "Main answer",
23701 "OVERALL: PASS\nCONFIDENCE: 0.9",
23702 ])))
23703 .build()
23704 .unwrap()
23705 };
23706 let (blocking, streamed, _) = assert_blocking_streaming_parity(build, "hello").await;
23707 assert_eq!(blocking.content, "Main answer");
23708 assert!(metadata_keys(&streamed).contains("reflection"));
23709 }
23710
23711 #[tokio::test]
23712 async fn parity_cot_hidden_thinking() {
23713 let yaml = r#"
23714name: CotHiddenAgent
23715system_prompt: "Think first."
23716reasoning:
23717 mode: cot
23718 output: hidden
23719"#;
23720 let build = || {
23721 AgentBuilder::from_yaml(yaml)
23722 .unwrap()
23723 .llm(Arc::new(mock_with_response(
23724 "<thinking>step by step</thinking>Visible answer",
23725 )))
23726 .build()
23727 .unwrap()
23728 };
23729 let (blocking, _, _) = assert_blocking_streaming_parity(build, "hello").await;
23731 assert_eq!(blocking.content, "Visible answer");
23732 }
23733
23734 #[tokio::test]
23735 async fn parity_skill_route() {
23736 let yaml = skills_with_parallel_transition_yaml(
23737 " max_speculative_llm_calls_per_turn: 2\n max_parallel_runtime_tasks: 2",
23738 "",
23739 );
23740 let build = || {
23741 build_skills_beside_transition_agent(
23742 &yaml,
23743 role_mocks(
23744 mock_with_response("Draft response"),
23745 mock_with_response("helper"),
23746 ),
23747 )
23748 };
23749 assert_blocking_streaming_parity(build, "please use helper").await;
23750 }
23751}