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, LLMError, LLMProvider,
246 LLMResponse, LLMToolDefinition, LLMToolRequest, PermissionOutcome, Result, ToolActorContext,
247 ToolApprovalRecord, ToolApprovalStatus, ToolCallClassification, ToolCallSource,
248 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 ClarificationObserver, ClarificationParseFuture, ClarificationQuestionFuture,
255 ConfirmationParseFuture, DisambiguationConfig, DisambiguationContext, DisambiguationManager,
256 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
333#[derive(Clone)]
335struct ActiveNativeExchange {
336 exchange_id: String,
337 call_ids: Vec<String>,
338}
339
340struct CommittedTextResponse<'a> {
344 processed_input: &'a str,
345 input_context: &'a HashMap<String, Value>,
346 answer: String,
347 reasoning_mode: ReasoningMode,
348 auto_detected: bool,
349 iterations: u32,
350 thinking_content: Option<String>,
351 all_tool_calls: Vec<ToolCall>,
352}
353
354struct AgentResponseParts {
358 content: String,
359 all_tool_calls: Vec<ToolCall>,
360 reasoning_mode: ReasoningMode,
361 auto_detected: bool,
362 iterations: u32,
363 thinking: Option<String>,
364 reflection_metadata: Option<ReflectionMetadata>,
365}
366
367type RuntimeStreamTerminalSlot = Arc<RwLock<Option<AgentResponse>>>;
371
372fn new_runtime_stream_terminal_slot() -> RuntimeStreamTerminalSlot {
376 Arc::new(RwLock::new(None))
377}
378
379fn record_runtime_stream_final(slot: &RuntimeStreamTerminalSlot, response: AgentResponse) {
383 *slot.write() = Some(response);
384}
385
386#[derive(Clone, Copy)]
387struct DisambiguationOwnership {
388 epoch: u64,
389 state_generation: Option<u64>,
390}
391
392enum SkillRouteResult {
394 NoMatch,
396 Response { skill_id: String, content: String },
398 NeedsClarification {
400 response: AgentResponse,
401 ownership: Option<DisambiguationOwnership>,
402 },
403}
404
405enum ParallelTransitionSelection {
407 Candidate(TransitionCandidate),
409 NoMatch,
411 ReservationExhausted,
413}
414
415enum PostLoopResult {
417 NoTransition(String),
419 Transitioned(String),
421 NeedsRedispatch,
424}
425
426struct StateTransitionReservation<'a> {
427 reserved: &'a AtomicBool,
428}
429
430impl Drop for StateTransitionReservation<'_> {
431 fn drop(&mut self) {
432 self.reserved.store(false, Ordering::SeqCst);
433 }
434}
435
436struct RootTurnCleanup<'a> {
437 agent: &'a RuntimeAgent,
438}
439
440impl<'a> RootTurnCleanup<'a> {
441 fn new(agent: &'a RuntimeAgent) -> Self {
442 Self { agent }
443 }
444}
445
446impl Drop for RootTurnCleanup<'_> {
447 fn drop(&mut self) {
448 self.agent.end_root_turn();
449 }
450}
451
452#[derive(Debug)]
454struct RuntimeControlState {
455 snapshot_guard: RwLock<()>,
457 version: AtomicU64,
459 emergency_deny: Arc<AtomicBool>,
461 tool_security_override: RwLock<Option<ToolSecurityEngine>>,
463 tool_scope_override: RwLock<Option<Vec<String>>>,
465}
466
467impl Default for RuntimeControlState {
468 fn default() -> Self {
469 Self {
470 snapshot_guard: RwLock::new(()),
471 version: AtomicU64::new(1),
472 emergency_deny: Arc::new(AtomicBool::new(false)),
473 tool_security_override: RwLock::new(None),
474 tool_scope_override: RwLock::new(None),
475 }
476 }
477}
478
479#[derive(Clone)]
481pub struct RuntimeControlHandle {
482 state: Arc<RuntimeControlState>,
483}
484
485impl RuntimeControlHandle {
486 pub fn version(&self) -> u64 {
488 self.state.version.load(Ordering::SeqCst)
489 }
490
491 fn bump(&self) -> u64 {
492 self.state.version.fetch_add(1, Ordering::SeqCst) + 1
493 }
494
495 pub fn set_tool_security(&self, config: ToolSecurityConfig) -> u64 {
497 self.try_set_tool_security(config)
498 .expect("invalid tool security configuration")
499 }
500
501 pub fn try_set_tool_security(&self, config: ToolSecurityConfig) -> Result<u64> {
503 config.validate()?;
504 let _guard = self.state.snapshot_guard.write();
505 let generation = self.bump();
506 *self.state.tool_security_override.write() = Some(
507 ToolSecurityEngine::new_with_policy_version(config, generation),
508 );
509 Ok(generation)
510 }
511
512 pub fn clear_tool_security_override(&self) -> u64 {
514 let _guard = self.state.snapshot_guard.write();
515 *self.state.tool_security_override.write() = None;
516 self.bump()
517 }
518
519 pub fn set_tool_scope(&self, tool_ids: Vec<String>) -> u64 {
521 let _guard = self.state.snapshot_guard.write();
522 *self.state.tool_scope_override.write() = Some(tool_ids);
523 self.bump()
524 }
525
526 pub fn clear_tool_scope_override(&self) -> u64 {
528 let _guard = self.state.snapshot_guard.write();
529 *self.state.tool_scope_override.write() = None;
530 self.bump()
531 }
532
533 pub fn set_emergency_deny(&self, enabled: bool) -> u64 {
535 let _guard = self.state.snapshot_guard.write();
536 self.state.emergency_deny.store(enabled, Ordering::SeqCst);
537 self.bump()
538 }
539
540 pub fn cancel_all(&self) -> u64 {
542 self.set_emergency_deny(true)
543 }
544}
545
546pub struct RuntimeAgent {
547 info: AgentInfo,
548 llm_registry: Arc<LLMRegistry>,
549 memory: Arc<dyn Memory>,
550 tools: Arc<ToolRegistry>,
551 skills: Vec<SkillDefinition>,
552 skill_router: Option<SkillRouter>,
553 skill_executor: Option<SkillExecutor>,
554 base_system_prompt: String,
555 max_iterations: u32,
556 iteration_count: RwLock<u32>,
557 max_context_tokens: u32,
558 memory_token_budget: Option<MemoryTokenBudget>,
559 recovery_manager: RecoveryManager,
560 tool_security: ToolSecurityEngine,
561 process_processor: Option<ProcessProcessor>,
562 message_filters: RwLock<HashMap<String, Arc<dyn MessageFilter>>>,
563 state_machine: Option<Arc<StateMachine>>,
564 transition_evaluator: Option<Arc<dyn TransitionEvaluator>>,
565 context_manager: Arc<ContextManager>,
566 template_renderer: TemplateRenderer,
567 tool_call_history: RwLock<Vec<ToolCallRecord>>,
568 parallel_tools: ParallelToolsConfig,
569 streaming: StreamingConfig,
570 hooks: Arc<dyn AgentHooks>,
571 hitl_engine: Option<HITLEngine>,
572 approval_handler: Arc<dyn ApprovalHandler>,
573 storage_config: StorageConfig,
574 storage: RwLock<Option<Arc<dyn AgentStorage>>>,
575 storage_init: tokio::sync::Mutex<()>,
576 reasoning_config: ReasoningConfig,
577 reflection_config: ReflectionConfig,
578 disambiguation_manager: Option<DisambiguationManager>,
579 disambiguation_epoch: AtomicU64,
581 disambiguation_admission: tokio::sync::RwLock<()>,
583 state_transition_reserved: AtomicBool,
585 persona_manager: Option<Arc<ai_agents_persona::PersonaManager>>,
587 pending_skill_id: RwLock<Option<String>>,
591 current_plan: RwLock<Option<Plan>>,
592 declared_tool_ids: Option<Vec<String>>,
594 context_initialized: AtomicBool,
596 spawner: Option<Arc<crate::spawner::AgentSpawner>>,
598 spawner_registry: Option<Arc<crate::spawner::AgentRegistry>>,
600 redispatch_depth: RwLock<u32>,
603 active_turn_context: RwLock<Option<TurnOptimizationContext>>,
605 root_user_message_committed: AtomicBool,
607 active_native_exchanges: RwLock<Vec<ActiveNativeExchange>>,
609 actor_id: RwLock<Option<String>>,
611 fact_store: RwLock<Option<Arc<ai_agents_facts::FactStore>>>,
613 fact_extractor: RwLock<Option<Arc<dyn ai_agents_facts::FactExtractor>>>,
616 actor_facts_cache: Arc<RwLock<HashMap<String, Vec<ai_agents_core::KeyFact>>>>,
618 messages_since_extraction: Arc<RwLock<usize>>,
620 actor_memory_config: Option<ai_agents_facts::ActorMemoryConfig>,
622 facts_config: Option<ai_agents_facts::FactsConfig>,
624 session_metadata: RwLock<ai_agents_core::SessionMetadata>,
626 current_session_id: RwLock<Option<String>>,
628 relationship_manager: Option<Arc<RelationshipManager>>,
630 observability_manager: Option<Arc<ObservabilityManager>>,
632 runtime_config: RuntimeConfig,
634 background_maintenance: Arc<BackgroundMaintenanceQueue>,
636 resource_locks: ToolResourceLocks,
638 runtime_control: Arc<RuntimeControlState>,
640 root_turn_gate: RootTurnGate,
642}
643
644impl std::fmt::Debug for RuntimeAgent {
645 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
646 f.debug_struct("RuntimeAgent")
647 .field("info", &self.info)
648 .field("base_system_prompt", &self.base_system_prompt)
649 .field("max_iterations", &self.max_iterations)
650 .field("skills_count", &self.skills.len())
651 .field("max_context_tokens", &self.max_context_tokens)
652 .field("has_state_machine", &self.state_machine.is_some())
653 .field("parallel_tools", &self.parallel_tools)
654 .field("streaming", &self.streaming)
655 .field("has_hooks", &true)
656 .field("has_hitl", &self.hitl_engine.is_some())
657 .field("storage_type", &self.storage_config.storage_type())
658 .field("reasoning_mode", &self.reasoning_config.mode)
659 .field("reflection_enabled", &self.reflection_config.enabled)
660 .field("declared_tool_ids", &self.declared_tool_ids)
661 .field("has_persona", &self.persona_manager.is_some())
662 .field("has_observability", &self.observability_manager.is_some())
663 .finish_non_exhaustive()
664 }
665}
666
667struct ObservabilityClarificationObserver;
668
669impl ClarificationObserver for ObservabilityClarificationObserver {
670 fn observe_question<'a>(
672 &'a self,
673 future: ClarificationQuestionFuture<'a>,
674 ) -> ClarificationQuestionFuture<'a> {
675 Box::pin(async move {
676 with_observation_purpose(ObservationPurpose::DisambiguationClarification, future).await
677 })
678 }
679
680 fn observe_parse<'a>(
682 &'a self,
683 future: ClarificationParseFuture<'a>,
684 ) -> ClarificationParseFuture<'a> {
685 Box::pin(async move {
686 with_observation_purpose(ObservationPurpose::DisambiguationClarification, future).await
687 })
688 }
689
690 fn observe_confirmation_parse<'a>(
692 &'a self,
693 future: ConfirmationParseFuture<'a>,
694 ) -> ConfirmationParseFuture<'a> {
695 Box::pin(async move {
696 with_observation_purpose(ObservationPurpose::DisambiguationClarification, future).await
697 })
698 }
699}
700
701struct ObservabilityProcessStageObserver;
702
703impl ProcessStageObserver for ObservabilityProcessStageObserver {
704 fn observe<'a>(
706 &'a self,
707 hint: ProcessPurposeHint,
708 future: ProcessStageFuture<'a>,
709 ) -> ProcessStageFuture<'a> {
710 Box::pin(async move {
711 with_observation_purpose(observation_purpose_for_process(hint), future).await
712 })
713 }
714}
715
716struct RegistryLLMGetter {
717 registry: Arc<LLMRegistry>,
718}
719
720impl LLMGetter for RegistryLLMGetter {
721 fn get_llm(&self, alias: &str) -> Option<Arc<dyn LLMProvider>> {
722 self.registry.get(alias).ok()
723 }
724}
725
726impl RuntimeAgent {
727 #[allow(clippy::too_many_arguments)]
729 pub fn new(
730 info: AgentInfo,
731 llm_registry: Arc<LLMRegistry>,
732 memory: Arc<dyn Memory>,
733 tools: Arc<ToolRegistry>,
734 skills: Vec<SkillDefinition>,
735 system_prompt: String,
736 max_iterations: u32,
737 ) -> Self {
738 let (skill_router, skill_executor) = if !skills.is_empty() {
739 let router_llm = llm_registry.router().ok();
740 let router = router_llm.map(|llm| SkillRouter::new(llm, skills.clone()));
741 let executor = SkillExecutor::new(llm_registry.clone(), tools.clone());
742 (router, Some(executor))
743 } else {
744 (None, None)
745 };
746
747 let context_manager =
748 ContextManager::new(HashMap::new(), info.name.clone(), info.version.clone());
749
750 Self {
751 info,
752 llm_registry,
753 memory,
754 tools,
755 skills,
756 skill_router,
757 skill_executor,
758 base_system_prompt: system_prompt,
759 max_iterations,
760 iteration_count: RwLock::new(0),
761 max_context_tokens: 128000,
762 memory_token_budget: None,
763 recovery_manager: RecoveryManager::default(),
764 tool_security: ToolSecurityEngine::default(),
765 process_processor: None,
766 message_filters: RwLock::new(HashMap::new()),
767 state_machine: None,
768 transition_evaluator: None,
769 context_manager: Arc::new(context_manager),
770 template_renderer: TemplateRenderer::new(),
771 tool_call_history: RwLock::new(Vec::new()),
772 parallel_tools: ParallelToolsConfig::default(),
773 streaming: StreamingConfig::default(),
774 hooks: Arc::new(NoopHooks),
775 hitl_engine: None,
776 approval_handler: Arc::new(RejectAllHandler::new()),
777 storage_config: StorageConfig::default(),
778 storage: RwLock::new(None),
779 storage_init: tokio::sync::Mutex::new(()),
780 reasoning_config: ReasoningConfig::default(),
781 reflection_config: ReflectionConfig::default(),
782 disambiguation_manager: None,
783 disambiguation_epoch: AtomicU64::new(0),
784 disambiguation_admission: tokio::sync::RwLock::new(()),
785 state_transition_reserved: AtomicBool::new(false),
786 persona_manager: None,
787 pending_skill_id: RwLock::new(None),
788 current_plan: RwLock::new(None),
789 declared_tool_ids: None,
790 context_initialized: AtomicBool::new(false),
791 spawner: None,
792 spawner_registry: None,
793 redispatch_depth: RwLock::new(0),
794 active_turn_context: RwLock::new(None),
795 root_user_message_committed: AtomicBool::new(false),
796 active_native_exchanges: RwLock::new(Vec::new()),
797 actor_id: RwLock::new(None),
798 fact_store: RwLock::new(None),
799 fact_extractor: RwLock::new(None),
800 actor_facts_cache: Arc::new(RwLock::new(HashMap::new())),
801 messages_since_extraction: Arc::new(RwLock::new(0)),
802 actor_memory_config: None,
803 facts_config: None,
804 session_metadata: RwLock::new(ai_agents_core::SessionMetadata::default()),
805 current_session_id: RwLock::new(None),
806 relationship_manager: None,
807 observability_manager: None,
808 runtime_config: RuntimeConfig::default(),
809 background_maintenance: Arc::new(BackgroundMaintenanceQueue::default()),
810 resource_locks: new_tool_resource_locks(),
811 runtime_control: Arc::new(RuntimeControlState::default()),
812 root_turn_gate: Arc::new(tokio::sync::Mutex::new(())),
813 }
814 }
815
816 pub fn with_declared_tool_ids(mut self, ids: Option<Vec<String>>) -> Self {
817 self.declared_tool_ids = ids;
818 self
819 }
820
821 pub fn with_storage_config(mut self, config: StorageConfig) -> Self {
822 self.storage_config = config;
823 self
824 }
825
826 pub fn with_storage(self, storage: Arc<dyn AgentStorage>) -> Self {
827 *self.storage.write() = Some(storage);
828 self
829 }
830
831 pub(crate) fn with_shared_resource_locks(mut self, locks: ToolResourceLocks) -> Self {
832 self.resource_locks = locks;
833 self
834 }
835
836 pub fn with_reasoning(mut self, config: ReasoningConfig) -> Self {
837 self.reasoning_config = config;
838 self
839 }
840
841 pub fn with_reflection(mut self, config: ReflectionConfig) -> Self {
842 self.reflection_config = config;
843 self
844 }
845
846 pub fn with_relationships(mut self, manager: Arc<RelationshipManager>) -> Self {
848 self.relationship_manager = Some(manager);
849 self
850 }
851
852 pub fn with_observability(mut self, manager: Arc<ObservabilityManager>) -> Self {
854 self.observability_manager = Some(manager);
855 self
856 }
857
858 pub fn with_runtime_config(mut self, config: RuntimeConfig) -> Self {
860 let max_tasks = config.optimization.post_turn.max_background_tasks;
861 self.background_maintenance = Arc::new(BackgroundMaintenanceQueue::new(max_tasks));
862 self.runtime_config = config;
863 self
864 }
865
866 pub fn runtime_config(&self) -> &RuntimeConfig {
868 &self.runtime_config
869 }
870
871 pub async fn flush_background_tasks(&self) -> Result<()> {
873 self.background_maintenance.flush_all().await
874 }
875
876 pub async fn flush_background_tasks_for_actor(&self, actor_id: &str) -> Result<()> {
878 self.background_maintenance.flush_scope(actor_id).await
879 }
880
881 pub async fn flush_background_tasks_for_purpose(
883 &self,
884 purpose: RuntimeTaskPurpose,
885 ) -> Result<()> {
886 self.background_maintenance.flush_purpose(purpose).await
887 }
888
889 pub async fn flush_background_tasks_for_actor_purpose(
891 &self,
892 actor_id: &str,
893 purpose: RuntimeTaskPurpose,
894 ) -> Result<()> {
895 self.background_maintenance
896 .flush_scope_purpose(actor_id, purpose)
897 .await
898 }
899
900 pub async fn shutdown_background_tasks(&self) -> Result<()> {
902 self.flush_background_tasks().await
903 }
904
905 pub fn observability(&self) -> Option<Arc<ObservabilityManager>> {
907 self.observability_manager.clone()
908 }
909
910 async fn export_observability_if_configured(&self) {
912 let Some(manager) = self.observability_manager.as_ref() else {
913 return;
914 };
915 let export = &manager.config().export;
916 if !export.write_report && !export.write_raw_events {
917 return;
918 }
919 if let Err(error) = manager.export().await {
920 warn!(error = %error, "Observability export failed");
921 }
922 }
923
924 pub fn relationship_manager(&self) -> Option<Arc<RelationshipManager>> {
926 self.relationship_manager.clone()
927 }
928
929 fn current_turn_actor_context(&self) -> Option<crate::TurnActorContext> {
930 current_turn_actor_context()
931 }
932
933 fn effective_actor_id(&self) -> Option<String> {
934 self.current_turn_actor_context()
935 .and_then(|ctx| ctx.effective_actor_id().map(|id| id.to_string()))
936 .or_else(|| self.actor_id.read().clone())
937 }
938
939 fn effective_origin_actor_id(&self) -> Option<String> {
940 self.current_turn_actor_context()
941 .and_then(|ctx| ctx.origin_actor_id.clone())
942 .or_else(|| self.actor_id.read().clone())
943 }
944
945 fn record_session_actor_if_needed(&self) {
946 if let Some(actor_id) = self.effective_origin_actor_id() {
947 let mut meta = self.session_metadata.write();
948 meta.actor_id = Some(actor_id.clone());
949 if !meta.actors.iter().any(|a| a == &actor_id) {
950 meta.actors.push(actor_id);
951 }
952 }
953 }
954
955 fn outbound_actor_context(&self) -> crate::TurnActorContext {
956 let mut context = self.current_turn_actor_context().unwrap_or_default();
957 if context.origin_actor_id.is_none() {
958 context.origin_actor_id = self.effective_origin_actor_id();
959 }
960 context.sender_agent_id = Some(self.info.id.clone());
961 context
962 }
963
964 fn observation_session_id(&self) -> Option<String> {
966 let mut current = self.current_session_id.write();
967 if current.is_none() {
968 *current = Some(new_observation_session_id());
969 }
970 current.clone()
971 }
972
973 fn build_observation_context(&self, actor_id: Option<String>) -> Option<SpanContext> {
975 let manager = self.observability_manager.as_ref()?;
976 let context = self.build_context_with_overlays();
977 let language = resolve_language_from_context(manager.config(), &context);
978 let context = current_observation_context()
979 .map(|parent| parent.child_for_agent(self.info.id.clone()).with_new_turn())
980 .unwrap_or_else(|| SpanContext::new_root(self.info.id.clone()));
981 Some(
982 context
983 .with_actor(actor_id.or_else(|| self.effective_actor_id()))
984 .with_session(self.observation_session_id())
985 .with_state(self.current_state())
986 .with_language(Some(language)),
987 )
988 }
989
990 fn current_runtime_observation_context(
992 &self,
993 purpose: ObservationPurpose,
994 ) -> Option<SpanContext> {
995 let manager = self.observability_manager.as_ref()?;
996 let context = self.build_context_with_overlays();
997 let language = resolve_language_from_context(manager.config(), &context);
998 let mut observation = current_observation_context()
999 .unwrap_or_else(|| SpanContext::new_root(self.info.id.clone()));
1000 observation.agent_id = self.info.id.clone();
1001 observation.actor_id = self.effective_actor_id();
1002 observation.session_id = self.observation_session_id();
1003 observation.state = self.current_state();
1004 observation.language = Some(language);
1005 observation.purpose = purpose;
1006 Some(observation)
1007 }
1008
1009 async fn observe_purpose<F, T>(&self, purpose: ObservationPurpose, future: F) -> T
1011 where
1012 F: Future<Output = T>,
1013 {
1014 if let Some(context) = self.current_runtime_observation_context(purpose) {
1015 with_observation_context(context, future).await
1016 } else {
1017 future.await
1018 }
1019 }
1020
1021 fn chat_with_actor_context_boxed<'a>(
1025 &'a self,
1026 input: &'a str,
1027 actor_context: crate::TurnActorContext,
1028 ) -> Pin<Box<dyn Future<Output = Result<AgentResponse>> + Send + 'a>> {
1029 Box::pin(async move {
1030 let RootTurnAdmission {
1031 guard,
1032 identity_stack,
1033 } = self.acquire_root_turn().await?;
1034 let result = scope_runtime_gate_identity_stack(&identity_stack, async move {
1035 let actor_id = actor_context.effective_actor_id().map(str::to_string);
1036 let run = async move {
1037 scope_actor_context(
1038 actor_context,
1039 Box::pin(async move { self.run_loop(input).await }),
1040 )
1041 .await
1042 };
1043 let result = if let Some(context) = self.build_observation_context(actor_id) {
1044 with_observation_context(context, run).await
1045 } else {
1046 run.await
1047 };
1048 self.export_observability_if_configured().await;
1049 result
1050 })
1051 .await;
1052 drop(guard);
1053 result
1054 })
1055 }
1056
1057 async fn acquire_root_turn(&self) -> Result<RootTurnAdmission> {
1059 let gate_identity = Arc::clone(&self.root_turn_gate);
1060 let current_identity_stack = current_runtime_gate_identity_stack();
1061 if current_identity_stack
1062 .iter()
1063 .any(|owned_gate| Arc::ptr_eq(owned_gate, &gate_identity))
1064 {
1065 return Err(AgentError::Other(format!(
1066 "RuntimeAgent '{}' rejected reentrant root turn ownership",
1067 self.info.id
1068 )));
1069 }
1070 let guard = Arc::clone(&gate_identity).lock_owned().await;
1071 let mut identity_stack = Vec::with_capacity(current_identity_stack.len() + 1);
1075 identity_stack.extend(current_identity_stack.iter().cloned());
1076 identity_stack.push(gate_identity);
1077 Ok(RootTurnAdmission {
1078 guard,
1079 identity_stack: identity_stack.into(),
1080 })
1081 }
1082
1083 pub async fn chat_with_actor_context(
1087 &self,
1088 input: &str,
1089 actor_context: crate::TurnActorContext,
1090 ) -> Result<AgentResponse> {
1091 self.chat_with_actor_context_boxed(input, actor_context)
1092 .await
1093 }
1094
1095 pub async fn chat_as_actor(&self, actor_id: &str, input: &str) -> Result<AgentResponse> {
1097 let actor_context = crate::TurnActorContext::new().with_origin_actor(actor_id);
1098 self.chat_with_actor_context(input, actor_context).await
1099 }
1100
1101 pub async fn load_actor_relationship(&self) -> Result<()> {
1103 self.maybe_load_actor_relationship().await;
1104 Ok(())
1105 }
1106
1107 pub async fn update_relationship_dimension(
1109 &self,
1110 dimension: &str,
1111 delta: f64,
1112 reason: Option<&str>,
1113 ) -> Result<ai_agents_relationships::DimensionChange> {
1114 self.update_relationship_dimension_for_perspective(
1115 ai_agents_relationships::RelationshipPerspective::AgentToActor,
1116 dimension,
1117 delta,
1118 reason,
1119 )
1120 .await
1121 }
1122
1123 pub async fn update_relationship_dimension_for_perspective(
1127 &self,
1128 perspective: ai_agents_relationships::RelationshipPerspective,
1129 dimension: &str,
1130 delta: f64,
1131 reason: Option<&str>,
1132 ) -> Result<ai_agents_relationships::DimensionChange> {
1133 let manager = self
1134 .relationship_manager
1135 .as_ref()
1136 .ok_or_else(|| AgentError::Config("Relationship memory is not configured".into()))?;
1137 let actor_id = self.effective_actor_id().ok_or_else(|| {
1138 AgentError::Config("No actor ID set. Use set_actor_id() first".into())
1139 })?;
1140 let change = manager.update_dimension_for_perspective(
1141 &actor_id,
1142 perspective,
1143 dimension,
1144 delta,
1145 1.0,
1146 reason.unwrap_or("manual relationship update"),
1147 )?;
1148 self.persist_actor_relationship(&actor_id).await?;
1149 info!(
1150 actor_id = %actor_id,
1151 perspective = %change.perspective,
1152 dimension = %change.dimension,
1153 delta = change.delta,
1154 current = change.current,
1155 "relationship updated manually"
1156 );
1157 self.hooks
1158 .on_relationship_change(&actor_id, std::slice::from_ref(&change))
1159 .await;
1160 Ok(change)
1161 }
1162
1163 pub fn reasoning_config(&self) -> &ReasoningConfig {
1164 &self.reasoning_config
1165 }
1166
1167 pub fn reflection_config(&self) -> &ReflectionConfig {
1168 &self.reflection_config
1169 }
1170
1171 pub fn with_facts_config(
1174 mut self,
1175 actor_memory_config: Option<ai_agents_facts::ActorMemoryConfig>,
1176 facts_config: Option<ai_agents_facts::FactsConfig>,
1177 ) -> Self {
1178 self.actor_memory_config = actor_memory_config;
1179 self.facts_config = facts_config;
1180 self
1181 }
1182
1183 pub fn with_facts(
1186 mut self,
1187 store: Arc<ai_agents_facts::FactStore>,
1188 extractor: Option<Arc<dyn ai_agents_facts::FactExtractor>>,
1189 actor_memory_config: Option<ai_agents_facts::ActorMemoryConfig>,
1190 facts_config: Option<ai_agents_facts::FactsConfig>,
1191 ) -> Self {
1192 *self.fact_store.write() = Some(store);
1193 *self.fact_extractor.write() = extractor;
1194 self.actor_memory_config = actor_memory_config;
1195 self.facts_config = facts_config;
1196 self
1197 }
1198
1199 pub fn fact_store(&self) -> Option<Arc<ai_agents_facts::FactStore>> {
1201 self.fact_store.read().clone()
1202 }
1203
1204 pub fn actor_id(&self) -> Option<String> {
1206 self.actor_id.read().clone()
1207 }
1208
1209 pub fn set_actor_id(&self, actor_id: &str) -> ai_agents_core::Result<()> {
1211 *self.actor_id.write() = Some(actor_id.to_string());
1212 {
1213 let mut meta = self.session_metadata.write();
1214 meta.actor_id = Some(actor_id.to_string());
1215 if !meta.actors.iter().any(|a| a == actor_id) {
1216 meta.actors.push(actor_id.to_string());
1217 }
1218 }
1219 Ok(())
1220 }
1221
1222 pub fn clear_actor_id(&self) {
1224 *self.actor_id.write() = None;
1225 self.session_metadata.write().actor_id = None;
1226 }
1227
1228 pub fn set_user_id(&self, user_id: &str) -> ai_agents_core::Result<()> {
1230 self.set_actor_id(user_id)
1231 }
1232
1233 pub async fn load_actor_memory(&self) -> ai_agents_core::Result<()> {
1235 let actor_id = match self.effective_actor_id() {
1236 Some(id) => id,
1237 None => return Ok(()),
1238 };
1239
1240 let store_opt = self.fact_store.read().clone();
1241 if let Some(store) = store_opt {
1242 let facts = store.get_facts(&actor_id).await?;
1243 let count = facts.len();
1244 self.actor_facts_cache
1245 .write()
1246 .insert(actor_id.clone(), facts);
1247 self.hooks.on_actor_memory_loaded(&actor_id, count).await;
1248 tracing::debug!("loaded {} facts for actor {}", count, actor_id);
1249 }
1250
1251 Ok(())
1252 }
1253
1254 async fn maybe_load_actor_memory(&self) {
1256 let Some(actor_id) = self.effective_actor_id() else {
1257 return;
1258 };
1259 if self.actor_facts_cache.read().contains_key(&actor_id) {
1260 return;
1261 }
1262 let _ = self.load_actor_memory().await;
1263 }
1264
1265 async fn pre_turn_session_lifecycle(&self) {
1267 if *self.redispatch_depth.read() > 0 {
1268 return;
1269 }
1270 self.resolve_actor_id_from_context();
1271 self.await_background_before_next_turn().await;
1272 self.record_session_actor_if_needed();
1273 self.maybe_load_actor_memory().await;
1274 self.maybe_load_actor_relationship().await;
1275 *self.messages_since_extraction.write() += 1;
1276 }
1277
1278 async fn post_turn_session_lifecycle(&self) -> Result<()> {
1280 if *self.redispatch_depth.read() > 0 {
1281 return Ok(());
1282 }
1283 *self.messages_since_extraction.write() += 1;
1284 self.run_post_turn_maintenance().await
1285 }
1286
1287 fn begin_root_turn(&self) {
1289 if *self.redispatch_depth.read() == 0 {
1290 let mut guard = self.active_turn_context.write();
1291 if guard.is_none() {
1292 self.root_user_message_committed
1293 .store(false, Ordering::SeqCst);
1294 self.active_native_exchanges.write().clear();
1295 let max_calls = self
1296 .runtime_config
1297 .optimization
1298 .max_speculative_llm_calls_per_turn;
1299 *guard = Some(TurnOptimizationContext::new(
1300 String::new(),
1301 HashMap::new(),
1302 max_calls,
1303 ));
1304 }
1305 }
1306 }
1307
1308 fn update_active_turn_context(
1309 &self,
1310 processed_input: &str,
1311 input_context: HashMap<String, Value>,
1312 ) {
1313 if *self.redispatch_depth.read() > 0 {
1314 return;
1315 }
1316 let max_calls = self
1317 .runtime_config
1318 .optimization
1319 .max_speculative_llm_calls_per_turn;
1320 let mut guard = self.active_turn_context.write();
1321 match guard.as_mut() {
1322 Some(context) => {
1323 context.processed_input = processed_input.to_string();
1324 context.input_context = input_context;
1325 context.max_speculative_llm_calls = max_calls;
1326 }
1327 None => {
1328 *guard = Some(TurnOptimizationContext::new(
1329 processed_input,
1330 input_context,
1331 max_calls,
1332 ));
1333 }
1334 }
1335 }
1336
1337 async fn commit_root_user_message(&self, processed_input: &str) -> Result<()> {
1339 if *self.redispatch_depth.read() > 0 {
1340 return Ok(());
1341 }
1342 if !self
1343 .root_user_message_committed
1344 .swap(true, Ordering::SeqCst)
1345 {
1346 self.memory
1347 .add_message(ChatMessage::user(processed_input))
1348 .await?;
1349 if let Some(context) = self.active_turn_context.write().as_mut() {
1350 context.mark_user_message_committed();
1351 }
1352 }
1353 Ok(())
1354 }
1355
1356 fn end_root_turn(&self) {
1358 if *self.redispatch_depth.read() == 0 {
1359 self.root_user_message_committed
1360 .store(false, Ordering::SeqCst);
1361 *self.active_turn_context.write() = None;
1362 self.active_native_exchanges.write().clear();
1363 }
1364 }
1365
1366 fn reserve_active_speculative_llm_call(&self, kind: RuntimeOptimizationKind) -> bool {
1367 self.begin_root_turn();
1368 let mut guard = self.active_turn_context.write();
1369 let Some(context) = guard.as_mut() else {
1370 return false;
1371 };
1372 context.reserve_speculative_llm_call_for(kind)
1373 }
1374
1375 fn branch_context_preview(&self) -> String {
1376 let context = self.build_context_with_overlays();
1377 let mut value = serde_json::to_string_pretty(&context).unwrap_or_else(|_| "{}".to_string());
1378 const MAX_CONTEXT_PREVIEW_CHARS: usize = 2048;
1379 if value.chars().count() > MAX_CONTEXT_PREVIEW_CHARS {
1380 value = value
1381 .chars()
1382 .take(MAX_CONTEXT_PREVIEW_CHARS)
1383 .collect::<String>();
1384 value.push_str("...");
1385 }
1386 value
1387 }
1388
1389 async fn await_background_before_next_turn(&self) {
1391 let optimization = &self.runtime_config.optimization;
1392 if !optimization.enabled {
1393 return;
1394 }
1395 let actor_id = self.effective_actor_id();
1396 let post = &optimization.post_turn;
1397 self.await_background_task(
1398 post.facts.await_before_next_turn,
1399 RuntimeTaskPurpose::PostTurnFacts,
1400 actor_id.as_deref(),
1401 "facts",
1402 )
1403 .await;
1404 self.await_background_task(
1405 post.relationships.await_before_next_turn,
1406 RuntimeTaskPurpose::PostTurnRelationship,
1407 actor_id.as_deref(),
1408 "relationships",
1409 )
1410 .await;
1411 }
1412
1413 async fn await_background_task(
1414 &self,
1415 policy: AwaitBeforeNextTurn,
1416 purpose: RuntimeTaskPurpose,
1417 actor_id: Option<&str>,
1418 label: &str,
1419 ) {
1420 match policy {
1421 AwaitBeforeNextTurn::Never => {}
1422 AwaitBeforeNextTurn::Always => {
1423 if let Err(error) = self.flush_background_tasks_for_purpose(purpose).await {
1424 warn!(label = label, error = %error, "background maintenance flush failed");
1425 }
1426 }
1427 AwaitBeforeNextTurn::SameActor => {
1428 if let Some(actor_id) = actor_id
1429 && let Err(error) = self
1430 .flush_background_tasks_for_actor_purpose(actor_id, purpose)
1431 .await
1432 {
1433 warn!(label = label, actor_id = %actor_id, error = %error, "actor background maintenance flush failed");
1434 }
1435 }
1436 }
1437 }
1438
1439 async fn run_post_turn_maintenance(&self) -> Result<()> {
1441 let optimization = &self.runtime_config.optimization;
1442 if !optimization.enabled {
1443 self.auto_extract_facts().await;
1444 self.auto_update_relationship().await;
1445 return Ok(());
1446 }
1447
1448 let facts_mode = effective_maintenance_mode(
1449 optimization.post_turn.facts.mode,
1450 optimization.parallel_post_turn_memory,
1451 );
1452 let relationships_mode = effective_maintenance_mode(
1453 optimization.post_turn.relationships.mode,
1454 optimization.parallel_post_turn_memory,
1455 );
1456
1457 match (facts_mode, relationships_mode) {
1458 (MaintenanceMode::InlineSerial, MaintenanceMode::InlineSerial) => {
1459 self.auto_extract_facts().await;
1460 self.auto_update_relationship().await;
1461 }
1462 (MaintenanceMode::InlineParallel, MaintenanceMode::InlineParallel) => {
1463 let facts = self.auto_extract_facts();
1464 let relationships = self.auto_update_relationship();
1465 tokio::join!(facts, relationships);
1466 }
1467 (MaintenanceMode::Background, MaintenanceMode::Background) => {
1468 self.schedule_facts_background().await?;
1469 self.schedule_relationship_background().await?;
1470 }
1471 (MaintenanceMode::Background, MaintenanceMode::InlineParallel)
1472 | (MaintenanceMode::Background, MaintenanceMode::InlineSerial) => {
1473 self.schedule_facts_background().await?;
1474 self.auto_update_relationship().await;
1475 }
1476 (MaintenanceMode::InlineParallel, MaintenanceMode::Background)
1477 | (MaintenanceMode::InlineSerial, MaintenanceMode::Background) => {
1478 self.auto_extract_facts().await;
1479 self.schedule_relationship_background().await?;
1480 }
1481 _ => {
1482 self.auto_extract_facts().await;
1483 self.auto_update_relationship().await;
1484 }
1485 }
1486 Ok(())
1487 }
1488
1489 async fn schedule_facts_background(&self) -> Result<()> {
1490 let policy = self.runtime_config.optimization.post_turn.facts.clone();
1491 let should_extract = self
1492 .facts_config
1493 .as_ref()
1494 .map(|c| c.enabled && c.auto_extract)
1495 .unwrap_or(false);
1496 if !should_extract {
1497 return Ok(());
1498 }
1499 let msgs_since = *self.messages_since_extraction.read();
1500 if msgs_since < 2 {
1501 return Ok(());
1502 }
1503 let Some(actor_id) = self.effective_actor_id() else {
1504 self.record_skipped_maintenance(
1505 "facts",
1506 ObservationPurpose::FactsExtraction,
1507 "missing_actor",
1508 Some(&policy),
1509 );
1510 return Ok(());
1511 };
1512 let Some(extractor) = self.fact_extractor.read().clone() else {
1513 return Ok(());
1514 };
1515 let messages = match self.memory.get_messages(None).await {
1516 Ok(messages) => messages,
1517 Err(error) => {
1518 warn!(error = %error, "failed to snapshot messages for fact extraction");
1519 return Ok(());
1520 }
1521 };
1522 let messages = Self::readable_native_messages(messages)?;
1523 let recent: Vec<_> = messages
1524 .iter()
1525 .rev()
1526 .take(msgs_since)
1527 .rev()
1528 .cloned()
1529 .collect();
1530 if recent.is_empty() {
1531 return Ok(());
1532 }
1533 let existing = self
1534 .actor_facts_cache
1535 .read()
1536 .get(&actor_id)
1537 .cloned()
1538 .unwrap_or_default();
1539 let categories = self
1540 .facts_config
1541 .as_ref()
1542 .map(|c| c.custom_categories.clone())
1543 .unwrap_or_default();
1544 let store = self.fact_store.read().clone();
1545 let cache = Arc::clone(&self.actor_facts_cache);
1546 let counter = Arc::clone(&self.messages_since_extraction);
1547 let hooks = Arc::clone(&self.hooks);
1548 let agent_id = self.info.id.clone();
1549 let observation = current_observation_context();
1550 let key = MaintenanceSequenceKey::actor(
1551 agent_id,
1552 actor_id.clone(),
1553 RuntimeTaskPurpose::PostTurnFacts,
1554 );
1555 let actor_for_task = actor_id.clone();
1556 let task = async move {
1557 let run = async move {
1558 let facts = extractor
1559 .extract(&recent, &existing, Some(&actor_for_task), &categories)
1560 .await?;
1561 if !facts.is_empty() {
1562 if let Some(store) = store {
1563 let authoritative = store.add_facts(&actor_for_task, facts.clone()).await?;
1564 cache.write().insert(actor_for_task.clone(), authoritative);
1565 } else {
1566 cache
1567 .write()
1568 .entry(actor_for_task.clone())
1569 .or_default()
1570 .extend(facts.clone());
1571 }
1572 {
1573 let mut count = counter.write();
1574 if *count <= msgs_since {
1575 *count = 0;
1576 } else {
1577 *count -= msgs_since;
1578 }
1579 }
1580 hooks.on_facts_extracted(&actor_for_task, &facts).await;
1581 }
1582 Ok(())
1583 };
1584 if let Some(context) = observation {
1585 with_observation_context(
1586 context.with_purpose(ObservationPurpose::FactsExtraction),
1587 run,
1588 )
1589 .await
1590 } else {
1591 run.await
1592 }
1593 };
1594 self.spawn_or_handle_background(Some(key), task, "facts", &policy)
1595 .await
1596 }
1597
1598 async fn schedule_relationship_background(&self) -> Result<()> {
1599 let policy = self
1600 .runtime_config
1601 .optimization
1602 .post_turn
1603 .relationships
1604 .clone();
1605 let Some(manager) = self.relationship_manager.as_ref().cloned() else {
1606 return Ok(());
1607 };
1608 let Some(actor_id) = self.effective_actor_id() else {
1609 self.record_skipped_maintenance(
1610 "relationships",
1611 ObservationPurpose::RelationshipUpdate,
1612 "missing_actor",
1613 Some(&policy),
1614 );
1615 return Ok(());
1616 };
1617 let recent_messages = manager.config().auto_update.recent_messages;
1618 let messages = match self.memory.get_messages(Some(recent_messages)).await {
1619 Ok(messages) => messages,
1620 Err(error) => {
1621 warn!(actor = %actor_id, error = %error, "failed to snapshot messages for relationship update");
1622 return Ok(());
1623 }
1624 };
1625 let messages = Self::readable_native_messages(messages)?;
1626 let storage = self.storage.read().clone();
1627 let hooks = Arc::clone(&self.hooks);
1628 let agent_id = self.info.id.clone();
1629 let observation = current_observation_context();
1630 let key = MaintenanceSequenceKey::actor(
1631 agent_id.clone(),
1632 actor_id.clone(),
1633 RuntimeTaskPurpose::PostTurnRelationship,
1634 );
1635 let actor_for_task = actor_id.clone();
1636 let task = async move {
1637 let run = async move {
1638 if manager.config().auto_update.enabled {
1639 let update = manager.auto_update(&actor_for_task, &messages).await?;
1640 if !update.changes.is_empty() {
1641 hooks
1642 .on_relationship_change(&actor_for_task, &update.changes)
1643 .await;
1644 }
1645 if let Some(ref event) = update.event {
1646 hooks.on_notable_event(&actor_for_task, event).await;
1647 }
1648 }
1649 if manager.config().persistence.enabled
1650 && let (Some(storage), Some(value)) =
1651 (storage, manager.relationship_as_value(&actor_for_task)?)
1652 {
1653 storage
1654 .save_relationship(&agent_id, &actor_for_task, &value)
1655 .await?;
1656 }
1657 Ok(())
1658 };
1659 if let Some(context) = observation {
1660 with_observation_context(
1661 context.with_purpose(ObservationPurpose::RelationshipUpdate),
1662 run,
1663 )
1664 .await
1665 } else {
1666 run.await
1667 }
1668 };
1669 self.spawn_or_handle_background(Some(key), task, "relationships", &policy)
1670 .await
1671 }
1672
1673 async fn spawn_or_handle_background<F>(
1675 &self,
1676 key: Option<MaintenanceSequenceKey>,
1677 task: F,
1678 label: &'static str,
1679 policy: &crate::optimization::config::MaintenanceTaskPolicy,
1680 ) -> Result<()>
1681 where
1682 F: Future<Output = Result<()>> + Send + 'static,
1683 {
1684 if self.background_maintenance.is_full() {
1685 match self
1686 .runtime_config
1687 .optimization
1688 .post_turn
1689 .on_background_overflow
1690 {
1691 BackgroundOverflowPolicy::RunInline => {
1692 record_background_maintenance_event(
1693 self.observability_manager.as_ref(),
1694 label,
1695 EventStatus::Success,
1696 0,
1697 "inline_overflow",
1698 None,
1699 Some(policy),
1700 );
1701 let start = Instant::now();
1702 match task.await {
1703 Ok(()) => record_background_maintenance_event(
1704 self.observability_manager.as_ref(),
1705 label,
1706 EventStatus::Success,
1707 start.elapsed().as_millis() as u64,
1708 "inline_completed",
1709 None,
1710 Some(policy),
1711 ),
1712 Err(error) => {
1713 warn!(label = label, error = %error, "inline maintenance fallback failed");
1714 record_background_maintenance_event(
1715 self.observability_manager.as_ref(),
1716 label,
1717 EventStatus::Error,
1718 start.elapsed().as_millis() as u64,
1719 "inline_failed",
1720 Some(error.to_string()),
1721 Some(policy),
1722 );
1723 return Err(error);
1724 }
1725 }
1726 }
1727 BackgroundOverflowPolicy::Drop => {
1728 self.record_skipped_maintenance(
1729 label,
1730 ObservationPurpose::Other(label.to_string()),
1731 "queue_full",
1732 Some(policy),
1733 );
1734 }
1735 BackgroundOverflowPolicy::Error => {
1736 record_background_maintenance_event(
1737 self.observability_manager.as_ref(),
1738 label,
1739 EventStatus::Error,
1740 0,
1741 "queue_full",
1742 None,
1743 Some(policy),
1744 );
1745 warn!(label = label, "background maintenance queue full");
1746 return Err(AgentError::Other(format!(
1747 "background maintenance queue is full for {}",
1748 label
1749 )));
1750 }
1751 }
1752 return Ok(());
1753 }
1754
1755 record_background_maintenance_event(
1756 self.observability_manager.as_ref(),
1757 label,
1758 EventStatus::Success,
1759 0,
1760 "scheduled",
1761 None,
1762 Some(policy),
1763 );
1764 let manager = self.observability_manager.clone();
1765 let policy_for_task = policy.clone();
1766 let observed_task = async move {
1767 let start = Instant::now();
1768 let result = task.await;
1769 match &result {
1770 Ok(()) => record_background_maintenance_event(
1771 manager.as_ref(),
1772 label,
1773 EventStatus::Success,
1774 start.elapsed().as_millis() as u64,
1775 "completed",
1776 None,
1777 Some(&policy_for_task),
1778 ),
1779 Err(error) => record_background_maintenance_event(
1780 manager.as_ref(),
1781 label,
1782 EventStatus::Error,
1783 start.elapsed().as_millis() as u64,
1784 "failed",
1785 Some(error.to_string()),
1786 Some(&policy_for_task),
1787 ),
1788 }
1789 result
1790 };
1791
1792 if let Err(error) = self.background_maintenance.spawn(key, observed_task) {
1793 record_background_maintenance_event(
1794 self.observability_manager.as_ref(),
1795 label,
1796 EventStatus::Error,
1797 0,
1798 "spawn_failed",
1799 Some(error.to_string()),
1800 Some(policy),
1801 );
1802 warn!(label = label, error = %error, "background maintenance spawn failed");
1803 return Err(error);
1804 }
1805 Ok(())
1806 }
1807
1808 fn record_skipped_maintenance(
1810 &self,
1811 label: &str,
1812 purpose: ObservationPurpose,
1813 reason: &str,
1814 policy: Option<&crate::optimization::config::MaintenanceTaskPolicy>,
1815 ) {
1816 if let Some(manager) = self.observability_manager.as_ref() {
1817 let mut tags = background_maintenance_tags(label, "skipped", Some(reason), policy);
1818 tags.insert("runtime.skip_reason".to_string(), reason.to_string());
1819 manager.record_lifecycle_event(
1820 EventType::MemoryOperation {
1821 operation: format!("{}_maintenance", label),
1822 },
1823 purpose,
1824 EventStatus::Skipped,
1825 0,
1826 tags,
1827 None,
1828 );
1829 }
1830 }
1831
1832 pub fn actor_facts(&self) -> Vec<ai_agents_core::KeyFact> {
1834 let Some(actor_id) = self.effective_actor_id() else {
1835 return Vec::new();
1836 };
1837 self.actor_facts_cache
1838 .read()
1839 .get(&actor_id)
1840 .cloned()
1841 .unwrap_or_default()
1842 }
1843
1844 pub fn relationship_memory_text(&self) -> Option<String> {
1846 self.format_relationship_for_context().map(|(_, text)| text)
1847 }
1848
1849 pub async fn extract_facts(
1851 &self,
1852 last_n: usize,
1853 ) -> ai_agents_core::Result<Vec<ai_agents_core::KeyFact>> {
1854 self.extract_facts_with_source(last_n, "manual").await
1855 }
1856
1857 async fn extract_facts_with_source(
1858 &self,
1859 last_n: usize,
1860 source: &'static str,
1861 ) -> ai_agents_core::Result<Vec<ai_agents_core::KeyFact>> {
1862 let extractor = match self.fact_extractor.read().clone() {
1863 Some(e) => e,
1864 None => return Ok(vec![]),
1865 };
1866
1867 let messages = Self::readable_native_messages(self.memory.get_messages(None).await?)?;
1868 let recent: Vec<_> = messages.iter().rev().take(last_n).rev().cloned().collect();
1869
1870 if recent.is_empty() {
1871 return Ok(vec![]);
1872 }
1873
1874 let actor_id = self.effective_actor_id();
1875 let existing = actor_id
1876 .as_ref()
1877 .and_then(|aid| self.actor_facts_cache.read().get(aid).cloned())
1878 .unwrap_or_default();
1879
1880 let categories = self
1881 .facts_config
1882 .as_ref()
1883 .map(|c| c.custom_categories.clone())
1884 .unwrap_or_default();
1885
1886 let facts = self
1887 .observe_purpose(
1888 ObservationPurpose::FactsExtraction,
1889 extractor.extract(&recent, &existing, actor_id.as_deref(), &categories),
1890 )
1891 .await?;
1892
1893 if !facts.is_empty() {
1895 let fact_store_opt = self.fact_store.read().clone();
1896 let mut stored_total = 0usize;
1897 let mut cache_updated = false;
1898 if let (Some(store), Some(aid)) = (fact_store_opt, &actor_id) {
1899 let authoritative = store.add_facts(aid, facts.clone()).await?;
1901 stored_total = authoritative.len();
1902 self.actor_facts_cache
1903 .write()
1904 .insert(aid.clone(), authoritative);
1905 cache_updated = true;
1906 } else if let Some(aid) = &actor_id {
1907 let mut cache = self.actor_facts_cache.write();
1908 let entry = cache.entry(aid.clone()).or_default();
1909 entry.extend(facts.clone());
1910 stored_total = entry.len();
1911 cache_updated = true;
1912 }
1913
1914 info!(
1915 actor_id = %actor_id.as_deref().unwrap_or("<none>"),
1916 source = source,
1917 requested_messages = last_n,
1918 message_count = recent.len(),
1919 extracted_count = facts.len(),
1920 cache_updated = cache_updated,
1921 stored_total = stored_total,
1922 "facts extracted"
1923 );
1924
1925 if let Some(ref aid) = actor_id {
1926 self.hooks.on_facts_extracted(aid, &facts).await;
1927 }
1928 }
1929
1930 Ok(facts)
1931 }
1932
1933 fn resolve_actor_id_from_context(&self) {
1936 if self
1937 .current_turn_actor_context()
1938 .and_then(|ctx| ctx.effective_actor_id().map(str::to_string))
1939 .is_some()
1940 {
1941 return;
1942 }
1943
1944 if let Some(ref am_config) = self.actor_memory_config
1945 && am_config.identification.method == ai_agents_facts::IdentificationMethod::FromContext
1946 && let Some(ref path) = am_config.identification.context_path
1947 {
1948 let val = self
1950 .context_manager
1951 .get_path(path)
1952 .or_else(|| self.context_manager.get(path));
1953 if let Some(val) = val
1954 && let Some(id_str) = val.as_str()
1955 {
1956 let current = self.actor_id.read().clone();
1957 if current.as_deref() != Some(id_str) {
1958 *self.actor_id.write() = Some(id_str.to_string());
1959 let mut meta = self.session_metadata.write();
1960 meta.actor_id = Some(id_str.to_string());
1961 if !meta.actors.iter().any(|a| a == id_str) {
1962 meta.actors.push(id_str.to_string());
1963 }
1964 }
1965 }
1966 }
1967 }
1968
1969 fn format_actor_facts_for_context(&self) -> String {
1971 let should_inject = self
1973 .facts_config
1974 .as_ref()
1975 .map(|c| c.inject_in_context)
1976 .unwrap_or(true);
1977 if !should_inject {
1978 return String::new();
1979 }
1980
1981 let Some(actor_id) = self.effective_actor_id() else {
1982 return String::new();
1983 };
1984
1985 let facts = self
1986 .actor_facts_cache
1987 .read()
1988 .get(&actor_id)
1989 .cloned()
1990 .unwrap_or_default();
1991 if facts.is_empty() {
1992 return String::new();
1993 }
1994
1995 let am_config = self.actor_memory_config.as_ref();
1996 let facts_budget = self
1999 .memory_token_budget
2000 .as_ref()
2001 .map(|b| b.allocation.facts as usize)
2002 .filter(|n| *n > 0);
2003 let default_max = am_config.map(|c| c.injection.max_tokens).unwrap_or(800);
2004 let max_tokens = facts_budget.unwrap_or(default_max);
2005
2006 let filtered: Vec<ai_agents_core::KeyFact> = if let Some(cfg) = am_config {
2008 if cfg.injection.mode == ai_agents_facts::InjectionMode::OnDemand {
2009 return String::new();
2010 }
2011 if cfg.injection.mode == ai_agents_facts::InjectionMode::Category
2012 && !cfg.injection.categories.is_empty()
2013 {
2014 facts
2015 .iter()
2016 .filter(|f| {
2017 cfg.injection
2018 .categories
2019 .iter()
2020 .any(|c| f.category.to_string() == *c)
2021 })
2022 .cloned()
2023 .collect()
2024 } else {
2025 facts.clone()
2026 }
2027 } else {
2028 facts.clone()
2029 };
2030
2031 if filtered.is_empty() {
2032 return String::new();
2033 }
2034
2035 if let Some(store) = self.fact_store.read().clone() {
2036 store.format_for_context(&filtered, max_tokens)
2037 } else {
2038 String::new()
2039 }
2040 }
2041
2042 fn build_context_with_staged(&self, staged: &HashMap<String, Value>) -> HashMap<String, Value> {
2043 let context = self.build_context_with_overlays();
2044 let mut root = Value::Object(context.into_iter().collect());
2045 for (path, value) in staged {
2046 if let Ok(updated) = ai_agents_core::set_dot_path(root.clone(), path, value.clone()) {
2047 root = updated;
2048 }
2049 }
2050 match root {
2051 Value::Object(obj) => obj.into_iter().collect(),
2052 _ => HashMap::new(),
2053 }
2054 }
2055
2056 fn build_context_with_overlays(&self) -> HashMap<String, Value> {
2057 let mut context = self.context_manager.get_all();
2058 let mut root = Value::Object(context.clone().into_iter().collect());
2059
2060 if let Some(turn_ctx) = self.current_turn_actor_context() {
2061 if let Some(ref origin_actor_id) = turn_ctx.origin_actor_id
2062 && let Ok(updated) = ai_agents_core::set_dot_path(
2063 root.clone(),
2064 "interaction.origin_actor_id",
2065 serde_json::json!(origin_actor_id),
2066 )
2067 {
2068 root = updated;
2069 }
2070 if let Some(ref sender_agent_id) = turn_ctx.sender_agent_id
2071 && let Ok(updated) = ai_agents_core::set_dot_path(
2072 root.clone(),
2073 "interaction.sender_agent_id",
2074 serde_json::json!(sender_agent_id),
2075 )
2076 {
2077 root = updated;
2078 }
2079 }
2080
2081 if let Some(ref actor_id) = self.effective_actor_id()
2082 && let Ok(updated) = ai_agents_core::set_dot_path(
2083 root.clone(),
2084 "interaction.actor_id",
2085 serde_json::json!(actor_id),
2086 )
2087 {
2088 root = updated;
2089 }
2090
2091 if let Some(manager) = self.relationship_manager.as_ref()
2092 && let Some(actor_id) = self.effective_actor_id()
2093 && let Some(value) = manager.to_context_value(&actor_id)
2094 && let Ok(updated) = ai_agents_core::set_dot_path(
2095 root.clone(),
2096 &manager.config().injection.context_path,
2097 value,
2098 )
2099 {
2100 root = updated;
2101 }
2102
2103 if let Value::Object(obj) = root {
2104 context = obj.into_iter().collect();
2105 }
2106
2107 context
2108 }
2109
2110 fn resolve_actor_name_from_context(&self) -> Option<String> {
2111 for path in ["actor.name", "user.name", "player.name", "customer.name"] {
2112 if let Some(value) = self.context_manager.get_path(path)
2113 && let Some(name) = value.as_str()
2114 {
2115 return Some(name.to_string());
2116 }
2117 }
2118 None
2119 }
2120
2121 async fn maybe_load_actor_relationship(&self) {
2122 let Some(manager) = self.relationship_manager.as_ref() else {
2123 return;
2124 };
2125 let Some(actor_id) = self.effective_actor_id() else {
2126 return;
2127 };
2128
2129 let mut should_fire_loaded = false;
2130 if manager.get(&actor_id).is_none() {
2131 let mut loaded = false;
2132 if manager.config().persistence.enabled {
2133 let storage = self.storage.read().clone();
2134 if let Some(storage) = storage {
2135 match storage.load_relationship(&self.info.id, &actor_id).await {
2136 Ok(Some(value)) => match manager.insert_from_value(value) {
2137 Ok(_) => loaded = true,
2138 Err(e) => {
2139 warn!(actor = %actor_id, error = %e, "failed to restore relationship")
2140 }
2141 },
2142 Ok(None) => {}
2143 Err(e) => {
2144 warn!(actor = %actor_id, error = %e, "failed to load relationship")
2145 }
2146 }
2147 }
2148 }
2149
2150 if !loaded {
2151 manager.get_or_create(&actor_id, self.resolve_actor_name_from_context().as_deref());
2152 }
2153 should_fire_loaded = true;
2154 }
2155
2156 let actor_name = self.resolve_actor_name_from_context();
2157 let relationship = manager.touch_interaction(&actor_id, actor_name.as_deref());
2158 if should_fire_loaded {
2159 self.hooks
2160 .on_relationship_loaded(&actor_id, &relationship)
2161 .await;
2162 }
2163 }
2164
2165 fn format_relationship_for_context(&self) -> Option<(String, String)> {
2166 let manager = self.relationship_manager.as_ref()?;
2167 if !manager.config().injection.enabled {
2168 return None;
2169 }
2170 let actor_id = self.effective_actor_id()?;
2171 let relationship = manager.get(&actor_id)?;
2172 let local_cap = manager.config().injection.max_tokens;
2173 let global_cap = self
2174 .memory_token_budget
2175 .as_ref()
2176 .map(|b| b.allocation.relationships as usize)
2177 .filter(|n| *n > 0);
2178 let max_tokens = global_cap.map(|g| g.min(local_cap)).unwrap_or(local_cap);
2179 let text = ai_agents_relationships::format_relationship(
2180 &relationship,
2181 &manager.config().injection.format,
2182 max_tokens,
2183 );
2184 if text.is_empty() {
2185 None
2186 } else {
2187 Some((manager.config().injection.prompt_variable.clone(), text))
2188 }
2189 }
2190
2191 async fn persist_actor_relationship(&self, actor_id: &str) -> Result<()> {
2192 let Some(manager) = self.relationship_manager.as_ref() else {
2193 return Ok(());
2194 };
2195 if !manager.config().persistence.enabled {
2196 return Ok(());
2197 }
2198 let storage = self.storage.read().clone();
2199 let Some(storage) = storage else {
2200 return Ok(());
2201 };
2202 if let Some(value) = manager.relationship_as_value(actor_id)? {
2203 storage
2204 .save_relationship(&self.info.id, actor_id, &value)
2205 .await?;
2206 }
2207 Ok(())
2208 }
2209
2210 async fn auto_update_relationship(&self) {
2211 let Some(manager) = self.relationship_manager.as_ref() else {
2212 return;
2213 };
2214 let Some(actor_id) = self.effective_actor_id() else {
2215 return;
2216 };
2217 if !manager.config().auto_update.enabled {
2218 let _ = self.persist_actor_relationship(&actor_id).await;
2219 return;
2220 }
2221
2222 let recent_messages = manager.config().auto_update.recent_messages;
2223 let messages = match self.memory.get_messages(Some(recent_messages)).await {
2224 Ok(messages) => messages,
2225 Err(e) => {
2226 warn!(actor = %actor_id, error = %e, "failed to read messages for relationship update");
2227 return;
2228 }
2229 };
2230 let messages = match Self::readable_native_messages(messages) {
2231 Ok(messages) => messages,
2232 Err(error) => {
2233 warn!(actor = %actor_id, error = %error, "failed to project native history for relationship update");
2234 return;
2235 }
2236 };
2237
2238 match self
2239 .observe_purpose(
2240 ObservationPurpose::RelationshipUpdate,
2241 manager.auto_update(&actor_id, &messages),
2242 )
2243 .await
2244 {
2245 Ok(update) => {
2246 if !update.changes.is_empty() {
2247 self.hooks
2248 .on_relationship_change(&actor_id, &update.changes)
2249 .await;
2250 }
2251 if let Some(ref event) = update.event {
2252 self.hooks.on_notable_event(&actor_id, event).await;
2253 }
2254 let persisted = match self.persist_actor_relationship(&actor_id).await {
2255 Ok(()) => true,
2256 Err(e) => {
2257 warn!(actor = %actor_id, error = %e, "failed to persist relationship");
2258 false
2259 }
2260 };
2261 if !update.changes.is_empty() || update.event.is_some() {
2262 let changed_dimensions: Vec<String> = update
2263 .changes
2264 .iter()
2265 .map(|change| format!("{}:{}", change.perspective, change.dimension))
2266 .collect();
2267 info!(
2268 actor_id = %actor_id,
2269 change_count = update.changes.len(),
2270 changed_dimensions = ?changed_dimensions,
2271 event_present = update.event.is_some(),
2272 persisted = persisted,
2273 "relationship updated"
2274 );
2275 } else {
2276 debug!(actor_id = %actor_id, persisted = persisted, "relationship evaluation ran but found no changes");
2277 }
2278 }
2279 Err(e) => warn!(actor = %actor_id, error = %e, "relationship update failed"),
2280 }
2281 }
2282
2283 async fn auto_extract_facts(&self) {
2285 let should_extract = self
2286 .facts_config
2287 .as_ref()
2288 .map(|c| c.enabled && c.auto_extract)
2289 .unwrap_or(false);
2290
2291 if !should_extract {
2292 debug!("fact extraction skipped because auto extraction is disabled");
2293 return;
2294 }
2295
2296 let msgs_since = *self.messages_since_extraction.read();
2297 if msgs_since < 2 {
2298 debug!(
2299 messages_since_extraction = msgs_since,
2300 "fact extraction skipped until threshold is reached"
2301 );
2302 return;
2303 }
2304
2305 match self.extract_facts_with_source(msgs_since, "auto").await {
2306 Ok(facts) => {
2307 if !facts.is_empty() {
2308 *self.messages_since_extraction.write() = 0;
2309 } else {
2310 debug!("fact extraction ran but found no new facts");
2311 }
2312 }
2313 Err(e) => {
2314 warn!("fact extraction failed: {}", e);
2315 }
2316 }
2317 }
2318
2319 pub fn with_persona(mut self, manager: Arc<ai_agents_persona::PersonaManager>) -> Self {
2320 self.persona_manager = Some(manager);
2321 self
2322 }
2323
2324 pub fn persona_manager(&self) -> Option<&Arc<ai_agents_persona::PersonaManager>> {
2325 self.persona_manager.as_ref()
2326 }
2327
2328 pub fn with_disambiguation(mut self, config: DisambiguationConfig) -> Self {
2329 if config.is_enabled() {
2330 let manager = DisambiguationManager::new(config, Arc::clone(&self.llm_registry))
2331 .with_clarification_observer(Arc::new(ObservabilityClarificationObserver));
2332 self.disambiguation_manager = Some(manager);
2333 }
2334 self
2335 }
2336
2337 pub fn disambiguation_manager(&self) -> Option<&DisambiguationManager> {
2338 self.disambiguation_manager.as_ref()
2339 }
2340
2341 pub fn has_disambiguation(&self) -> bool {
2342 self.disambiguation_manager
2343 .as_ref()
2344 .is_some_and(|m| m.is_enabled())
2345 }
2346
2347 pub async fn init_storage(&self) -> Result<()> {
2348 let _guard = self.storage_init.lock().await;
2352 let mut storage = self.storage.read().clone();
2353 if storage.is_none() && !self.storage_config.is_none() {
2354 let storage_config = self.convert_storage_config();
2355 storage = create_storage(&storage_config).await?;
2356 *self.storage.write() = storage.clone();
2357 }
2358
2359 self.validate_storage_requirements(storage.as_deref())?;
2360 self.complete_facts_init().await;
2361 Ok(())
2362 }
2363
2364 fn validate_storage_requirements(&self, storage: Option<&dyn AgentStorage>) -> Result<()> {
2365 let facts_required = self
2366 .facts_config
2367 .as_ref()
2368 .is_some_and(|config| config.enabled)
2369 || self
2370 .actor_memory_config
2371 .as_ref()
2372 .is_some_and(|config| config.enabled);
2373 let relationships_required = self
2374 .relationship_manager
2375 .as_ref()
2376 .is_some_and(|manager| manager.config().persistence.enabled);
2377
2378 let Some(storage) = storage else {
2379 let mut requirements = Vec::new();
2380 if facts_required {
2381 requirements.push("actor facts or actor memory");
2382 }
2383 if relationships_required {
2384 requirements.push("persistent relationships");
2385 }
2386 if requirements.is_empty() {
2387 return Ok(());
2388 }
2389 return Err(AgentError::Config(format!(
2390 "Storage is required for enabled {} but none is configured or injected",
2391 requirements.join(" and ")
2392 )));
2393 };
2394
2395 if facts_required && !storage.supports(StorageCapability::ActorFacts) {
2399 return Err(AgentError::UnsupportedStorageCapability(
2400 StorageCapability::ActorFacts,
2401 ));
2402 }
2403 if relationships_required && !storage.supports(StorageCapability::ActorRelationships) {
2404 return Err(AgentError::UnsupportedStorageCapability(
2405 StorageCapability::ActorRelationships,
2406 ));
2407 }
2408 Ok(())
2409 }
2410
2411 async fn complete_facts_init(&self) {
2414 if self.fact_store.read().is_some() {
2415 return;
2416 }
2417 let storage = match self.storage.read().clone() {
2418 Some(s) => s,
2419 None => return,
2420 };
2421
2422 let facts_enabled = self
2423 .facts_config
2424 .as_ref()
2425 .map(|f| f.enabled)
2426 .unwrap_or(false);
2427 let actor_memory_enabled = self
2428 .actor_memory_config
2429 .as_ref()
2430 .map(|a| a.enabled)
2431 .unwrap_or(false);
2432
2433 if !facts_enabled && !actor_memory_enabled {
2434 return;
2435 }
2436
2437 let fc = self.facts_config.clone().unwrap_or_default();
2438 let store = Arc::new(ai_agents_facts::FactStore::new(
2439 storage,
2440 self.info.id.clone(),
2441 fc.clone(),
2442 ));
2443
2444 let extractor: Option<Arc<dyn ai_agents_facts::FactExtractor>> = if facts_enabled {
2445 let extractor_llm = fc
2446 .extractor_llm
2447 .as_ref()
2448 .and_then(|alias| self.llm_registry.get(alias).ok())
2449 .or_else(|| self.llm_registry.router().ok())
2450 .or_else(|| self.llm_registry.default().ok());
2451 extractor_llm.map(|llm| {
2452 Arc::new(ai_agents_facts::LLMFactExtractor::new(llm, fc.clone()))
2453 as Arc<dyn ai_agents_facts::FactExtractor>
2454 })
2455 } else {
2456 None
2457 };
2458
2459 *self.fact_store.write() = Some(store);
2460 *self.fact_extractor.write() = extractor;
2461 debug!(
2462 agent = %self.info.id,
2463 facts_enabled,
2464 actor_memory_enabled,
2465 "facts storage initialized"
2466 );
2467 }
2468
2469 fn convert_storage_config(&self) -> StorageStorageConfig {
2470 crate::spec::storage::to_storage_config(&self.storage_config)
2471 }
2472
2473 pub fn storage(&self) -> Option<Arc<dyn AgentStorage>> {
2474 self.storage.read().clone()
2475 }
2476
2477 pub fn storage_config(&self) -> &StorageConfig {
2478 &self.storage_config
2479 }
2480
2481 pub fn spawner(&self) -> Option<&Arc<crate::spawner::AgentSpawner>> {
2483 self.spawner.as_ref()
2484 }
2485
2486 pub fn spawner_registry(&self) -> Option<&Arc<crate::spawner::AgentRegistry>> {
2488 self.spawner_registry.as_ref()
2489 }
2490
2491 pub fn has_spawner(&self) -> bool {
2492 self.spawner_registry.is_some()
2493 }
2494
2495 pub fn with_spawner_handles(
2496 mut self,
2497 spawner: Arc<crate::spawner::AgentSpawner>,
2498 registry: Arc<crate::spawner::AgentRegistry>,
2499 ) -> Self {
2500 self.spawner = Some(spawner);
2501 self.spawner_registry = Some(registry);
2502 self
2503 }
2504
2505 pub fn with_hooks(mut self, hooks: Arc<dyn AgentHooks>) -> Self {
2506 self.hooks = hooks;
2507 self
2508 }
2509
2510 pub fn with_parallel_tools(mut self, config: ParallelToolsConfig) -> Self {
2511 self.parallel_tools = config;
2512 self
2513 }
2514
2515 pub fn with_streaming(mut self, config: StreamingConfig) -> Self {
2516 self.streaming = config;
2517 self
2518 }
2519
2520 pub fn with_hitl(mut self, engine: HITLEngine, handler: Arc<dyn ApprovalHandler>) -> Self {
2521 self.hitl_engine = Some(engine);
2522 self.approval_handler = handler;
2523 self
2524 }
2525
2526 pub fn with_max_context_tokens(mut self, tokens: u32) -> Self {
2527 self.max_context_tokens = tokens;
2528 self
2529 }
2530
2531 pub fn with_memory_token_budget(mut self, budget: MemoryTokenBudget) -> Self {
2532 self.memory_token_budget = Some(budget);
2533 self
2534 }
2535
2536 pub fn with_recovery_manager(mut self, manager: RecoveryManager) -> Self {
2537 self.recovery_manager = manager;
2538 self
2539 }
2540
2541 pub fn with_tool_security(mut self, engine: ToolSecurityEngine) -> Self {
2542 self.tool_security = engine;
2543 self
2544 }
2545
2546 pub fn runtime_control(&self) -> RuntimeControlHandle {
2548 RuntimeControlHandle {
2549 state: Arc::clone(&self.runtime_control),
2550 }
2551 }
2552
2553 pub fn set_question_handler(&self, handler: Option<Arc<dyn QuestionHandler>>) {
2555 self.tools.set_question_handler(handler);
2556 }
2557
2558 pub fn set_diagnostics_provider(&self, provider: Arc<dyn DiagnosticsProvider>) {
2560 self.tools.set_diagnostics_provider(provider);
2561 }
2562
2563 pub fn set_command_runner(&self, runner: Arc<dyn CommandRunner>) {
2565 self.tools.set_command_runner(runner);
2566 }
2567
2568 pub fn set_web_search_provider(&self, provider: Arc<dyn ai_agents_tools::WebSearchProvider>) {
2570 self.tools.set_web_search_provider(provider);
2571 }
2572
2573 pub fn todos(&self) -> Vec<TodoItem> {
2575 self.tools.todos()
2576 }
2577
2578 fn active_tool_security(&self) -> ToolSecurityEngine {
2580 self.runtime_control
2581 .tool_security_override
2582 .read()
2583 .clone()
2584 .unwrap_or_else(|| self.tool_security.clone())
2585 }
2586
2587 fn runtime_safety_snapshot(&self) -> RuntimeSafetySnapshot {
2589 let _guard = self.runtime_control.snapshot_guard.read();
2590 RuntimeSafetySnapshot {
2591 version: self.runtime_control.version.load(Ordering::SeqCst),
2592 emergency_deny: self.runtime_control.emergency_deny.load(Ordering::SeqCst),
2593 tool_security: self
2594 .runtime_control
2595 .tool_security_override
2596 .read()
2597 .clone()
2598 .unwrap_or_else(|| self.tool_security.clone()),
2599 tool_scope_override: self.runtime_control.tool_scope_override.read().clone(),
2600 }
2601 }
2602
2603 fn admit_tool_execution(
2605 &self,
2606 expected_runtime_version: u64,
2607 expected_policy_version: u64,
2608 expected_state_generation: Option<u64>,
2609 canonical_id: &str,
2610 ) -> SecurityCheckResult {
2611 let _guard = self.runtime_control.snapshot_guard.read();
2612 if self.runtime_control.emergency_deny.load(Ordering::SeqCst) {
2613 return SecurityCheckResult::Block {
2614 reason: "runtime emergency deny is enabled".to_string(),
2615 };
2616 }
2617 let runtime_version = self.runtime_control.version.load(Ordering::SeqCst);
2618 let security_engine = self
2619 .runtime_control
2620 .tool_security_override
2621 .read()
2622 .clone()
2623 .unwrap_or_else(|| self.tool_security.clone());
2624 if runtime_version != expected_runtime_version
2625 || security_engine.policy_version() != expected_policy_version
2626 {
2627 return SecurityCheckResult::Block {
2628 reason: "runtime safety controls changed before admission".to_string(),
2629 };
2630 }
2631 let current_state_generation = self
2632 .state_machine
2633 .as_ref()
2634 .map(|state_machine| state_machine.generation());
2635 if current_state_generation != expected_state_generation {
2636 return SecurityCheckResult::Block {
2637 reason: "state scope changed before admission".to_string(),
2638 };
2639 }
2640 security_engine.admit_tool_execution(canonical_id)
2641 }
2642
2643 pub fn with_process_processor(mut self, processor: ProcessProcessor) -> Self {
2644 let processor = processor.with_stage_observer(Arc::new(ObservabilityProcessStageObserver));
2645 self.process_processor = Some(processor);
2646 self
2647 }
2648
2649 pub fn with_state_machine(
2650 mut self,
2651 state_machine: Arc<StateMachine>,
2652 evaluator: Arc<dyn TransitionEvaluator>,
2653 ) -> Self {
2654 self.state_machine = Some(state_machine);
2655 self.transition_evaluator = Some(evaluator);
2656 self
2657 }
2658
2659 pub fn with_context_manager(mut self, manager: Arc<ContextManager>) -> Self {
2660 self.context_manager = manager;
2661 self
2662 }
2663
2664 pub fn register_message_filter(&self, name: impl Into<String>, filter: Arc<dyn MessageFilter>) {
2665 self.message_filters.write().insert(name.into(), filter);
2666 }
2667
2668 pub fn set_context(&self, key: &str, value: Value) -> Result<()> {
2669 self.context_manager.update(key, value)
2670 }
2671
2672 pub fn update_context(&self, path: &str, value: Value) -> Result<()> {
2673 self.context_manager.update(path, value)
2674 }
2675
2676 pub fn get_context(&self) -> HashMap<String, Value> {
2677 self.build_context_with_overlays()
2678 }
2679
2680 pub fn remove_context(&self, key: &str) -> Option<Value> {
2681 self.context_manager.remove(key)
2682 }
2683
2684 pub async fn refresh_context(&self, key: &str) -> Result<()> {
2685 self.context_manager.refresh(key).await
2686 }
2687
2688 pub fn register_context_provider(&self, name: &str, provider: Arc<dyn ContextProvider>) {
2689 self.context_manager.register_provider(name, provider);
2690 }
2691
2692 pub fn current_state(&self) -> Option<String> {
2693 self.state_machine.as_ref().map(|sm| sm.current())
2694 }
2695
2696 async fn invalidate_pending_confirmation(&self, reason: &'static str) {
2698 self.disambiguation_epoch.fetch_add(1, Ordering::SeqCst);
2699 let Some(disambiguator) = self.disambiguation_manager.as_ref() else {
2700 return;
2701 };
2702 if disambiguator.has_pending_confirmation().await {
2703 disambiguator.clear_pending().await;
2704 *self.pending_skill_id.write() = None;
2705 info!(
2706 confirmation_event = "invalidated",
2707 invalidation_reason = reason,
2708 "Runtime invalidated pending confirmation"
2709 );
2710 }
2711 }
2712
2713 async fn admit_disambiguation_redispatch(
2715 &self,
2716 expected_epoch: u64,
2717 expected_state_generation: Option<u64>,
2718 ) -> Result<tokio::sync::RwLockReadGuard<'_, ()>> {
2719 let admission = self.disambiguation_admission.read().await;
2720 let state_generation = self
2721 .state_machine
2722 .as_ref()
2723 .map(|state_machine| state_machine.generation());
2724 if self.disambiguation_epoch.load(Ordering::SeqCst) != expected_epoch
2725 || state_generation != expected_state_generation
2726 {
2727 return Err(AgentError::Other(
2728 "Disambiguation ownership changed before redispatch admission".to_string(),
2729 ));
2730 }
2731 Ok(admission)
2732 }
2733
2734 fn reserve_state_transition(&self) -> Option<StateTransitionReservation<'_>> {
2736 self.state_transition_reserved
2737 .compare_exchange(false, true, Ordering::SeqCst, Ordering::SeqCst)
2738 .ok()
2739 .map(|_| StateTransitionReservation {
2740 reserved: &self.state_transition_reserved,
2741 })
2742 }
2743
2744 async fn admit_optional_disambiguation_ownership(
2746 &self,
2747 ownership: Option<DisambiguationOwnership>,
2748 ) -> Result<Option<tokio::sync::RwLockReadGuard<'_, ()>>> {
2749 match ownership {
2750 Some(ownership) => self
2751 .admit_disambiguation_redispatch(ownership.epoch, ownership.state_generation)
2752 .await
2753 .map(Some),
2754 None => Ok(None),
2755 }
2756 }
2757
2758 pub async fn transition_to(&self, state: &str) -> Result<()> {
2760 let Some(ref sm) = self.state_machine else {
2761 return Ok(());
2762 };
2763 let claim_admission = self.disambiguation_admission.write().await;
2764 let reservation = self.reserve_state_transition().ok_or_else(|| {
2765 AgentError::Other("Another state transition is already in progress".to_string())
2766 })?;
2767 let from_state = sm.current();
2768 let expected_state_generation = sm.generation();
2769 let expected_disambiguation_epoch = self.disambiguation_epoch.load(Ordering::SeqCst);
2770 let history_before = sm.history();
2771 drop(claim_admission);
2772
2773 self.execute_state_exit_actions(&from_state).await;
2774
2775 let admission = self.disambiguation_admission.write().await;
2776 if sm.current() != from_state
2777 || sm.generation() != expected_state_generation
2778 || self.disambiguation_epoch.load(Ordering::SeqCst) != expected_disambiguation_epoch
2779 {
2780 return Err(AgentError::Other(
2781 "State ownership changed during manual transition preparation".to_string(),
2782 ));
2783 }
2784 sm.transition_to(state, "manual transition")?;
2785 self.invalidate_pending_confirmation("state_transition")
2786 .await;
2787 let entered = sm.current();
2788 let is_reentry = Self::state_was_previously_entered(&entered, &from_state, &history_before);
2789 drop(admission);
2790
2791 self.execute_state_enter_actions(&entered, is_reentry).await;
2792 drop(reservation);
2793 info!(to = %entered, "Manual state transition");
2794 Ok(())
2795 }
2796
2797 pub fn state_history(&self) -> Vec<StateTransitionEvent> {
2798 self.state_machine
2799 .as_ref()
2800 .map(|sm| sm.history())
2801 .unwrap_or_default()
2802 }
2803
2804 pub fn session_metadata(&self) -> ai_agents_core::SessionMetadata {
2806 self.session_metadata.read().clone()
2807 }
2808
2809 pub async fn delete_actor_data(&self, actor_id: &str) -> Result<()> {
2812 let allowed = self
2813 .actor_memory_config
2814 .as_ref()
2815 .map(|c| c.privacy.allow_deletion)
2816 .unwrap_or(true);
2817 if !allowed {
2818 return Err(AgentError::Config(
2819 "privacy.allow_deletion is false; actor data deletion is not permitted".into(),
2820 ));
2821 }
2822 let storage = self.storage.read().clone();
2823 if let Some(storage) = storage {
2824 if !storage.supports(StorageCapability::ActorDataDeletion) {
2828 return Err(AgentError::UnsupportedStorageCapability(
2829 StorageCapability::ActorDataDeletion,
2830 ));
2831 }
2832 storage.delete_actor_data(&self.info.id, actor_id).await?;
2833 } else {
2834 let store = { self.fact_store.read().clone() };
2838 if let Some(store) = store {
2839 store.delete_actor_data(actor_id).await?;
2840 }
2841 }
2842 if let Some(manager) = self.relationship_manager.as_ref() {
2843 manager.remove(actor_id);
2844 }
2845 self.actor_facts_cache.write().remove(actor_id);
2846 Ok(())
2847 }
2848
2849 pub fn set_session_metadata(&self, meta: ai_agents_core::SessionMetadata) {
2851 *self.session_metadata.write() = meta;
2852 }
2853
2854 pub async fn cleanup_expired_sessions(&self) -> Result<usize> {
2856 let storage = self.storage.read().clone();
2857 match storage {
2858 Some(s) => {
2859 let count = s.cleanup_expired().await?;
2860 if count > 0 {
2861 self.hooks.on_sessions_expired(count).await;
2862 }
2863 Ok(count)
2864 }
2865 None => Err(AgentError::Config(
2866 "No storage configured. Use with_storage_config() or with_storage() first".into(),
2867 )),
2868 }
2869 }
2870
2871 pub async fn list_sessions_filtered(
2873 &self,
2874 filter: &ai_agents_core::SessionFilter,
2875 ) -> Result<Vec<ai_agents_core::SessionSummary>> {
2876 let storage = self.storage.read().clone();
2877 match storage {
2878 Some(s) => s.list_sessions_filtered(filter).await,
2879 None => Err(AgentError::Config(
2880 "No storage configured. Use with_storage_config() or with_storage() first".into(),
2881 )),
2882 }
2883 }
2884
2885 pub async fn save_state(&self) -> Result<AgentSnapshot> {
2886 let memory_snapshot = self.memory.snapshot().await?;
2887 let state_machine_snapshot = self.state_machine.as_ref().map(|sm| sm.snapshot());
2888 let context_snapshot = self.context_manager.snapshot();
2889
2890 let mut snapshot = AgentSnapshot::new(self.info.id.clone())
2891 .with_memory(memory_snapshot)
2892 .with_context(context_snapshot)
2893 .with_state_machine(
2894 state_machine_snapshot.unwrap_or_else(|| StateMachineSnapshot {
2895 current_state: String::new(),
2896 previous_state: None,
2897 turn_count: 0,
2898 no_transition_count: 0,
2899 history: vec![],
2900 }),
2901 );
2902
2903 if let Some(ref persona) = self.persona_manager {
2904 snapshot.persona = Some(persona.snapshot_as_value()?);
2905 }
2906
2907 if let Some(ref relationships) = self.relationship_manager {
2908 snapshot.relationships = Some(relationships.snapshot_as_value()?);
2909 }
2910
2911 Ok(snapshot)
2912 }
2913
2914 pub async fn save_state_full(&self) -> Result<AgentSnapshot> {
2916 let mut snapshot = self.save_state().await?;
2917 if let Some(ref registry) = self.spawner_registry {
2918 let entries = registry.list_with_specs();
2919 if !entries.is_empty() {
2920 snapshot = snapshot.with_spawned_agents(entries);
2921 }
2922 }
2923 Ok(snapshot)
2924 }
2925
2926 pub async fn restore_state(&self, snapshot: AgentSnapshot) -> Result<()> {
2928 let _admission = self.disambiguation_admission.write().await;
2929 if self.state_transition_reserved.load(Ordering::SeqCst) {
2930 return Err(AgentError::Other(
2931 "Cannot restore state while a state transition is in progress".to_string(),
2932 ));
2933 }
2934 self.invalidate_pending_confirmation("state_restore").await;
2935 *self.pending_skill_id.write() = None;
2936 if let Some(disambiguator) = self.disambiguation_manager.as_ref() {
2937 disambiguator.clear_pending().await;
2938 }
2939 self.memory.restore(snapshot.memory).await?;
2940 self.active_native_exchanges.write().clear();
2941
2942 if let (Some(sm), Some(sm_snapshot)) = (&self.state_machine, snapshot.state_machine)
2943 && !sm_snapshot.current_state.is_empty()
2944 {
2945 sm.restore(sm_snapshot)?;
2946 }
2947
2948 self.context_manager.restore(snapshot.context);
2949
2950 if let (Some(persona_value), Some(persona_manager)) =
2951 (snapshot.persona, &self.persona_manager)
2952 {
2953 persona_manager.restore_from_value(persona_value)?;
2954 }
2955
2956 if let (Some(relationship_value), Some(relationship_manager)) =
2957 (snapshot.relationships, &self.relationship_manager)
2958 {
2959 relationship_manager.restore_from_value(relationship_value)?;
2960 }
2961
2962 info!(agent_id = %snapshot.agent_id, "State restored");
2963 Ok(())
2964 }
2965
2966 pub async fn save_to(&self, storage: &dyn AgentStorage, session_id: &str) -> Result<()> {
2967 let snapshot = self.save_state().await?;
2968 storage.save(session_id, &snapshot).await
2969 }
2970
2971 async fn load_session_restore(
2972 storage: &dyn AgentStorage,
2973 session_id: &str,
2974 ) -> Result<Option<StoredSessionRestore>> {
2975 let Some(snapshot) = storage.load(session_id).await? else {
2976 return Ok(None);
2977 };
2978 let metadata = if storage.supports(StorageCapability::SessionMetadata) {
2982 storage.load_metadata(session_id).await?
2983 } else {
2984 None
2985 };
2986 Ok(Some(StoredSessionRestore { snapshot, metadata }))
2987 }
2988
2989 async fn capture_session_restore_point(&self) -> Result<RuntimeSessionRestorePoint> {
2990 Ok(RuntimeSessionRestorePoint {
2991 snapshot: self.save_state().await?,
2992 metadata: self.session_metadata(),
2993 actor_id: self.actor_id(),
2994 session_id: self.current_session_id.read().clone(),
2995 })
2996 }
2997
2998 async fn apply_session_restore_unchecked(
2999 &self,
3000 session_id: &str,
3001 stored: StoredSessionRestore,
3002 ) -> Result<()> {
3003 self.restore_state(stored.snapshot).await?;
3004 let metadata = stored.metadata.unwrap_or_default();
3005 if let Some(actor_id) = metadata.actor_id.as_deref() {
3006 self.set_actor_id(actor_id)?;
3007 } else {
3008 self.clear_actor_id();
3009 }
3010 self.set_session_metadata(metadata);
3011 *self.current_session_id.write() = Some(session_id.to_string());
3012 Ok(())
3013 }
3014
3015 async fn restore_session_restore_point(
3016 &self,
3017 restore_point: &RuntimeSessionRestorePoint,
3018 ) -> Result<()> {
3019 self.restore_state(restore_point.snapshot.clone()).await?;
3020 if let Some(actor_id) = restore_point.actor_id.as_deref() {
3021 self.set_actor_id(actor_id)?;
3022 } else {
3023 self.clear_actor_id();
3024 }
3025 self.set_session_metadata(restore_point.metadata.clone());
3026 *self.current_session_id.write() = restore_point.session_id.clone();
3027 Ok(())
3028 }
3029
3030 async fn apply_session_restore(
3031 &self,
3032 session_id: &str,
3033 stored: StoredSessionRestore,
3034 ) -> Result<()> {
3035 let before = self.capture_session_restore_point().await?;
3036 if let Err(error) = self
3037 .apply_session_restore_unchecked(session_id, stored)
3038 .await
3039 {
3040 return match self.restore_session_restore_point(&before).await {
3041 Ok(()) => Err(error),
3042 Err(rollback_error) => Err(AgentError::Other(format!(
3043 "Session restore failed: {error}; rollback failed: {rollback_error}"
3044 ))),
3045 };
3046 }
3047 Ok(())
3048 }
3049
3050 async fn rollback_session_restore_set(
3051 parent: Option<(&RuntimeAgent, &RuntimeSessionRestorePoint)>,
3052 children: &[(String, Arc<RuntimeAgent>, RuntimeSessionRestorePoint)],
3053 ) -> Vec<String> {
3054 let mut errors = Vec::new();
3055 if let Some((agent, restore_point)) = parent
3056 && let Err(error) = agent.restore_session_restore_point(restore_point).await
3057 {
3058 errors.push(format!("parent: {error}"));
3059 }
3060 for (id, agent, restore_point) in children {
3061 if let Err(error) = agent.restore_session_restore_point(restore_point).await {
3062 errors.push(format!("child '{id}': {error}"));
3063 }
3064 }
3065 errors
3066 }
3067
3068 fn restore_failure(error: impl std::fmt::Display, rollback_errors: Vec<String>) -> AgentError {
3069 if rollback_errors.is_empty() {
3070 AgentError::Other(format!(
3071 "Session restore failed: {error}; runtime state was rolled back"
3072 ))
3073 } else {
3074 AgentError::Other(format!(
3075 "Session restore failed: {error}; rollback also failed for {}",
3076 rollback_errors.join(", ")
3077 ))
3078 }
3079 }
3080
3081 pub async fn load_from(&self, storage: &dyn AgentStorage, session_id: &str) -> Result<bool> {
3082 let Some(stored) = Self::load_session_restore(storage, session_id).await? else {
3083 return Ok(false);
3084 };
3085 self.apply_session_restore(session_id, stored).await?;
3086 Ok(true)
3087 }
3088
3089 pub async fn save_session(&self, session_id: &str) -> Result<()> {
3090 let storage = self.storage.read().clone();
3091 match storage {
3092 Some(s) => {
3093 let is_new = {
3095 let cur = self.current_session_id.read().clone();
3096 cur.as_deref() != Some(session_id)
3097 };
3098 if is_new {
3099 *self.current_session_id.write() = Some(session_id.to_string());
3100 self.hooks.on_session_created(session_id).await;
3101 }
3102
3103 {
3105 let now = chrono::Utc::now();
3106 let msg_count = self
3107 .memory
3108 .get_messages(None)
3109 .await
3110 .map(|v| v.len())
3111 .unwrap_or(0);
3112 let mut meta = self.session_metadata.write();
3113 meta.last_active = now;
3114 meta.message_count = msg_count;
3115 if meta.actor_id.is_none() {
3116 meta.actor_id = self.actor_id.read().clone();
3117 }
3118 }
3119
3120 let snapshot = self.save_state().await?;
3121 if s.supports(StorageCapability::SessionMetadata) {
3125 let metadata = self.session_metadata.read().clone();
3126 s.save_snapshot_with_metadata(session_id, &snapshot, &metadata)
3127 .await
3128 } else {
3129 s.save(session_id, &snapshot).await
3130 }
3131 }
3132 None => Err(AgentError::Config(
3133 "No storage configured. Use with_storage_config() or with_storage() first".into(),
3134 )),
3135 }
3136 }
3137
3138 pub async fn load_session(&self, session_id: &str) -> Result<bool> {
3139 let storage = self.storage.read().clone();
3140 match storage {
3141 Some(storage) => self.load_from(storage.as_ref(), session_id).await,
3142 None => Err(AgentError::Config(
3143 "No storage configured. Use with_storage_config() or with_storage() first".into(),
3144 )),
3145 }
3146 }
3147
3148 pub async fn restore_session_full(&self, session_id: &str) -> Result<usize> {
3150 self.init_storage().await?;
3151 let storage = self.storage.read().clone().ok_or_else(|| {
3152 AgentError::Config(
3153 "No storage configured. Use with_storage_config() or with_storage() first".into(),
3154 )
3155 })?;
3156 let target_parent = Self::load_session_restore(storage.as_ref(), session_id)
3157 .await?
3158 .ok_or_else(|| AgentError::Persistence(format!("Session not found: {session_id}")))?;
3159 let manifest = target_parent
3160 .snapshot
3161 .spawned_agents
3162 .clone()
3163 .unwrap_or_default();
3164
3165 let registry = self.spawner_registry.as_ref().cloned();
3166 let spawner = if manifest.is_empty() {
3167 self.spawner.as_ref().cloned()
3168 } else {
3169 Some(self.spawner.as_ref().cloned().ok_or_else(|| {
3170 AgentError::Config(
3171 "Saved session contains child agents but this runtime has no spawner".into(),
3172 )
3173 })?)
3174 };
3175 let registry = if manifest.is_empty() {
3176 registry
3177 } else {
3178 Some(registry.ok_or_else(|| {
3179 AgentError::Config(
3180 "Saved session contains child agents but this runtime has no registry".into(),
3181 )
3182 })?)
3183 };
3184
3185 let mut target_ids = HashSet::with_capacity(manifest.len());
3186 let mut prepared = Vec::with_capacity(manifest.len());
3187 for entry in manifest {
3188 if !target_ids.insert(entry.id.clone()) {
3189 return Err(AgentError::InvalidSpec(format!(
3190 "Saved child manifest contains duplicate ID: {}",
3191 entry.id
3192 )));
3193 }
3194 let spec = crate::spec::AgentSpec::from_yaml_strict(&entry.spec_yaml)?;
3195 spawner
3196 .as_ref()
3197 .expect("non-empty manifests require a spawner")
3198 .validate_explicit_child(&entry.id, &spec)?;
3199 prepared.push((entry.id, spec));
3200 }
3201
3202 let current_ids = registry
3203 .as_ref()
3204 .map(|registry| {
3205 registry
3206 .list()
3207 .into_iter()
3208 .map(|info| info.id)
3209 .collect::<HashSet<_>>()
3210 })
3211 .unwrap_or_default();
3212 let removal_count = current_ids.difference(&target_ids).count();
3213 let additions = prepared
3214 .iter()
3215 .filter(|(id, _)| !current_ids.contains(id))
3216 .cloned()
3217 .collect::<Vec<_>>();
3218
3219 let mut existing = Vec::new();
3220 if let Some(registry) = registry.as_ref() {
3221 for (id, _) in prepared.iter().filter(|(id, _)| current_ids.contains(id)) {
3222 let agent = registry.get(id).ok_or_else(|| {
3223 AgentError::Config(format!("Retained child disappeared during restore: {id}"))
3224 })?;
3225 let child_storage = agent.storage().ok_or_else(|| {
3226 AgentError::Config(format!("Child '{id}' has no storage for session restore"))
3227 })?;
3228 let stored = Self::load_session_restore(child_storage.as_ref(), session_id)
3229 .await?
3230 .ok_or_else(|| {
3231 AgentError::Persistence(format!(
3232 "Child '{id}' has no saved session '{session_id}'"
3233 ))
3234 })?;
3235 existing.push((id.clone(), agent, stored));
3236 }
3237 }
3238
3239 let mut staged = Vec::with_capacity(additions.len());
3240 if !additions.is_empty() {
3241 let spawner = spawner
3242 .as_ref()
3243 .expect("restored additions require a spawner");
3244 let reservations = spawner.reserve_restore_capacity(additions.len(), removal_count)?;
3245 for ((id, spec), reservation) in additions.into_iter().zip(reservations) {
3246 let spawned = spawner
3247 .spawn_with_reserved_capacity(id.clone(), spec, reservation)
3248 .await?;
3249 let child_storage = spawned.agent.storage().ok_or_else(|| {
3250 AgentError::Config(format!("Child '{id}' has no storage for session restore"))
3251 })?;
3252 let stored = Self::load_session_restore(child_storage.as_ref(), session_id)
3253 .await?
3254 .ok_or_else(|| {
3255 AgentError::Persistence(format!(
3256 "Child '{id}' has no saved session '{session_id}'"
3257 ))
3258 })?;
3259 staged.push((spawned, stored));
3260 }
3261 } else if let Some(spawner) = spawner.as_ref() {
3262 spawner.reserve_restore_capacity(0, removal_count)?;
3263 }
3264
3265 let parent_before = self.capture_session_restore_point().await?;
3266 let mut existing_before = Vec::with_capacity(existing.len());
3267 for (id, agent, _) in &existing {
3268 existing_before.push((
3269 id.clone(),
3270 Arc::clone(agent),
3271 agent.capture_session_restore_point().await?,
3272 ));
3273 }
3274
3275 for (_, agent, stored) in &existing {
3279 if let Err(error) = agent
3280 .apply_session_restore_unchecked(session_id, stored.clone())
3281 .await
3282 {
3283 drop(staged);
3284 let rollback_errors =
3285 Self::rollback_session_restore_set(None, &existing_before).await;
3286 return Err(Self::restore_failure(error, rollback_errors));
3287 }
3288 }
3289 for (spawned, stored) in &staged {
3290 if let Err(error) = spawned
3291 .agent
3292 .apply_session_restore_unchecked(session_id, stored.clone())
3293 .await
3294 {
3295 drop(staged);
3296 let rollback_errors =
3297 Self::rollback_session_restore_set(None, &existing_before).await;
3298 return Err(Self::restore_failure(error, rollback_errors));
3299 }
3300 }
3301 if let Err(error) = self
3302 .apply_session_restore_unchecked(session_id, target_parent)
3303 .await
3304 {
3305 drop(staged);
3306 let rollback_errors =
3307 Self::rollback_session_restore_set(Some((self, &parent_before)), &existing_before)
3308 .await;
3309 return Err(Self::restore_failure(error, rollback_errors));
3310 }
3311
3312 if let Some(registry) = registry.as_ref()
3313 && let Err(error) = registry
3314 .reconcile(
3315 &target_ids,
3316 staged.into_iter().map(|(spawned, _)| spawned).collect(),
3317 )
3318 .await
3319 {
3320 let rollback_errors =
3321 Self::rollback_session_restore_set(Some((self, &parent_before)), &existing_before)
3322 .await;
3323 return Err(Self::restore_failure(error, rollback_errors));
3324 }
3325
3326 Ok(target_ids.len())
3327 }
3328
3329 pub async fn delete_session(&self, session_id: &str) -> Result<()> {
3330 let storage = self.storage.read().clone();
3331 match storage {
3332 Some(s) => s.delete(session_id).await,
3333 None => Err(AgentError::Config(
3334 "No storage configured. Use with_storage_config() or with_storage() first".into(),
3335 )),
3336 }
3337 }
3338
3339 pub async fn list_sessions(&self) -> Result<Vec<String>> {
3340 let storage = self.storage.read().clone();
3341 match storage {
3342 Some(s) => s.list_sessions().await,
3343 None => Err(AgentError::Config(
3344 "No storage configured. Use with_storage_config() or with_storage() first".into(),
3345 )),
3346 }
3347 }
3348
3349 fn estimate_tokens(&self, text: &str) -> u32 {
3350 (text.len() as f32 / 4.0).ceil() as u32
3351 }
3352
3353 fn estimate_total_tokens(&self, messages: &[ChatMessage]) -> u32 {
3354 messages
3355 .iter()
3356 .map(|m| self.estimate_tokens(&m.content))
3357 .sum()
3358 }
3359
3360 fn native_safe_prefix_at_least(messages: &[ChatMessage], required: usize) -> Result<usize> {
3362 let inspection =
3363 inspect_native_history(messages).map_err(|error| AgentError::LLM(error.to_string()))?;
3364 let has_signed_history = !inspection.exchanges().is_empty();
3365 for count in required.min(messages.len())..=messages.len() {
3366 if has_signed_history
3367 && count < messages.len()
3368 && messages[count].role != ai_agents_core::Role::User
3369 {
3370 continue;
3371 }
3372 if inspection.is_safe_prefix_len(count)
3373 && inspect_native_history(&messages[count..]).is_ok()
3374 {
3375 return Ok(count);
3376 }
3377 }
3378 Err(AgentError::LLM(
3379 "Context limits cannot remove a complete native history prefix".to_string(),
3380 ))
3381 }
3382
3383 fn truncate_context(&self, messages: &mut Vec<ChatMessage>, keep_recent: usize) -> Result<()> {
3385 if messages.len() <= keep_recent + 1 {
3386 return Ok(());
3387 }
3388 let system_msg = messages.remove(0);
3389 let required = messages.len().saturating_sub(keep_recent);
3390 let to_remove = Self::native_safe_prefix_at_least(messages, required)?;
3391 messages.drain(..to_remove);
3392 messages.insert(0, system_msg);
3393 Ok(())
3394 }
3395
3396 fn get_filter(&self, config: &FilterConfig) -> Arc<dyn MessageFilter> {
3397 match config {
3398 FilterConfig::KeepRecent(n) => Arc::new(KeepRecentFilter::new(*n)),
3399 FilterConfig::ByRole { keep_roles } => Arc::new(ByRoleFilter::new(keep_roles.clone())),
3400 FilterConfig::SkipPattern { skip_if_contains } => {
3401 Arc::new(SkipPatternFilter::new(skip_if_contains.clone()))
3402 }
3403 FilterConfig::Custom { name } => {
3404 let filters = self.message_filters.read();
3405 filters
3406 .get(name)
3407 .cloned()
3408 .unwrap_or_else(|| Arc::new(KeepRecentFilter::new(10)))
3409 }
3410 }
3411 }
3412
3413 async fn summarize_context(
3414 &self,
3415 messages: &mut Vec<ChatMessage>,
3416 summarizer_llm: Option<&str>,
3417 max_summary_tokens: u32,
3418 custom_prompt: Option<&str>,
3419 keep_recent: usize,
3420 filter: Option<&FilterConfig>,
3421 ) -> Result<()> {
3422 let system_msg = messages.remove(0);
3423
3424 let required = messages.len().saturating_sub(keep_recent);
3425 if required == 0 {
3426 messages.insert(0, system_msg);
3427 return Ok(());
3428 }
3429 let to_summarize_count = Self::native_safe_prefix_at_least(messages, required)?;
3430
3431 let recent_msgs: Vec<ChatMessage> = messages.drain(to_summarize_count..).collect();
3432 let mut to_summarize = std::mem::take(messages);
3433
3434 if let Some(filter_config) = filter {
3435 let filter = self.get_filter(filter_config);
3436 to_summarize = filter.filter(to_summarize);
3437 }
3438
3439 if to_summarize.is_empty() {
3440 *messages = recent_msgs;
3441 messages.insert(0, system_msg);
3442 return Ok(());
3443 }
3444
3445 let to_summarize = Self::readable_native_messages(to_summarize)?;
3446 let conversation_text = to_summarize
3447 .iter()
3448 .map(|m| format!("{:?}: {}", m.role, m.content))
3449 .collect::<Vec<_>>()
3450 .join("\n");
3451
3452 let default_prompt = format!(
3453 "Summarize the following conversation in under {} tokens, preserving key information:\n\n{}",
3454 max_summary_tokens, conversation_text
3455 );
3456
3457 let summary_prompt = custom_prompt
3458 .map(|p| format!("{}\n\n{}", p, conversation_text))
3459 .unwrap_or(default_prompt);
3460
3461 let summarizer = if let Some(alias) = summarizer_llm {
3462 self.llm_registry
3463 .get(alias)
3464 .map_err(|e| AgentError::Config(e.to_string()))?
3465 } else {
3466 self.llm_registry
3467 .router()
3468 .or_else(|_| self.llm_registry.default())
3469 .map_err(|e| AgentError::Config(e.to_string()))?
3470 };
3471
3472 let summary_msgs = vec![ChatMessage::user(&summary_prompt)];
3473 let response = self
3474 .observe_purpose(
3475 ObservationPurpose::Summarization,
3476 summarizer.complete(&summary_msgs, None),
3477 )
3478 .await?;
3479
3480 let summary_message = ChatMessage::system(format!(
3481 "[Previous conversation summary]\n{}",
3482 response.content
3483 ));
3484
3485 *messages = vec![system_msg, summary_message];
3486 messages.extend(recent_msgs);
3487
3488 debug!(
3489 summarized_count = to_summarize_count,
3490 kept_recent = keep_recent,
3491 "Context summarized"
3492 );
3493
3494 Ok(())
3495 }
3496
3497 fn render_system_prompt(&self) -> Result<String> {
3498 let mut context = self.build_context_with_overlays();
3499
3500 let facts_text = self.format_actor_facts_for_context();
3502 if !facts_text.is_empty() {
3503 context.insert(
3504 "actor_facts".to_string(),
3505 serde_json::Value::String(facts_text),
3506 );
3507 }
3508
3509 if let Some((key, text)) = self.format_relationship_for_context() {
3510 context.insert(key, serde_json::Value::String(text));
3511 }
3512
3513 self.template_renderer
3514 .render(&self.base_system_prompt, &context)
3515 }
3516
3517 fn canonical_unique_tool_ids(&self, ids: &[String]) -> Vec<String> {
3519 let mut seen = HashSet::new();
3520 ids.iter()
3521 .filter_map(|id| self.tools.canonical_id(id))
3522 .filter(|canonical_id| seen.insert(canonical_id.clone()))
3523 .collect()
3524 }
3525
3526 fn get_top_level_tool_ids_for_scope(&self, scope_override: Option<&[String]>) -> Vec<String> {
3528 let Some(declared) = self.declared_tool_ids.as_deref() else {
3529 return Vec::new();
3530 };
3531 let mut effective = self.canonical_unique_tool_ids(declared);
3532 if let Some(scope) = scope_override {
3533 let scope: HashSet<String> =
3534 self.canonical_unique_tool_ids(scope).into_iter().collect();
3535 effective.retain(|canonical_id| scope.contains(canonical_id));
3536 }
3537 effective
3538 }
3539
3540 async fn get_available_tool_ids(&self) -> Result<Vec<String>> {
3542 Ok(self.get_available_tool_ids_snapshot().await?.tool_ids)
3543 }
3544
3545 async fn get_available_tool_ids_snapshot(&self) -> Result<AvailableToolIdsSnapshot> {
3547 let scope_override = self.runtime_control.tool_scope_override.read().clone();
3548 self.get_available_tool_ids_snapshot_for_scope(scope_override.as_deref())
3549 .await
3550 }
3551
3552 async fn get_available_tool_ids_snapshot_for_scope(
3554 &self,
3555 scope_override: Option<&[String]>,
3556 ) -> Result<AvailableToolIdsSnapshot> {
3557 let mut available = self.get_top_level_tool_ids_for_scope(scope_override);
3558 let (state_generation, state_scopes) = self
3559 .state_machine
3560 .as_ref()
3561 .map(|state_machine| {
3562 let (generation, scopes) = state_machine.current_tool_scope_snapshot();
3563 (Some(generation), scopes)
3564 })
3565 .unwrap_or((None, Vec::new()));
3566
3567 if available.is_empty() || state_scopes.is_empty() {
3568 return Ok(AvailableToolIdsSnapshot {
3569 tool_ids: available,
3570 state_generation,
3571 });
3572 }
3573
3574 let eval_ctx = self.build_evaluation_context().await?;
3575 let llm_getter = RegistryLLMGetter {
3576 registry: self.llm_registry.clone(),
3577 };
3578 let evaluator = ConditionEvaluator::new(llm_getter);
3579
3580 for state_scope in state_scopes {
3581 if state_scope.is_empty() {
3582 available.clear();
3583 break;
3584 }
3585
3586 let mut allowed = HashSet::new();
3587 for tool_ref in &state_scope {
3588 let tool_id = tool_ref.id();
3589 let Some(canonical_id) = self.tools.canonical_id(tool_id) else {
3590 continue;
3591 };
3592 let condition_matches = if let Some(condition) = tool_ref.condition() {
3593 match evaluator.evaluate(condition, &eval_ctx).await {
3594 Ok(matches) => matches,
3595 Err(error) => {
3596 warn!(tool = tool_id, error = %error, "Error evaluating tool condition");
3597 false
3598 }
3599 }
3600 } else {
3601 true
3602 };
3603 if condition_matches {
3604 allowed.insert(canonical_id);
3605 } else {
3606 debug!(tool = tool_id, "Tool condition not met, skipping");
3607 }
3608 }
3609 available.retain(|canonical_id| allowed.contains(canonical_id));
3610 if available.is_empty() {
3611 break;
3612 }
3613 }
3614
3615 Ok(AvailableToolIdsSnapshot {
3616 tool_ids: available,
3617 state_generation,
3618 })
3619 }
3620
3621 async fn build_evaluation_context(&self) -> Result<EvaluationContext> {
3622 let context = self.build_context_with_overlays();
3623 let messages = Self::readable_native_messages(self.memory.get_messages(Some(10)).await?)?;
3624 let tool_history = self.tool_call_history.read().clone();
3625
3626 let (state_name, turn_count, previous_state) = if let Some(ref sm) = self.state_machine {
3627 (Some(sm.current()), sm.turn_count(), sm.previous())
3628 } else {
3629 (None, 0, None)
3630 };
3631
3632 Ok(EvaluationContext::default()
3633 .with_context(context)
3634 .with_state(state_name, turn_count, previous_state)
3635 .with_called_tools(tool_history)
3636 .with_messages(messages))
3637 }
3638
3639 fn record_tool_call(&self, tool_id: &str, result: Value) {
3640 self.tool_call_history.write().push(ToolCallRecord {
3641 tool_id: tool_id.to_string(),
3642 result,
3643 timestamp: chrono::Utc::now(),
3644 });
3645 }
3646
3647 async fn get_effective_system_prompt_with_persona_hooks(
3648 &self,
3649 fire_persona_hooks: bool,
3650 include_tool_prompt: bool,
3651 ) -> Result<String> {
3652 let rendered_base = self.render_system_prompt()?;
3653
3654 let persona_prefix = if let Some(ref persona) = self.persona_manager {
3655 let context = self.build_context_with_overlays();
3656 if fire_persona_hooks {
3657 let render_result = persona.render_prompt(&context)?;
3658 for content in &render_result.newly_revealed {
3659 self.hooks.on_secret_revealed(content).await;
3660 }
3661 render_result.prompt
3662 } else {
3663 persona.render_prompt_preview(&context)?
3664 }
3665 } else {
3666 String::new()
3667 };
3668
3669 if let Some(ref sm) = self.state_machine
3670 && let Some(state_def) = sm.current_definition()
3671 {
3672 let state_prompt = if let Some(ref prompt) = state_def.prompt {
3673 let context = self.build_context_with_overlays();
3674 self.template_renderer.render_with_state(
3675 prompt,
3676 &context,
3677 &sm.current(),
3678 sm.previous().as_deref(),
3679 sm.turn_count(),
3680 state_def.max_turns,
3681 )?
3682 } else {
3683 String::new()
3684 };
3685
3686 let combined = match state_def.prompt_mode {
3687 PromptMode::Append => {
3688 if state_prompt.is_empty() {
3689 rendered_base
3690 } else {
3691 format!(
3692 "{}\n\n[Current State: {}]\n{}",
3693 rendered_base,
3694 sm.current(),
3695 state_prompt
3696 )
3697 }
3698 }
3699 PromptMode::Replace => {
3700 if state_prompt.is_empty() {
3701 rendered_base
3702 } else {
3703 state_prompt
3704 }
3705 }
3706 PromptMode::Prepend => {
3707 if state_prompt.is_empty() {
3708 rendered_base
3709 } else {
3710 format!("{}\n\n{}", state_prompt, rendered_base)
3711 }
3712 }
3713 };
3714
3715 let with_persona = if persona_prefix.is_empty() {
3717 combined
3718 } else {
3719 format!("{}\n\n{}", persona_prefix, combined)
3720 };
3721
3722 if include_tool_prompt {
3723 let available_tool_ids = self.get_available_tool_ids().await?;
3724 if !available_tool_ids.is_empty() {
3725 let tools_prompt = self.tools.generate_scoped_prompt_with_mode(
3726 &available_tool_ids,
3727 None,
3728 self.parallel_tools.enabled,
3729 self.runtime_config.tool_schema_prompt_mode,
3730 );
3731 if !tools_prompt.is_empty() {
3732 return Ok(format!("{}\n\n{}", with_persona, tools_prompt));
3733 }
3734 }
3735 }
3736 return Ok(with_persona);
3737 }
3738
3739 let with_persona = if persona_prefix.is_empty() {
3741 rendered_base
3742 } else {
3743 format!("{}\n\n{}", persona_prefix, rendered_base)
3744 };
3745
3746 if include_tool_prompt {
3747 let available_tool_ids = self.get_available_tool_ids().await?;
3748 let tools_prompt = self.tools.generate_scoped_prompt_with_mode(
3749 &available_tool_ids,
3750 None,
3751 self.parallel_tools.enabled,
3752 self.runtime_config.tool_schema_prompt_mode,
3753 );
3754 if !tools_prompt.is_empty() {
3755 return Ok(format!("{}\n\n{}", with_persona, tools_prompt));
3756 }
3757 }
3758 Ok(with_persona)
3759 }
3760
3761 fn get_state_llm(&self) -> Result<Arc<dyn LLMProvider>> {
3762 if let Some(ref sm) = self.state_machine
3763 && let Some(state_def) = sm.current_definition()
3764 && let Some(ref llm_alias) = state_def.llm
3765 {
3766 return self
3767 .llm_registry
3768 .get(llm_alias)
3769 .map_err(|e| AgentError::Config(e.to_string()));
3770 }
3771 self.llm_registry
3772 .default()
3773 .map_err(|e| AgentError::Config(e.to_string()))
3774 }
3775
3776 fn get_effective_reasoning_config(&self) -> ReasoningConfig {
3777 if let Some(ref sm) = self.state_machine
3778 && let Some(state_def) = sm.current_definition()
3779 && let Some(ref state_reasoning) = state_def.reasoning
3780 {
3781 return state_reasoning.clone();
3782 }
3783 self.reasoning_config.clone()
3784 }
3785
3786 fn get_effective_reflection_config(&self) -> ReflectionConfig {
3787 if let Some(ref sm) = self.state_machine
3788 && let Some(state_def) = sm.current_definition()
3789 && let Some(ref state_reflection) = state_def.reflection
3790 {
3791 return state_reflection.clone();
3792 }
3793 self.reflection_config.clone()
3794 }
3795
3796 fn get_skill_reasoning_config(&self, skill: &SkillDefinition) -> ReasoningConfig {
3797 skill
3798 .reasoning
3799 .clone()
3800 .unwrap_or_else(|| self.get_effective_reasoning_config())
3801 }
3802
3803 fn get_skill_reflection_config(&self, skill: &SkillDefinition) -> ReflectionConfig {
3804 skill
3805 .reflection
3806 .clone()
3807 .unwrap_or_else(|| self.get_effective_reflection_config())
3808 }
3809
3810 async fn build_disambiguation_context(&self) -> Result<DisambiguationContext> {
3811 let recent_messages: Vec<String> =
3812 Self::readable_native_messages(self.memory.get_messages(Some(5)).await?)?
3813 .iter()
3814 .rev()
3815 .map(|m| format!("{:?}: {}", m.role, m.content))
3816 .collect();
3817
3818 let current_state = self.current_state().map(|s| s.to_string());
3819
3820 let state_prompt: Option<String> = self
3823 .state_machine
3824 .as_ref()
3825 .and_then(|sm| sm.current_definition())
3826 .and_then(|def| def.prompt.clone());
3827
3828 let available_tools: Vec<String> = self
3829 .get_available_tool_ids()
3830 .await
3831 .unwrap_or_else(|_| self.tools.list_ids());
3832
3833 let available_skills: Vec<String> = self.skills.iter().map(|s| s.id.clone()).collect();
3834
3835 let mut user_context = self.build_context_with_overlays();
3836 user_context.remove(DISAMBIGUATION_STATE_GENERATION_KEY);
3837 if let Some(state_generation) = self
3838 .state_machine
3839 .as_ref()
3840 .map(|state_machine| state_machine.generation())
3841 {
3842 user_context.insert(
3843 DISAMBIGUATION_STATE_GENERATION_KEY.to_string(),
3844 serde_json::json!(state_generation),
3845 );
3846 }
3847
3848 let available_intents: Vec<String> = if let Some(ref sm) = self.state_machine {
3850 sm.current_definition()
3851 .map(|def| {
3852 def.transitions
3853 .iter()
3854 .filter_map(|t| t.intent.clone())
3855 .collect()
3856 })
3857 .unwrap_or_default()
3858 } else {
3859 Vec::new()
3860 };
3861
3862 Ok(DisambiguationContext::from_agent_state(
3863 recent_messages,
3864 current_state,
3865 state_prompt,
3866 available_tools,
3867 available_skills,
3868 available_intents,
3869 user_context,
3870 ))
3871 }
3872
3873 fn get_available_skills(&self) -> Vec<&SkillDefinition> {
3874 if let Some(ref sm) = self.state_machine
3875 && let Some(state_def) = sm.current_definition()
3876 {
3877 let parent_def = sm.get_parent_definition();
3878 let effective_skills = state_def.get_effective_skills(parent_def.as_ref());
3879 if !effective_skills.is_empty() {
3880 return self
3881 .skills
3882 .iter()
3883 .filter(|s| effective_skills.contains(&&s.id))
3884 .collect();
3885 }
3886 }
3887 self.skills.iter().collect()
3888 }
3889
3890 async fn build_messages(&self) -> Result<Vec<ChatMessage>> {
3891 self.build_messages_internal(true, None, true).await
3892 }
3893
3894 async fn build_messages_for_draft(&self, user_message: &str) -> Result<Vec<ChatMessage>> {
3895 self.build_messages_internal(false, Some(user_message), true)
3896 .await
3897 }
3898
3899 async fn build_messages_internal(
3900 &self,
3901 fire_persona_hooks: bool,
3902 ephemeral_user_message: Option<&str>,
3903 include_tool_prompt: bool,
3904 ) -> Result<Vec<ChatMessage>> {
3905 let system_prompt = self
3906 .get_effective_system_prompt_with_persona_hooks(fire_persona_hooks, include_tool_prompt)
3907 .await?;
3908 let mut messages = vec![ChatMessage::system(&system_prompt)];
3909
3910 let context = self.memory.get_context().await?;
3911 let history = if let Some(ref budget) = self.memory_token_budget {
3912 context.to_llm_messages_with_allocation(&budget.allocation)
3913 } else {
3914 context.to_llm_messages()
3915 };
3916 messages.extend(history);
3917 if let Some(user_message) = ephemeral_user_message {
3918 messages.push(ChatMessage::user(user_message));
3919 }
3920
3921 let total_tokens = self.estimate_total_tokens(&messages);
3922
3923 if total_tokens > self.max_context_tokens {
3924 debug!(
3925 total = total_tokens,
3926 limit = self.max_context_tokens,
3927 "Context overflow"
3928 );
3929
3930 match &self.recovery_manager.config().llm.on_context_overflow {
3931 ContextOverflowAction::Error => {
3932 return Err(AgentError::LLM(format!(
3933 "Context overflow: {} tokens > {} limit",
3934 total_tokens, self.max_context_tokens
3935 )));
3936 }
3937 ContextOverflowAction::Truncate { keep_recent } => {
3938 self.truncate_context(&mut messages, *keep_recent)?;
3939 }
3940 ContextOverflowAction::Summarize {
3941 summarizer_llm,
3942 max_summary_tokens,
3943 custom_prompt,
3944 keep_recent,
3945 filter,
3946 } => {
3947 self.summarize_context(
3948 &mut messages,
3949 summarizer_llm.as_deref(),
3950 *max_summary_tokens,
3951 custom_prompt.as_deref(),
3952 *keep_recent,
3953 filter.as_ref(),
3954 )
3955 .await?;
3956 }
3957 }
3958 }
3959
3960 self.validate_active_native_history(&messages, true)?;
3961 Ok(messages)
3962 }
3963
3964 async fn main_tool_protocol(
3965 &self,
3966 llm: &dyn LLMProvider,
3967 ephemeral_new_turn: bool,
3968 ) -> Result<MainToolProtocol> {
3969 let mut choice = llm.configured_tool_choice();
3970 if matches!(choice.as_ref(), Some(ToolChoice::None)) {
3971 return Ok(MainToolProtocol {
3972 choice,
3973 tool_ids: Vec::new(),
3974 definitions: Vec::new(),
3975 });
3976 }
3977
3978 let mut tool_ids = self.get_available_tool_ids().await?;
3979 tool_ids.sort();
3980 tool_ids.dedup();
3981 if let Some(ToolChoice::Specific(expected)) = choice.as_ref() {
3982 let canonical = self.tools.canonical_id(expected).ok_or_else(|| {
3983 AgentError::Config(format!(
3984 "specific tool choice '{expected}' is not registered"
3985 ))
3986 })?;
3987 if canonical != *expected {
3988 return Err(AgentError::Config(format!(
3989 "specific tool choice must use canonical ID '{canonical}', not '{expected}'"
3990 )));
3991 }
3992 if !tool_ids.iter().any(|tool_id| tool_id == expected) {
3993 return Err(AgentError::Config(format!(
3994 "specific tool choice '{expected}' is outside the effective tool grant"
3995 )));
3996 }
3997 }
3998 if matches!(
3999 choice.as_ref(),
4000 Some(ToolChoice::Required | ToolChoice::Specific(_))
4001 ) && tool_ids.is_empty()
4002 {
4003 return Err(AgentError::Config(
4004 "required tool choice has no tool inside the effective grant".to_string(),
4005 ));
4006 }
4007 if !ephemeral_new_turn
4008 && let Some(configured_choice) = choice.as_ref()
4009 && matches!(
4010 configured_choice,
4011 ToolChoice::Required | ToolChoice::Specific(_)
4012 )
4013 && self
4014 .tool_choice_satisfied_in_current_turn(configured_choice, &tool_ids)
4015 .await?
4016 {
4017 choice = Some(ToolChoice::Auto);
4018 }
4019 if let Some(ToolChoice::Specific(expected)) = choice.as_ref() {
4020 tool_ids.retain(|tool_id| tool_id == expected);
4021 }
4022
4023 let definitions = tool_ids
4024 .iter()
4025 .map(|tool_id| {
4026 let tool = self.tools.get(tool_id).ok_or_else(|| {
4027 AgentError::Config(format!(
4028 "effective tool '{tool_id}' disappeared before provider exposure"
4029 ))
4030 })?;
4031 Ok(LLMToolDefinition {
4032 name: tool_id.clone(),
4033 description: tool.description().to_string(),
4034 input_schema: tool.input_schema(),
4035 })
4036 })
4037 .collect::<Result<Vec<_>>>()?;
4038
4039 Ok(MainToolProtocol {
4043 choice,
4044 tool_ids,
4045 definitions,
4046 })
4047 }
4048
4049 async fn tool_choice_satisfied_in_current_turn(
4050 &self,
4051 choice: &ToolChoice,
4052 effective_tool_ids: &[String],
4053 ) -> Result<bool> {
4054 let messages = self.memory.get_messages(None).await?;
4055 let mut saw_tool_result = false;
4056 for message in messages.iter().rev() {
4057 match message.role {
4058 ai_agents_core::Role::Tool | ai_agents_core::Role::Function => {
4059 saw_tool_result = true;
4060 }
4061 ai_agents_core::Role::Assistant if saw_tool_result => {
4062 let Some(calls) = self.parse_tool_calls(&message.content)? else {
4063 continue;
4064 };
4065 let calls_are_effective = !calls.is_empty()
4066 && calls.iter().all(|call| {
4067 self.tools
4068 .canonical_id(&call.name)
4069 .is_some_and(|canonical| effective_tool_ids.contains(&canonical))
4070 });
4071 return Ok(calls_are_effective
4072 && match choice {
4073 ToolChoice::Required => true,
4074 ToolChoice::Specific(expected) => calls.iter().all(|call| {
4075 self.tools.canonical_id(&call.name).as_deref()
4076 == Some(expected.as_str())
4077 }),
4078 _ => false,
4079 });
4080 }
4081 ai_agents_core::Role::User => return Ok(false),
4082 _ => {}
4083 }
4084 }
4085 Ok(false)
4086 }
4087
4088 fn provider_can_use_native_tools(
4089 &self,
4090 llm: &dyn LLMProvider,
4091 protocol: &MainToolProtocol,
4092 ) -> bool {
4093 let Some(choice) = protocol.choice.as_ref() else {
4094 return false;
4095 };
4096 if matches!(choice, ToolChoice::None) || protocol.definitions.is_empty() {
4097 return false;
4098 }
4099 llm.supports_tool_choice(choice)
4100 && protocol.definitions.iter().all(|definition| {
4101 !definition.name.is_empty()
4102 && definition.name.len() <= 64
4103 && definition
4104 .name
4105 .bytes()
4106 .all(|byte| byte.is_ascii_alphanumeric() || matches!(byte, b'_' | b'-'))
4107 })
4108 }
4109
4110 fn prompt_messages_for_tool_protocol(
4111 &self,
4112 messages: &[ChatMessage],
4113 protocol: &MainToolProtocol,
4114 corrective: bool,
4115 ) -> Vec<ChatMessage> {
4116 let mut messages = messages.to_vec();
4117 let Some(choice) = protocol.choice.as_ref() else {
4118 return messages;
4119 };
4120 if matches!(choice, ToolChoice::None) || protocol.tool_ids.is_empty() {
4121 return messages;
4122 }
4123
4124 let mut tool_prompt = self.tools.generate_scoped_prompt_with_mode(
4125 &protocol.tool_ids,
4126 None,
4127 self.parallel_tools.enabled,
4128 self.runtime_config.tool_schema_prompt_mode,
4129 );
4130 match choice {
4131 ToolChoice::Required => tool_prompt.push_str(
4132 "\n\nYou must call at least one listed tool before giving a final answer.",
4133 ),
4134 ToolChoice::Specific(tool_id) => tool_prompt.push_str(&format!(
4135 "\n\nYou must call the '{tool_id}' tool before giving a final answer."
4136 )),
4137 ToolChoice::Auto => {}
4138 ToolChoice::None => return messages,
4139 _ => return messages,
4140 }
4141 if let Some(system) = messages
4142 .iter_mut()
4143 .find(|message| message.role == ai_agents_core::Role::System)
4144 {
4145 system.content.push_str("\n\n");
4146 system.content.push_str(&tool_prompt);
4147 } else {
4148 messages.insert(0, ChatMessage::system(tool_prompt));
4149 }
4150 if corrective {
4151 let instruction = match choice {
4152 ToolChoice::Required => {
4153 "Your previous response did not call a required tool. Call at least one listed tool now and return only the JSON tool call."
4154 }
4155 ToolChoice::Specific(tool_id) => {
4156 messages.push(ChatMessage::user(format!(
4157 "Your previous response did not call the required '{tool_id}' tool. Call it now and return only the JSON tool call."
4158 )));
4159 return messages;
4160 }
4161 _ => return messages,
4162 };
4163 messages.push(ChatMessage::user(instruction));
4164 }
4165 messages
4166 }
4167
4168 async fn invoke_main_provider(
4169 &self,
4170 llm: Arc<dyn LLMProvider>,
4171 messages: &[ChatMessage],
4172 protocol: &MainToolProtocol,
4173 corrective: bool,
4174 ) -> std::result::Result<MainProviderResponse, LLMError> {
4175 let use_native = self.provider_can_use_native_tools(llm.as_ref(), protocol);
4176 let response = if use_native {
4177 let request = LLMToolRequest {
4178 tools: protocol.definitions.clone(),
4179 choice: protocol
4180 .choice
4181 .clone()
4182 .expect("native tool requests require an explicit choice"),
4183 };
4184 self.observe_purpose(
4185 ObservationPurpose::MainResponse,
4186 llm.complete_with_tools(messages, None, &request),
4187 )
4188 .await?
4189 } else {
4190 let prompt_messages =
4191 self.prompt_messages_for_tool_protocol(messages, protocol, corrective);
4192 self.observe_purpose(
4193 ObservationPurpose::MainResponse,
4194 llm.complete(&prompt_messages, None),
4195 )
4196 .await?
4197 };
4198 Ok(MainProviderResponse {
4199 response,
4200 used_native_tools: use_native,
4201 })
4202 }
4203
4204 async fn complete_main_attempt_with_recovery(
4205 &self,
4206 llm: Arc<dyn LLMProvider>,
4207 messages: &[ChatMessage],
4208 protocol: &MainToolProtocol,
4209 corrective: bool,
4210 ) -> Result<MainProviderResponse> {
4211 let primary_result = self
4213 .recovery_manager
4214 .with_llm_retry(
4215 "llm_call",
4216 None,
4217 || {
4218 let llm = Arc::clone(&llm);
4219 async move {
4220 self.invoke_main_provider(llm, messages, protocol, corrective)
4221 .await
4222 }
4223 },
4224 |error| llm.is_terminal_error(error),
4225 )
4226 .await;
4227
4228 match primary_result {
4229 Ok(response) => Ok(response),
4230 Err(ai_agents_recovery::RetryFailure::Terminal { error, .. }) => {
4231 Err(AgentError::LLM(error.to_string()))
4232 }
4233 Err(failure) => {
4234 let primary_error = AgentError::LLM(failure.into_error().to_string());
4235 match &self.recovery_manager.config().llm.on_failure {
4236 LLMFailureAction::FallbackLlm { fallback_llm } => {
4237 let fallback = self.llm_registry.get(fallback_llm).map_err(|error| {
4238 AgentError::Config(format!(
4239 "Fallback LLM '{fallback_llm}' not found: {error}"
4240 ))
4241 })?;
4242 self.invoke_main_provider(fallback, messages, protocol, corrective)
4243 .await
4244 .map_err(|error| AgentError::LLM(error.to_string()))
4245 }
4246 LLMFailureAction::FallbackResponse { message } => {
4247 if matches!(
4248 protocol.choice.as_ref(),
4249 Some(ToolChoice::Required | ToolChoice::Specific(_))
4250 ) {
4251 Err(AgentError::LLM(format!(
4252 "Required tool selection failed and cannot be satisfied by a static fallback response: {primary_error}"
4253 )))
4254 } else {
4255 Ok(MainProviderResponse {
4256 response: LLMResponse::new(message.clone(), FinishReason::Stop),
4257 used_native_tools: false,
4258 })
4259 }
4260 }
4261 LLMFailureAction::Error => Err(primary_error),
4262 }
4263 }
4264 }
4265 }
4266
4267 fn normalize_main_provider_response(
4268 &self,
4269 mut response: LLMResponse,
4270 protocol: &MainToolProtocol,
4271 ) -> Result<(LLMResponse, bool)> {
4272 let provider_state = response
4273 .take_provider_state()
4274 .map_err(|error| AgentError::LLM(error.to_string()))?;
4275 let native_calls = response
4276 .tool_calls()
4277 .map_err(|error| AgentError::LLM(error.to_string()))?;
4278 let calls = match native_calls {
4279 Some(calls) => {
4280 response.content = encode_native_tool_call_markers(&calls, provider_state.as_ref())
4281 .map_err(|error| AgentError::LLM(error.to_string()))?;
4282 Some(calls)
4283 }
4284 None if provider_state.is_some() => {
4285 return Err(AgentError::LLM(
4286 "Provider returned replay state without native tool calls".to_string(),
4287 ));
4288 }
4289 None if !matches!(protocol.choice.as_ref(), Some(ToolChoice::None)) => {
4290 self.parse_tool_calls(response.content.trim())?
4291 }
4292 None => None,
4293 };
4294
4295 if protocol.choice.is_some()
4296 && let Some(calls) = calls.as_ref()
4297 && calls.iter().any(|call| {
4298 self.tools
4299 .canonical_id(&call.name)
4300 .is_none_or(|canonical| !protocol.tool_ids.contains(&canonical))
4301 })
4302 {
4303 return Err(AgentError::LLM(
4304 "Provider returned a tool call outside the effective grant".to_string(),
4305 ));
4306 }
4307
4308 let compliant = match protocol.choice.as_ref() {
4309 Some(ToolChoice::Required) => calls.as_ref().is_some_and(|calls| !calls.is_empty()),
4310 Some(ToolChoice::Specific(expected)) => calls.as_ref().is_some_and(|calls| {
4311 !calls.is_empty()
4312 && calls.iter().all(|call| {
4313 self.tools.canonical_id(&call.name).as_deref() == Some(expected.as_str())
4314 })
4315 }),
4316 _ => true,
4317 };
4318 Ok((response, compliant))
4319 }
4320
4321 async fn complete_main_llm_with_recovery(
4322 &self,
4323 llm: Arc<dyn LLMProvider>,
4324 messages: &[ChatMessage],
4325 protocol: &MainToolProtocol,
4326 ) -> Result<LLMResponse> {
4327 let first = self
4328 .complete_main_attempt_with_recovery(Arc::clone(&llm), messages, protocol, false)
4329 .await?;
4330 let (response, compliant) =
4331 self.normalize_main_provider_response(first.response, protocol)?;
4332 if compliant {
4333 return Ok(response);
4334 }
4335 if first.used_native_tools {
4336 return Err(AgentError::LLM(
4337 "Provider returned no compliant native call for required tool choice".to_string(),
4338 ));
4339 }
4340
4341 let corrected = self
4342 .complete_main_attempt_with_recovery(llm, messages, protocol, true)
4343 .await?;
4344 let (response, compliant) =
4345 self.normalize_main_provider_response(corrected.response, protocol)?;
4346 if compliant {
4347 return Ok(response);
4348 }
4349 Err(AgentError::LLM(
4350 "Provider returned no compliant tool call after one corrective retry".to_string(),
4351 ))
4352 }
4353
4354 fn is_native_tool_call_content(content: &str) -> Result<bool> {
4356 decode_native_tool_call_markers(content)
4357 .map(|batch| batch.is_some())
4358 .map_err(|error| AgentError::LLM(error.to_string()))
4359 }
4360
4361 fn tool_result_message(
4363 tool_call: &ToolCall,
4364 output: &str,
4365 native_tool_call: bool,
4366 ) -> Result<ChatMessage> {
4367 if !native_tool_call {
4368 return Ok(ChatMessage::function(&tool_call.name, output));
4369 }
4370 let output = serde_json::from_str::<serde_json::Value>(output)
4371 .unwrap_or_else(|_| serde_json::Value::String(output.to_string()));
4372 let content = encode_native_tool_result_marker(tool_call, output)
4373 .map_err(|error| AgentError::LLM(error.to_string()))?;
4374 Ok(ChatMessage::function(&tool_call.name, content))
4375 }
4376
4377 fn remember_active_native_exchange(&self, content: &str) -> Result<()> {
4379 let Some(batch) = decode_native_tool_call_markers(content)
4380 .map_err(|error| AgentError::LLM(error.to_string()))?
4381 else {
4382 return Ok(());
4383 };
4384 let Some(state) = batch.provider_state() else {
4385 return Ok(());
4386 };
4387 let expected = ActiveNativeExchange {
4388 exchange_id: state.exchange_id().to_string(),
4389 call_ids: batch.calls().iter().map(|call| call.id.clone()).collect(),
4390 };
4391 let mut active = self.active_native_exchanges.write();
4392 if let Some(existing) = active
4393 .iter()
4394 .find(|existing| existing.exchange_id == expected.exchange_id)
4395 {
4396 if existing.call_ids != expected.call_ids {
4397 return Err(AgentError::LLM(format!(
4398 "Active native exchange '{}' changed its call identities",
4399 expected.exchange_id
4400 )));
4401 }
4402 } else {
4403 active.push(expected);
4404 }
4405 Ok(())
4406 }
4407
4408 fn validate_active_native_history(
4410 &self,
4411 messages: &[ChatMessage],
4412 require_complete: bool,
4413 ) -> Result<()> {
4414 let expected = self.active_native_exchanges.read().clone();
4415 if expected.is_empty() {
4416 return Ok(());
4417 }
4418 let inspection =
4419 inspect_native_history(messages).map_err(|error| AgentError::LLM(error.to_string()))?;
4420 let expected_count = expected.len();
4421 for (index, expected) in expected.iter().enumerate() {
4422 let Some(exchange) = inspection
4423 .exchanges()
4424 .iter()
4425 .find(|exchange| exchange.state().exchange_id() == expected.exchange_id)
4426 else {
4427 return Err(AgentError::LLM(format!(
4428 "Active native exchange '{}' was removed before provider continuation",
4429 expected.exchange_id
4430 )));
4431 };
4432 let must_be_complete = require_complete || index + 1 < expected_count;
4433 if exchange.call_ids() != expected.call_ids
4434 || (must_be_complete && !exchange.is_complete())
4435 {
4436 return Err(AgentError::LLM(format!(
4437 "Active native exchange '{}' is incomplete before provider continuation",
4438 expected.exchange_id
4439 )));
4440 }
4441 }
4442 Ok(())
4443 }
4444
4445 async fn remember_committed_native_exchange(&self, content: &str) -> Result<()> {
4447 self.remember_active_native_exchange(content)?;
4448 if !self.active_native_exchanges.read().is_empty() {
4449 let messages = self.memory.get_messages(None).await?;
4450 self.validate_active_native_history(&messages, false)?;
4451 }
4452 Ok(())
4453 }
4454
4455 fn readable_native_messages(mut messages: Vec<ChatMessage>) -> Result<Vec<ChatMessage>> {
4457 for message in &mut messages {
4458 if matches!(
4459 message.role,
4460 ai_agents_core::Role::Assistant
4461 | ai_agents_core::Role::Tool
4462 | ai_agents_core::Role::Function
4463 ) {
4464 message.content = native_readable_projection(&message.content)
4465 .map_err(|error| AgentError::LLM(error.to_string()))?;
4466 }
4467 }
4468 Ok(messages)
4469 }
4470
4471 fn parse_main_tool_calls(
4473 &self,
4474 content: &str,
4475 protocol: &MainToolProtocol,
4476 ) -> Result<Option<Vec<ToolCall>>> {
4477 if matches!(protocol.choice.as_ref(), Some(ToolChoice::None)) {
4478 Ok(None)
4479 } else {
4480 self.parse_tool_calls(content)
4481 }
4482 }
4483
4484 fn parse_tool_calls(&self, content: &str) -> Result<Option<Vec<ToolCall>>> {
4486 if let Some(batch) = decode_native_tool_call_markers(content)
4487 .map_err(|error| AgentError::LLM(error.to_string()))?
4488 {
4489 return Ok(Some(batch.into_parts().0));
4490 }
4491 if let Ok(parsed) = serde_json::from_str::<serde_json::Value>(content) {
4493 if let Some(arr) = parsed.as_array() {
4495 let calls: Vec<ToolCall> = arr
4496 .iter()
4497 .filter_map(|v| self.extract_tool_call_from_value(v))
4498 .collect();
4499 if !calls.is_empty() {
4500 return Ok(Some(calls));
4501 }
4502 }
4503 if let Some(tool_call) = self.extract_tool_call_from_value(&parsed) {
4505 return Ok(Some(vec![tool_call]));
4506 }
4507 }
4508
4509 if let Some(json_str) = self.extract_json_from_content(content)
4511 && let Ok(parsed) = serde_json::from_str::<serde_json::Value>(&json_str)
4512 {
4513 if let Some(arr) = parsed.as_array() {
4515 let calls: Vec<ToolCall> = arr
4516 .iter()
4517 .filter_map(|v| self.extract_tool_call_from_value(v))
4518 .collect();
4519 if !calls.is_empty() {
4520 return Ok(Some(calls));
4521 }
4522 }
4523 if let Some(tool_call) = self.extract_tool_call_from_value(&parsed) {
4525 return Ok(Some(vec![tool_call]));
4526 }
4527 }
4528
4529 Ok(None)
4530 }
4531
4532 fn extract_tool_call_from_value(&self, parsed: &serde_json::Value) -> Option<ToolCall> {
4533 if let Some(tool_name) = parsed.get("tool").and_then(|v| v.as_str()) {
4534 let arguments = parsed
4535 .get("arguments")
4536 .cloned()
4537 .unwrap_or(serde_json::json!({}));
4538 return Some(ToolCall {
4539 id: parsed
4540 .get("id")
4541 .and_then(|value| value.as_str())
4542 .filter(|id| !id.is_empty())
4543 .map(str::to_string)
4544 .unwrap_or_else(|| uuid::Uuid::new_v4().to_string()),
4545 name: tool_name.to_string(),
4546 arguments,
4547 });
4548 }
4549 None
4550 }
4551
4552 fn extract_json_from_content(&self, content: &str) -> Option<String> {
4554 if let Some(result) = self.extract_json_array_from_content(content) {
4556 return Some(result);
4557 }
4558 self.extract_json_object_from_content(content)
4559 }
4560
4561 fn extract_json_array_from_content(&self, content: &str) -> Option<String> {
4563 let start = content.find('[')?;
4564 let content_from_start = &content[start..];
4565
4566 let mut depth = 0;
4567 let mut end = 0;
4568 for (i, ch) in content_from_start.char_indices() {
4569 match ch {
4570 '[' => depth += 1,
4571 ']' => {
4572 depth -= 1;
4573 if depth == 0 {
4574 end = i + 1;
4575 break;
4576 }
4577 }
4578 _ => {}
4579 }
4580 }
4581
4582 if end > 0 {
4583 let json_str = &content_from_start[..end];
4584 if json_str.contains("\"tool\"") {
4586 return Some(json_str.to_string());
4587 }
4588 }
4589
4590 None
4591 }
4592
4593 fn extract_json_object_from_content(&self, content: &str) -> Option<String> {
4595 let start = content.find('{')?;
4596 let content_from_start = &content[start..];
4597
4598 let mut depth = 0;
4600 let mut end = 0;
4601 for (i, ch) in content_from_start.char_indices() {
4602 match ch {
4603 '{' => depth += 1,
4604 '}' => {
4605 depth -= 1;
4606 if depth == 0 {
4607 end = i + 1;
4608 break;
4609 }
4610 }
4611 _ => {}
4612 }
4613 }
4614
4615 if end > 0 {
4616 let json_str = &content_from_start[..end];
4617 if json_str.contains("\"tool\"") {
4619 return Some(json_str.to_string());
4620 }
4621 }
4622
4623 None
4624 }
4625
4626 #[allow(clippy::too_many_arguments)]
4630 fn record_from_parts(
4631 &self,
4632 request: &ToolExecutionRequest,
4633 canonical_id: String,
4634 executed_arguments: Value,
4635 started_at: chrono::DateTime<chrono::Utc>,
4636 start: Instant,
4637 executed: bool,
4638 success: bool,
4639 output: String,
4640 metadata: HashMap<String, Value>,
4641 policy: ToolPolicyDecisionRecord,
4642 approval: Option<ToolApprovalRecord>,
4643 timed_out: bool,
4644 output_truncated: bool,
4645 ) -> ToolExecutionRecord {
4646 let versions = ToolDecisionVersions {
4647 policy: self.active_tool_security().policy_version(),
4648 registry: self.tools.version(),
4649 runtime_control: self.runtime_control.version.load(Ordering::SeqCst),
4650 state: self
4651 .state_machine
4652 .as_ref()
4653 .map(|state_machine| state_machine.generation()),
4654 };
4655 self.record_from_parts_at(
4656 request,
4657 canonical_id,
4658 executed_arguments,
4659 started_at,
4660 start,
4661 executed,
4662 success,
4663 output,
4664 metadata,
4665 policy,
4666 approval,
4667 timed_out,
4668 output_truncated,
4669 versions,
4670 )
4671 }
4672
4673 #[allow(clippy::too_many_arguments)]
4675 fn record_from_parts_at(
4676 &self,
4677 request: &ToolExecutionRequest,
4678 canonical_id: String,
4679 executed_arguments: Value,
4680 started_at: chrono::DateTime<chrono::Utc>,
4681 start: Instant,
4682 executed: bool,
4683 success: bool,
4684 output: String,
4685 metadata: HashMap<String, Value>,
4686 policy: ToolPolicyDecisionRecord,
4687 approval: Option<ToolApprovalRecord>,
4688 timed_out: bool,
4689 output_truncated: bool,
4690 versions: ToolDecisionVersions,
4691 ) -> ToolExecutionRecord {
4692 ToolExecutionRecord {
4693 call_id: request.call_id.clone(),
4694 requested_name: request.requested_name.clone(),
4695 canonical_id,
4696 source: request.source.clone(),
4697 arguments: request.arguments.clone(),
4698 executed_arguments,
4699 policy_version: versions.policy,
4700 registry_version: versions.registry,
4701 runtime_config_version: versions.runtime_control,
4702 executed,
4703 success,
4704 output,
4705 metadata,
4706 policy,
4707 approval,
4708 started_at,
4709 duration_ms: start.elapsed().as_millis() as u64,
4710 timed_out,
4711 cancelled: false,
4712 cancellation_reason: None,
4713 output_truncated,
4714 }
4715 }
4716
4717 async fn finish_tool_record(&self, record: &ToolExecutionRecord) {
4719 let result = ToolResult {
4720 success: record.success,
4721 output: record.model_output_string(),
4722 metadata: if record.metadata.is_empty() {
4723 None
4724 } else {
4725 Some(record.metadata.clone())
4726 },
4727 };
4728 self.hooks
4729 .on_tool_complete(&record.canonical_id, &result, record.duration_ms)
4730 .await;
4731 self.hooks.on_tool_execution_record(record).await;
4732 self.record_tool_call(&record.canonical_id, record.model_output_value());
4733 if !record.success {
4734 self.hooks
4735 .on_error(&AgentError::Tool(record.output.clone()))
4736 .await;
4737 }
4738 }
4739
4740 async fn finish_tool_record_after_resource_guards(
4742 &self,
4743 resource_guards: ToolResourceGuards,
4744 record: &ToolExecutionRecord,
4745 ) {
4746 drop(resource_guards);
4747 self.finish_tool_record(record).await;
4748 }
4749
4750 fn validated_tool_timeout(timeout_ms: u64) -> Result<ValidatedToolTimeout> {
4754 if timeout_ms > MAX_TOOL_TIMEOUT_MS {
4755 return Err(AgentError::Config(format!(
4756 "effective tool timeout_ms must be no greater than {MAX_TOOL_TIMEOUT_MS} milliseconds"
4757 )));
4758 }
4759 let timer = Duration::from_millis(timeout_ms);
4760 let deadline_delta = chrono::Duration::from_std(timer).map_err(|_| {
4761 AgentError::Config(format!(
4762 "effective tool timeout_ms cannot be represented as a UTC deadline: {timeout_ms}"
4763 ))
4764 })?;
4765 Ok(ValidatedToolTimeout {
4766 timer,
4767 deadline_delta,
4768 })
4769 }
4770
4771 fn effective_tool_limits(
4775 security_engine: &ToolSecurityEngine,
4776 canonical_id: &str,
4777 safety: &ToolSafetyMetadata,
4778 classification: &ToolCallClassification,
4779 recovery_timeout_ms: Option<u64>,
4780 ) -> Result<(ToolExecutionLimits, ValidatedToolTimeout)> {
4781 if let Some(timeout_ms) = classification.timeout_ms {
4782 Self::validated_tool_timeout(timeout_ms)?;
4783 }
4784 if let Some(timeout_ms) = recovery_timeout_ms {
4785 Self::validated_tool_timeout(timeout_ms)?;
4786 }
4787
4788 let mut limits = security_engine.effective_limits(canonical_id, safety, classification);
4789 if let Some(recovery_timeout_ms) = recovery_timeout_ms {
4790 limits.timeout_ms = Some(limits.timeout_ms.map_or(recovery_timeout_ms, |timeout_ms| {
4791 timeout_ms.min(recovery_timeout_ms)
4792 }));
4793 }
4794 let timeout_ms = limits
4795 .timeout_ms
4796 .unwrap_or_else(|| security_engine.get_tool_timeout(canonical_id));
4797 let timeout = Self::validated_tool_timeout(timeout_ms)?;
4798 Ok((limits, timeout))
4799 }
4800
4801 async fn execute_resolved_tool_once(
4803 &self,
4804 tool: Arc<dyn ai_agents_core::Tool>,
4805 args: Value,
4806 mut ctx: ToolExecutionContext,
4807 timeout: ValidatedToolTimeout,
4808 ) -> Result<(ToolResult, bool, bool, bool)> {
4809 if self.runtime_control.emergency_deny.load(Ordering::SeqCst) {
4810 return Ok((
4811 ToolResult::error("Tool execution cancelled by runtime control"),
4812 false,
4813 true,
4814 false,
4815 ));
4816 }
4817 ctx.deadline = Some(
4822 chrono::Utc::now()
4823 .checked_add_signed(timeout.deadline_delta)
4824 .ok_or_else(|| {
4825 AgentError::Config(
4826 "effective tool timeout_ms exceeds the current UTC deadline range"
4827 .to_string(),
4828 )
4829 })?,
4830 );
4831 let invoked = Arc::new(AtomicBool::new(false));
4835 let invoked_by_future = Arc::clone(&invoked);
4836 let actor_context = current_turn_actor_context();
4837 let future = async move {
4838 invoked_by_future.store(true, Ordering::SeqCst);
4839 if let Some(actor_context) = actor_context {
4840 scope_actor_context(actor_context, tool.execute(args, ctx)).await
4841 } else {
4842 tool.execute(args, ctx).await
4843 }
4844 };
4845 tokio::pin!(future);
4846 let timer = tokio::time::sleep(timeout.timer);
4847 tokio::pin!(timer);
4848 let mut cancel_tick = tokio::time::interval(std::time::Duration::from_millis(50));
4849
4850 loop {
4851 tokio::select! {
4852 result = &mut future => return Ok((result, false, false, true)),
4853 _ = &mut timer => {
4854 return Ok((
4855 ToolResult::error("Tool execution timed out"),
4856 true,
4857 false,
4858 invoked.load(Ordering::SeqCst),
4859 ));
4860 }
4861 _ = cancel_tick.tick() => {
4862 if self.runtime_control.emergency_deny.load(Ordering::SeqCst) {
4863 return Ok((
4864 ToolResult::error("Tool execution cancelled by runtime control"),
4865 false,
4866 true,
4867 invoked.load(Ordering::SeqCst),
4868 ));
4869 }
4870 }
4871 }
4872 }
4873 }
4874
4875 fn truncate_tool_output(output: String, max_chars: Option<usize>) -> (String, bool) {
4877 let Some(max_chars) = max_chars else {
4878 return (output, false);
4879 };
4880 let mut chars = output.chars();
4881 let truncated: String = chars.by_ref().take(max_chars).collect();
4882 if chars.next().is_some() {
4883 (truncated, true)
4884 } else {
4885 (output, false)
4886 }
4887 }
4888
4889 async fn acquire_tool_resource_locks(&self, keys: &[String]) -> Option<ToolResourceGuards> {
4891 let locks = {
4892 let mut table = self.resource_locks.write();
4893 table.retain(|_, lock| lock.strong_count() > 0);
4894 keys.iter()
4895 .map(|key| {
4896 if let Some(lock) = table.get(key).and_then(Weak::upgrade) {
4897 lock
4898 } else {
4899 let lock = Arc::new(tokio::sync::Mutex::new(()));
4900 table.insert(key.clone(), Arc::downgrade(&lock));
4901 lock
4902 }
4903 })
4904 .collect::<Vec<_>>()
4905 };
4906 let mut resource_guards = ToolResourceGuards {
4907 guards: Vec::with_capacity(locks.len()),
4908 locks: Arc::clone(&self.resource_locks),
4909 };
4910 let mut locks = locks.into_iter();
4911 while let Some(lock) = locks.next() {
4912 let mut lock = Box::pin(lock.lock_owned());
4913 loop {
4914 tokio::select! {
4915 guard = &mut lock => {
4916 resource_guards.guards.push(guard);
4917 break;
4918 }
4919 _ = tokio::time::sleep(std::time::Duration::from_millis(10)) => {
4920 if self.runtime_control.emergency_deny.load(Ordering::SeqCst) {
4921 drop(lock);
4922 drop(locks);
4923 drop(resource_guards);
4924 return None;
4925 }
4926 }
4927 }
4928 }
4929 }
4930 Some(resource_guards)
4931 }
4932
4933 async fn run_tool_with_retries(
4937 &self,
4938 canonical_id: &str,
4939 tool: Arc<dyn ai_agents_core::Tool>,
4940 args: Value,
4941 ctx: ToolExecutionContext,
4942 timeout: ValidatedToolTimeout,
4943 max_retries: u32,
4944 ) -> Result<(ToolResult, bool, bool, bool)> {
4945 let max_retries = if ctx.classification.safely_retryable {
4946 max_retries
4947 } else {
4948 0
4949 };
4950 let mut attempts = 0;
4951 let mut invoked = false;
4952 loop {
4953 let (result, timed_out, cancelled, attempt_invoked) = self
4954 .execute_resolved_tool_once(tool.clone(), args.clone(), ctx.clone(), timeout)
4955 .await?;
4956 invoked |= attempt_invoked;
4957 if result.success || timed_out || cancelled || attempts >= max_retries {
4958 return Ok((result, timed_out, cancelled, invoked));
4959 }
4960 attempts += 1;
4961 warn!(tool = %canonical_id, attempt = attempts, error = %result.output, "Retrying failed tool call");
4962 }
4963 }
4964
4965 fn host_tool_unavailability(&self, canonical_id: &str) -> Option<(&'static str, &'static str)> {
4967 match canonical_id {
4968 "command" if !self.tools.command_runner_available() => Some((
4969 "Command runner is unavailable",
4970 "command runner is unavailable",
4971 )),
4972 "diagnostics" if !self.tools.diagnostics_available() => Some((
4973 "Diagnostics provider is unavailable",
4974 "diagnostics provider is unavailable",
4975 )),
4976 "web_search" if !self.tools.web_search_available() => Some((
4977 "Web search provider is unavailable",
4978 "web search provider is unavailable",
4979 )),
4980 _ => None,
4981 }
4982 }
4983
4984 fn execute_tool_record(
4986 &self,
4987 request: ToolExecutionRequest,
4988 ) -> Pin<Box<dyn Future<Output = Result<ToolExecutionRecord>> + Send + '_>> {
4989 Box::pin(self.execute_tool_record_inner(request, ToolFallbackState::default()))
4990 }
4991
4992 async fn execute_tool_record_inner(
4996 &self,
4997 request: ToolExecutionRequest,
4998 fallback_state: ToolFallbackState,
4999 ) -> Result<ToolExecutionRecord> {
5000 let started_at = chrono::Utc::now();
5001 let start = Instant::now();
5002 info!(tool = %request.requested_name, args = %request.arguments, "Executing tool");
5003
5004 if self.runtime_control.emergency_deny.load(Ordering::SeqCst) {
5005 let record = self.record_from_parts(
5006 &request,
5007 request.requested_name.clone(),
5008 request.arguments.clone(),
5009 started_at,
5010 start,
5011 false,
5012 false,
5013 "Tool execution is disabled by runtime control".to_string(),
5014 HashMap::new(),
5015 ToolPolicyDecisionRecord::deny("runtime emergency deny is enabled"),
5016 None,
5017 false,
5018 false,
5019 );
5020 self.finish_tool_record(&record).await;
5021 return Ok(record);
5022 }
5023
5024 let Some(resolved) = self.tools.resolve(&request.requested_name) else {
5025 let record = self.record_from_parts(
5026 &request,
5027 request.requested_name.clone(),
5028 request.arguments.clone(),
5029 started_at,
5030 start,
5031 false,
5032 false,
5033 format!("Tool '{}' is unavailable", request.requested_name),
5034 HashMap::new(),
5035 ToolPolicyDecisionRecord::unavailable(format!(
5036 "Tool '{}' is not registered",
5037 request.requested_name
5038 )),
5039 None,
5040 false,
5041 false,
5042 );
5043 self.finish_tool_record(&record).await;
5044 return Ok(record);
5045 };
5046
5047 let canonical_id = resolved.identity.canonical_id.clone();
5048
5049 let initial_scope_snapshot = self.get_available_tool_ids_snapshot().await?;
5050 if !initial_scope_snapshot
5051 .tool_ids
5052 .iter()
5053 .any(|id| id == &canonical_id)
5054 {
5055 let record = self.record_from_parts(
5056 &request,
5057 canonical_id.clone(),
5058 request.arguments.clone(),
5059 started_at,
5060 start,
5061 false,
5062 false,
5063 format!(
5064 "Tool '{}' is not available in the current scope",
5065 canonical_id
5066 ),
5067 HashMap::new(),
5068 ToolPolicyDecisionRecord::deny(format!(
5069 "Tool '{}' is not granted by the current top-level and state tool scope",
5070 canonical_id
5071 )),
5072 None,
5073 false,
5074 false,
5075 );
5076 self.finish_tool_record(&record).await;
5077 return Ok(record);
5078 }
5079
5080 let approval_control_snapshot = self.runtime_safety_snapshot();
5081 let security_engine = approval_control_snapshot.tool_security.clone();
5082 if let Some(reason) = fallback_state.rejection_reason(&canonical_id) {
5083 let mut metadata = HashMap::new();
5084 metadata.insert(
5085 "fallback_chain".to_string(),
5086 serde_json::to_value(&fallback_state.visited_canonical_ids).unwrap_or(Value::Null),
5087 );
5088 let record = self.record_from_parts(
5089 &request,
5090 canonical_id,
5091 request.arguments.clone(),
5092 started_at,
5093 start,
5094 false,
5095 false,
5096 format!("Denied: {reason}"),
5097 metadata,
5098 ToolPolicyDecisionRecord::deny(reason),
5099 None,
5100 false,
5101 false,
5102 );
5103 self.finish_tool_record(&record).await;
5104 return Ok(record);
5105 }
5106 let admitted_canonical_id = canonical_id.clone();
5107 let fallback_state = fallback_state.with_current(canonical_id.clone());
5108 let bindings = resolved.tool.policy_bindings();
5109 let mut executed_arguments = security_engine.prepare_tool_arguments_with_bindings(
5110 &canonical_id,
5111 &request.arguments,
5112 &bindings,
5113 );
5114 let mut metadata = HashMap::new();
5115 let safety = resolved.tool.safety_metadata();
5116 let classification = resolved.tool.classify_call(&executed_arguments);
5117 let initial_recovery_timeout_ms = self.recovery_manager.get_tool_timeout(&canonical_id);
5118 let (limits, _) = Self::effective_tool_limits(
5119 &security_engine,
5120 &canonical_id,
5121 &safety,
5122 &classification,
5123 initial_recovery_timeout_ms,
5124 )?;
5125 self.hooks
5126 .on_tool_start(&canonical_id, &executed_arguments)
5127 .await;
5128 metadata.insert(
5129 "classification".to_string(),
5130 serde_json::to_value(&classification).unwrap_or(Value::Null),
5131 );
5132 metadata.insert(
5133 "effective_limits".to_string(),
5134 serde_json::to_value(&limits).unwrap_or(Value::Null),
5135 );
5136 let policy_snapshot = security_engine.policy_snapshot(&canonical_id);
5137 if !policy_snapshot.is_null() {
5138 metadata.insert("policy_snapshot".to_string(), policy_snapshot.clone());
5139 }
5140
5141 let mut approval_record = Some(ToolApprovalRecord {
5142 status: ToolApprovalStatus::NotRequired,
5143 reason: None,
5144 modified_arguments: None,
5145 });
5146
5147 let mut security_result = security_engine
5148 .validate_tool_execution_with_bindings(&canonical_id, &executed_arguments, &bindings)
5149 .await?;
5150 if (security_result.is_allowed()
5155 || matches!(
5156 &security_result,
5157 SecurityCheckResult::RequireConfirmation { .. }
5158 ))
5159 && let Some((output, reason)) = self.host_tool_unavailability(&canonical_id)
5160 {
5161 let record = self.record_from_parts(
5162 &request,
5163 canonical_id,
5164 executed_arguments,
5165 started_at,
5166 start,
5167 false,
5168 false,
5169 output.to_string(),
5170 metadata,
5171 ToolPolicyDecisionRecord::unavailable(reason),
5172 Some(ToolApprovalRecord {
5173 status: ToolApprovalStatus::Unavailable,
5174 reason: Some(reason.to_string()),
5175 modified_arguments: None,
5176 }),
5177 false,
5178 false,
5179 );
5180 self.finish_tool_record(&record).await;
5181 return Ok(record);
5182 }
5183 match &security_result {
5184 SecurityCheckResult::Allow => {}
5185 SecurityCheckResult::Warn { message } => {
5186 warn!(tool = %canonical_id, message = %message, "Tool security warning");
5187 }
5188 SecurityCheckResult::Block { reason } => {
5189 let record = self.record_from_parts(
5190 &request,
5191 canonical_id,
5192 executed_arguments,
5193 started_at,
5194 start,
5195 false,
5196 false,
5197 format!("Denied: {}", reason),
5198 metadata,
5199 ToolPolicyDecisionRecord::deny(reason.clone()),
5200 approval_record,
5201 false,
5202 false,
5203 );
5204 self.finish_tool_record(&record).await;
5205 return Ok(record);
5206 }
5207 SecurityCheckResult::Unavailable { reason } => {
5208 let record = self.record_from_parts(
5209 &request,
5210 canonical_id,
5211 executed_arguments,
5212 started_at,
5213 start,
5214 false,
5215 false,
5216 format!("Unavailable: {}", reason),
5217 metadata,
5218 ToolPolicyDecisionRecord::unavailable(reason.clone()),
5219 approval_record,
5220 false,
5221 false,
5222 );
5223 self.finish_tool_record(&record).await;
5224 return Ok(record);
5225 }
5226 SecurityCheckResult::RequireConfirmation { message } => {
5227 if self.hitl_engine.is_none() {
5228 approval_record = Some(ToolApprovalRecord {
5229 status: ToolApprovalStatus::Unavailable,
5230 reason: Some("No HITL engine configured".to_string()),
5231 modified_arguments: None,
5232 });
5233 let record = self.record_from_parts(
5234 &request,
5235 canonical_id,
5236 executed_arguments,
5237 started_at,
5238 start,
5239 false,
5240 false,
5241 format!("Approval unavailable: {}", message),
5242 metadata,
5243 ToolPolicyDecisionRecord::approval(message.clone()),
5244 approval_record,
5245 false,
5246 false,
5247 );
5248 self.finish_tool_record(&record).await;
5249 return Ok(record);
5250 }
5251
5252 let check_result = HITLCheckResult::required(
5253 ApprovalTrigger::tool(&canonical_id, executed_arguments.clone()),
5254 HashMap::new(),
5255 message.clone(),
5256 None,
5257 );
5258 match self.request_hitl_approval(check_result).await? {
5259 ApprovalResult::Approved => {
5260 merge_approved_record(&mut approval_record);
5261 }
5262 ApprovalResult::Modified { changes } => {
5263 if let Some(obj) = executed_arguments.as_object_mut() {
5264 for (key, value) in changes {
5265 obj.insert(key, value);
5266 }
5267 }
5268 security_result = security_engine
5269 .validate_tool_execution_with_bindings(
5270 &canonical_id,
5271 &executed_arguments,
5272 &bindings,
5273 )
5274 .await?;
5275 if !matches!(
5276 security_result,
5277 SecurityCheckResult::Allow
5278 | SecurityCheckResult::Warn { .. }
5279 | SecurityCheckResult::RequireConfirmation { .. }
5280 ) {
5281 let reason = security_result
5282 .reason()
5283 .unwrap_or("modified arguments failed policy")
5284 .to_string();
5285 let record = self.record_from_parts(
5286 &request,
5287 canonical_id,
5288 executed_arguments.clone(),
5289 started_at,
5290 start,
5291 false,
5292 false,
5293 reason.clone(),
5294 metadata,
5295 ToolPolicyDecisionRecord::deny(reason),
5296 Some(ToolApprovalRecord {
5297 status: ToolApprovalStatus::Modified,
5298 reason: None,
5299 modified_arguments: Some(executed_arguments),
5300 }),
5301 false,
5302 false,
5303 );
5304 self.finish_tool_record(&record).await;
5305 return Ok(record);
5306 }
5307 approval_record = Some(ToolApprovalRecord {
5308 status: ToolApprovalStatus::Modified,
5309 reason: None,
5310 modified_arguments: Some(executed_arguments.clone()),
5311 });
5312 }
5313 ApprovalResult::Rejected { reason } => {
5314 let reason = reason.unwrap_or_else(|| "rejected".to_string());
5315 approval_record = Some(ToolApprovalRecord {
5316 status: ToolApprovalStatus::Rejected,
5317 reason: Some(reason.clone()),
5318 modified_arguments: None,
5319 });
5320 let record = self.record_from_parts(
5321 &request,
5322 canonical_id,
5323 executed_arguments,
5324 started_at,
5325 start,
5326 false,
5327 false,
5328 format!("Approval rejected: {}", reason),
5329 metadata,
5330 ToolPolicyDecisionRecord::approval(reason),
5331 approval_record,
5332 false,
5333 false,
5334 );
5335 self.finish_tool_record(&record).await;
5336 return Ok(record);
5337 }
5338 ApprovalResult::Timeout => {
5339 approval_record = Some(ToolApprovalRecord {
5340 status: ToolApprovalStatus::Timeout,
5341 reason: Some("approval timeout".to_string()),
5342 modified_arguments: None,
5343 });
5344 let record = self.record_from_parts(
5345 &request,
5346 canonical_id,
5347 executed_arguments,
5348 started_at,
5349 start,
5350 false,
5351 false,
5352 "Approval timed out".to_string(),
5353 metadata,
5354 ToolPolicyDecisionRecord::approval("approval timeout"),
5355 approval_record,
5356 false,
5357 false,
5358 );
5359 self.finish_tool_record(&record).await;
5360 return Ok(record);
5361 }
5362 }
5363 }
5364 }
5365
5366 if approval_record
5367 .as_ref()
5368 .is_some_and(|record| matches!(record.status, ToolApprovalStatus::NotRequired))
5369 && let Some(message) =
5370 security_engine.classification_approval_message(&canonical_id, &classification)
5371 {
5372 if self.hitl_engine.is_none() {
5373 approval_record = Some(ToolApprovalRecord {
5374 status: ToolApprovalStatus::Unavailable,
5375 reason: Some("No HITL engine configured".to_string()),
5376 modified_arguments: None,
5377 });
5378 let record = self.record_from_parts(
5379 &request,
5380 canonical_id,
5381 executed_arguments,
5382 started_at,
5383 start,
5384 false,
5385 false,
5386 format!("Approval unavailable: {}", message),
5387 metadata,
5388 ToolPolicyDecisionRecord::approval(message),
5389 approval_record,
5390 false,
5391 false,
5392 );
5393 self.finish_tool_record(&record).await;
5394 return Ok(record);
5395 }
5396 let check_result = HITLCheckResult::required(
5397 ApprovalTrigger::tool(&canonical_id, executed_arguments.clone()),
5398 HashMap::new(),
5399 message.clone(),
5400 None,
5401 );
5402 match self.request_hitl_approval(check_result).await? {
5403 ApprovalResult::Approved => {
5404 merge_approved_record(&mut approval_record);
5405 }
5406 ApprovalResult::Modified { changes } => {
5407 if let Some(obj) = executed_arguments.as_object_mut() {
5408 for (key, value) in changes {
5409 obj.insert(key, value);
5410 }
5411 }
5412 let modified_security = security_engine
5413 .validate_tool_execution_with_bindings(
5414 &canonical_id,
5415 &executed_arguments,
5416 &bindings,
5417 )
5418 .await?;
5419 if !matches!(
5420 modified_security,
5421 SecurityCheckResult::Allow | SecurityCheckResult::Warn { .. }
5422 ) {
5423 let reason = modified_security
5424 .reason()
5425 .unwrap_or("modified arguments failed policy")
5426 .to_string();
5427 let record = self.record_from_parts(
5428 &request,
5429 canonical_id,
5430 executed_arguments.clone(),
5431 started_at,
5432 start,
5433 false,
5434 false,
5435 reason.clone(),
5436 metadata,
5437 ToolPolicyDecisionRecord::deny(reason),
5438 Some(ToolApprovalRecord {
5439 status: ToolApprovalStatus::Modified,
5440 reason: None,
5441 modified_arguments: Some(executed_arguments),
5442 }),
5443 false,
5444 false,
5445 );
5446 self.finish_tool_record(&record).await;
5447 return Ok(record);
5448 }
5449 approval_record = Some(ToolApprovalRecord {
5450 status: ToolApprovalStatus::Modified,
5451 reason: None,
5452 modified_arguments: Some(executed_arguments.clone()),
5453 });
5454 }
5455 ApprovalResult::Rejected { reason } => {
5456 let reason = reason.unwrap_or_else(|| "rejected".to_string());
5457 let record = self.record_from_parts(
5458 &request,
5459 canonical_id,
5460 executed_arguments,
5461 started_at,
5462 start,
5463 false,
5464 false,
5465 format!("Approval rejected: {}", reason),
5466 metadata,
5467 ToolPolicyDecisionRecord::approval(reason.clone()),
5468 Some(ToolApprovalRecord {
5469 status: ToolApprovalStatus::Rejected,
5470 reason: Some(reason),
5471 modified_arguments: None,
5472 }),
5473 false,
5474 false,
5475 );
5476 self.finish_tool_record(&record).await;
5477 return Ok(record);
5478 }
5479 ApprovalResult::Timeout => {
5480 let record = self.record_from_parts(
5481 &request,
5482 canonical_id,
5483 executed_arguments,
5484 started_at,
5485 start,
5486 false,
5487 false,
5488 "Approval timed out".to_string(),
5489 metadata,
5490 ToolPolicyDecisionRecord::approval("approval timeout"),
5491 Some(ToolApprovalRecord {
5492 status: ToolApprovalStatus::Timeout,
5493 reason: Some("approval timeout".to_string()),
5494 modified_arguments: None,
5495 }),
5496 false,
5497 false,
5498 );
5499 self.finish_tool_record(&record).await;
5500 return Ok(record);
5501 }
5502 }
5503 }
5504
5505 let hitl_lang_ctx = self.build_hitl_language_context();
5506 if let Some(ref hitl_engine) = self.hitl_engine {
5507 let check_result = self
5508 .observe_purpose(
5509 ObservationPurpose::HitlLocalization,
5510 hitl_engine.check_tool_with_localization(
5511 &canonical_id,
5512 &executed_arguments,
5513 &hitl_lang_ctx,
5514 self.approval_handler.as_ref(),
5515 Some(&self.llm_registry),
5516 ),
5517 )
5518 .await?;
5519 if check_result.is_required() {
5520 match self.request_hitl_approval(check_result).await? {
5521 ApprovalResult::Approved => {
5522 merge_approved_record(&mut approval_record);
5523 }
5524 ApprovalResult::Modified { changes } => {
5525 if let Some(obj) = executed_arguments.as_object_mut() {
5526 for (key, value) in changes {
5527 obj.insert(key, value);
5528 }
5529 }
5530 let modified_security = security_engine
5531 .validate_tool_execution_with_bindings(
5532 &canonical_id,
5533 &executed_arguments,
5534 &bindings,
5535 )
5536 .await?;
5537 if !matches!(
5538 modified_security,
5539 SecurityCheckResult::Allow | SecurityCheckResult::Warn { .. }
5540 ) {
5541 let reason = modified_security
5542 .reason()
5543 .unwrap_or("modified arguments failed policy")
5544 .to_string();
5545 let record = self.record_from_parts(
5546 &request,
5547 canonical_id,
5548 executed_arguments.clone(),
5549 started_at,
5550 start,
5551 false,
5552 false,
5553 reason.clone(),
5554 metadata,
5555 ToolPolicyDecisionRecord::deny(reason),
5556 Some(ToolApprovalRecord {
5557 status: ToolApprovalStatus::Modified,
5558 reason: None,
5559 modified_arguments: Some(executed_arguments),
5560 }),
5561 false,
5562 false,
5563 );
5564 self.finish_tool_record(&record).await;
5565 return Ok(record);
5566 }
5567 approval_record = Some(ToolApprovalRecord {
5568 status: ToolApprovalStatus::Modified,
5569 reason: None,
5570 modified_arguments: Some(executed_arguments.clone()),
5571 });
5572 }
5573 ApprovalResult::Rejected { reason } => {
5574 let reason = reason.unwrap_or_else(|| "rejected".to_string());
5575 let record = self.record_from_parts(
5576 &request,
5577 canonical_id,
5578 executed_arguments,
5579 started_at,
5580 start,
5581 false,
5582 false,
5583 format!("Approval rejected: {}", reason),
5584 metadata,
5585 ToolPolicyDecisionRecord::approval(reason.clone()),
5586 Some(ToolApprovalRecord {
5587 status: ToolApprovalStatus::Rejected,
5588 reason: Some(reason),
5589 modified_arguments: None,
5590 }),
5591 false,
5592 false,
5593 );
5594 self.finish_tool_record(&record).await;
5595 return Ok(record);
5596 }
5597 ApprovalResult::Timeout => {
5598 let record = self.record_from_parts(
5599 &request,
5600 canonical_id,
5601 executed_arguments,
5602 started_at,
5603 start,
5604 false,
5605 false,
5606 "Approval timed out".to_string(),
5607 metadata,
5608 ToolPolicyDecisionRecord::approval("approval timeout"),
5609 Some(ToolApprovalRecord {
5610 status: ToolApprovalStatus::Timeout,
5611 reason: Some("approval timeout".to_string()),
5612 modified_arguments: None,
5613 }),
5614 false,
5615 false,
5616 );
5617 self.finish_tool_record(&record).await;
5618 return Ok(record);
5619 }
5620 }
5621 }
5622
5623 let condition_check = self
5624 .observe_purpose(
5625 ObservationPurpose::HitlLocalization,
5626 hitl_engine.check_conditions_with_localization(
5627 &executed_arguments,
5628 &hitl_lang_ctx,
5629 self.approval_handler.as_ref(),
5630 Some(&self.llm_registry),
5631 ),
5632 )
5633 .await?;
5634 if condition_check.is_required() {
5635 match self.request_hitl_approval(condition_check).await? {
5636 ApprovalResult::Approved => {
5637 merge_approved_record(&mut approval_record);
5638 }
5639 ApprovalResult::Modified { changes } => {
5640 if let Some(obj) = executed_arguments.as_object_mut() {
5641 for (key, value) in changes {
5642 obj.insert(key, value);
5643 }
5644 }
5645 let modified_security = security_engine
5646 .validate_tool_execution_with_bindings(
5647 &canonical_id,
5648 &executed_arguments,
5649 &bindings,
5650 )
5651 .await?;
5652 if !matches!(
5653 modified_security,
5654 SecurityCheckResult::Allow | SecurityCheckResult::Warn { .. }
5655 ) {
5656 let reason = modified_security
5657 .reason()
5658 .unwrap_or("modified arguments failed policy")
5659 .to_string();
5660 let record = self.record_from_parts(
5661 &request,
5662 canonical_id,
5663 executed_arguments,
5664 started_at,
5665 start,
5666 false,
5667 false,
5668 reason.clone(),
5669 metadata,
5670 ToolPolicyDecisionRecord::deny(reason),
5671 approval_record,
5672 false,
5673 false,
5674 );
5675 self.finish_tool_record(&record).await;
5676 return Ok(record);
5677 }
5678 approval_record = Some(ToolApprovalRecord {
5679 status: ToolApprovalStatus::Modified,
5680 reason: None,
5681 modified_arguments: Some(executed_arguments.clone()),
5682 });
5683 }
5684 ApprovalResult::Rejected { reason } => {
5685 let reason = reason.unwrap_or_else(|| "rejected".to_string());
5686 let record = self.record_from_parts(
5687 &request,
5688 canonical_id,
5689 executed_arguments,
5690 started_at,
5691 start,
5692 false,
5693 false,
5694 format!("Approval rejected: {}", reason),
5695 metadata,
5696 ToolPolicyDecisionRecord::approval(reason.clone()),
5697 Some(ToolApprovalRecord {
5698 status: ToolApprovalStatus::Rejected,
5699 reason: Some(reason),
5700 modified_arguments: None,
5701 }),
5702 false,
5703 false,
5704 );
5705 self.finish_tool_record(&record).await;
5706 return Ok(record);
5707 }
5708 ApprovalResult::Timeout => {
5709 let record = self.record_from_parts(
5710 &request,
5711 canonical_id,
5712 executed_arguments,
5713 started_at,
5714 start,
5715 false,
5716 false,
5717 "Approval timed out".to_string(),
5718 metadata,
5719 ToolPolicyDecisionRecord::approval("approval timeout"),
5720 Some(ToolApprovalRecord {
5721 status: ToolApprovalStatus::Timeout,
5722 reason: Some("approval timeout".to_string()),
5723 modified_arguments: None,
5724 }),
5725 false,
5726 false,
5727 );
5728 self.finish_tool_record(&record).await;
5729 return Ok(record);
5730 }
5731 }
5732 }
5733 }
5734
5735 executed_arguments = security_engine.prepare_tool_arguments_with_bindings(
5740 &canonical_id,
5741 &executed_arguments,
5742 &bindings,
5743 );
5744 if let Some(record) = approval_record.as_mut()
5745 && matches!(record.status, ToolApprovalStatus::Modified)
5746 {
5747 record.modified_arguments = Some(executed_arguments.clone());
5748 }
5749 let binding_security_result = security_engine
5750 .validate_tool_execution_with_bindings(&canonical_id, &executed_arguments, &bindings)
5751 .await?;
5752 let approval_confirmation_required = matches!(
5753 binding_security_result,
5754 SecurityCheckResult::RequireConfirmation { .. }
5755 ) || security_engine
5756 .classification_approval_message(
5757 &canonical_id,
5758 &resolved.tool.classify_call(&executed_arguments),
5759 )
5760 .is_some();
5761 let approval_binding = approval_record.as_ref().and_then(|record| {
5762 matches!(
5763 record.status,
5764 ToolApprovalStatus::Approved | ToolApprovalStatus::Modified
5765 )
5766 .then(|| ToolApprovalBinding {
5767 canonical_id: canonical_id.clone(),
5768 arguments: executed_arguments.clone(),
5769 confirmation_required: approval_confirmation_required,
5770 policy_version: security_engine.policy_version(),
5771 runtime_control_version: approval_control_snapshot.version,
5772 state_generation: initial_scope_snapshot.state_generation,
5773 reviewed_tool: Arc::clone(&resolved.tool),
5774 })
5775 });
5776
5777 let control_snapshot = self.runtime_safety_snapshot();
5782 let resolved = self.tools.resolve(&request.requested_name);
5783 let registry_version = self.tools.version();
5784 let mut versions = ToolDecisionVersions {
5785 policy: control_snapshot.tool_security.policy_version(),
5786 registry: registry_version,
5787 runtime_control: control_snapshot.version,
5788 state: None,
5789 };
5790 metadata.insert(
5791 "runtime_scope_snapshot".to_string(),
5792 serde_json::to_value(&control_snapshot.tool_scope_override).unwrap_or(Value::Null),
5793 );
5794 let resolved = match resolved {
5795 Some(resolved) => resolved,
5796 None => {
5797 let reason = format!(
5798 "Tool '{}' became unavailable after approval",
5799 request.requested_name
5800 );
5801 let record = self.record_from_parts_at(
5802 &request,
5803 request.requested_name.clone(),
5804 executed_arguments,
5805 started_at,
5806 start,
5807 false,
5808 false,
5809 reason.clone(),
5810 metadata,
5811 ToolPolicyDecisionRecord::unavailable(reason),
5812 approval_record,
5813 false,
5814 false,
5815 versions,
5816 );
5817 self.finish_tool_record(&record).await;
5818 return Ok(record);
5819 }
5820 };
5821
5822 let canonical_id = resolved.identity.canonical_id.clone();
5823 if let Some(reason) =
5824 fallback_state.final_rejection_reason(&admitted_canonical_id, &canonical_id)
5825 {
5826 metadata.insert(
5830 "fallback_chain".to_string(),
5831 serde_json::to_value(&fallback_state.visited_canonical_ids).unwrap_or(Value::Null),
5832 );
5833 metadata.insert(
5834 "final_resolved_canonical_id".to_string(),
5835 Value::String(canonical_id),
5836 );
5837 let record = self.record_from_parts_at(
5838 &request,
5839 admitted_canonical_id,
5840 executed_arguments,
5841 started_at,
5842 start,
5843 false,
5844 false,
5845 format!("Denied: {reason}"),
5846 metadata,
5847 ToolPolicyDecisionRecord::deny(reason),
5848 approval_record,
5849 false,
5850 false,
5851 versions,
5852 );
5853 self.finish_tool_record(&record).await;
5854 return Ok(record);
5855 }
5856 let bindings = resolved.tool.policy_bindings();
5857 let final_arguments = control_snapshot
5858 .tool_security
5859 .prepare_tool_arguments_with_bindings(&canonical_id, &executed_arguments, &bindings);
5860 if let Some(record) = approval_record.as_mut()
5861 && matches!(record.status, ToolApprovalStatus::Modified)
5862 {
5863 record.modified_arguments = Some(final_arguments.clone());
5864 }
5865 let classification = resolved.tool.classify_call(&final_arguments);
5866 let safety = resolved.tool.safety_metadata();
5867 let security_engine = control_snapshot.tool_security;
5868 let tool_config = self.recovery_manager.get_tool_config(&canonical_id).clone();
5869 let recovery_timeout_ms = self.recovery_manager.get_tool_timeout(&canonical_id);
5870 metadata.insert(
5871 "classification".to_string(),
5872 serde_json::to_value(&classification).unwrap_or(Value::Null),
5873 );
5874 let (limits, timeout) = match Self::effective_tool_limits(
5878 &security_engine,
5879 &canonical_id,
5880 &safety,
5881 &classification,
5882 recovery_timeout_ms,
5883 ) {
5884 Ok(effective) => effective,
5885 Err(error) => {
5886 let reason = error.to_string();
5887 metadata.insert(
5888 "configuration_error".to_string(),
5889 Value::String(reason.clone()),
5890 );
5891 let record = self.record_from_parts_at(
5892 &request,
5893 canonical_id,
5894 final_arguments,
5895 started_at,
5896 start,
5897 false,
5898 false,
5899 format!("Denied: {reason}"),
5900 metadata,
5901 ToolPolicyDecisionRecord::deny(reason),
5902 approval_record,
5903 false,
5904 false,
5905 versions,
5906 );
5907 self.finish_tool_record(&record).await;
5908 return Ok(record);
5909 }
5910 };
5911 let policy_snapshot = security_engine.policy_snapshot(&canonical_id);
5912 let resource_lock_keys =
5913 tool_resource_lock_keys(&canonical_id, &final_arguments, &bindings, &classification);
5914 metadata.insert(
5915 "effective_limits".to_string(),
5916 serde_json::to_value(&limits).unwrap_or(Value::Null),
5917 );
5918 metadata.insert(
5919 "resource_lock_keys".to_string(),
5920 serde_json::to_value(&resource_lock_keys).unwrap_or(Value::Null),
5921 );
5922 if policy_snapshot.is_null() {
5923 metadata.remove("policy_snapshot");
5924 } else {
5925 metadata.insert("policy_snapshot".to_string(), policy_snapshot.clone());
5926 }
5927
5928 let final_denial = |canonical_id: String,
5929 output: String,
5930 policy: ToolPolicyDecisionRecord,
5931 metadata: HashMap<String, Value>,
5932 decision_versions: ToolDecisionVersions| {
5933 self.record_from_parts_at(
5934 &request,
5935 canonical_id,
5936 final_arguments.clone(),
5937 started_at,
5938 start,
5939 false,
5940 false,
5941 output,
5942 metadata,
5943 policy,
5944 approval_record.clone(),
5945 false,
5946 false,
5947 decision_versions,
5948 )
5949 };
5950
5951 if control_snapshot.emergency_deny {
5952 let reason = "Tool execution is disabled by runtime control".to_string();
5953 let record = final_denial(
5954 canonical_id,
5955 reason.clone(),
5956 ToolPolicyDecisionRecord::deny(reason),
5957 metadata,
5958 versions,
5959 );
5960 self.finish_tool_record(&record).await;
5961 return Ok(record);
5962 }
5963
5964 let available_snapshot = self
5969 .get_available_tool_ids_snapshot_for_scope(
5970 control_snapshot.tool_scope_override.as_deref(),
5971 )
5972 .await?;
5973 versions.state = available_snapshot.state_generation;
5974 metadata.insert(
5975 "available_tool_ids_snapshot".to_string(),
5976 serde_json::to_value(&available_snapshot.tool_ids).unwrap_or(Value::Null),
5977 );
5978 metadata.insert(
5979 "state_generation_snapshot".to_string(),
5980 serde_json::to_value(available_snapshot.state_generation).unwrap_or(Value::Null),
5981 );
5982 if !available_snapshot
5983 .tool_ids
5984 .iter()
5985 .any(|tool_id| tool_id == &canonical_id)
5986 {
5987 let reason = format!(
5988 "Tool '{}' is not available in the final runtime scope",
5989 canonical_id
5990 );
5991 let record = final_denial(
5992 canonical_id,
5993 reason.clone(),
5994 ToolPolicyDecisionRecord::deny(reason),
5995 metadata,
5996 versions,
5997 );
5998 self.finish_tool_record(&record).await;
5999 return Ok(record);
6000 }
6001
6002 let final_security_result = security_engine
6007 .validate_tool_execution_with_bindings(&canonical_id, &final_arguments, &bindings)
6008 .await?;
6009 match &final_security_result {
6010 SecurityCheckResult::Block { reason } => {
6011 let record = final_denial(
6012 canonical_id,
6013 format!("Denied: {}", reason),
6014 ToolPolicyDecisionRecord::deny(reason.clone()),
6015 metadata,
6016 versions,
6017 );
6018 self.finish_tool_record(&record).await;
6019 return Ok(record);
6020 }
6021 SecurityCheckResult::Unavailable { reason } => {
6022 let record = final_denial(
6023 canonical_id,
6024 format!("Unavailable: {}", reason),
6025 ToolPolicyDecisionRecord::unavailable(reason.clone()),
6026 metadata,
6027 versions,
6028 );
6029 self.finish_tool_record(&record).await;
6030 return Ok(record);
6031 }
6032 SecurityCheckResult::Warn { message } => {
6033 warn!(tool = %canonical_id, message = %message, "Tool security warning after approval");
6034 }
6035 SecurityCheckResult::Allow | SecurityCheckResult::RequireConfirmation { .. } => {}
6036 }
6037 let final_confirmation_required = matches!(
6038 final_security_result,
6039 SecurityCheckResult::RequireConfirmation { .. }
6040 ) || security_engine
6041 .classification_approval_message(&canonical_id, &classification)
6042 .is_some();
6043 let stale_approval = approval_binding.as_ref().is_some_and(|binding| {
6044 binding.is_stale(
6045 &canonical_id,
6046 &final_arguments,
6047 final_confirmation_required,
6048 versions,
6049 &resolved.tool,
6050 )
6051 });
6052 if stale_approval {
6053 let reason = "Approval became stale before final admission".to_string();
6054 let record = final_denial(
6055 canonical_id,
6056 reason.clone(),
6057 ToolPolicyDecisionRecord::deny(reason),
6058 metadata,
6059 versions,
6060 );
6061 self.finish_tool_record(&record).await;
6062 return Ok(record);
6063 }
6064 if final_confirmation_required && approval_binding.is_none() {
6065 let reason = "Final policy requires fresh approval".to_string();
6066 let record = final_denial(
6067 canonical_id,
6068 reason.clone(),
6069 ToolPolicyDecisionRecord::approval(reason),
6070 metadata,
6071 versions,
6072 );
6073 self.finish_tool_record(&record).await;
6074 return Ok(record);
6075 }
6076
6077 if let Some((_, reason)) = self.host_tool_unavailability(&canonical_id) {
6078 let record = final_denial(
6079 canonical_id,
6080 reason.to_string(),
6081 ToolPolicyDecisionRecord::unavailable(reason),
6082 metadata,
6083 versions,
6084 );
6085 self.finish_tool_record(&record).await;
6086 return Ok(record);
6087 }
6088
6089 let Some(resource_guards) = self.acquire_tool_resource_locks(&resource_lock_keys).await
6094 else {
6095 let reason = "Tool execution cancelled while waiting for resource locks".to_string();
6099 let mut record = final_denial(
6100 canonical_id,
6101 reason.clone(),
6102 ToolPolicyDecisionRecord::deny(reason),
6103 metadata,
6104 versions,
6105 );
6106 record.cancelled = true;
6107 record.cancellation_reason = Some("runtime control cancellation".to_string());
6108 self.finish_tool_record(&record).await;
6109 return Ok(record);
6110 };
6111
6112 let admission = self.admit_tool_execution(
6117 versions.runtime_control,
6118 versions.policy,
6119 versions.state,
6120 &canonical_id,
6121 );
6122 if !matches!(admission, SecurityCheckResult::Allow) {
6123 let latest_control = self.runtime_safety_snapshot();
6124 let reason = admission
6125 .reason()
6126 .unwrap_or("tool admission was denied")
6127 .to_string();
6128 let policy = if admission.is_unavailable() {
6129 ToolPolicyDecisionRecord::unavailable(reason.clone())
6130 } else {
6131 ToolPolicyDecisionRecord::deny(reason.clone())
6132 };
6133 let record = self.record_from_parts_at(
6134 &request,
6135 canonical_id,
6136 final_arguments,
6137 started_at,
6138 start,
6139 false,
6140 false,
6141 reason,
6142 metadata,
6143 policy,
6144 approval_record,
6145 false,
6146 false,
6147 ToolDecisionVersions {
6148 policy: latest_control.tool_security.policy_version(),
6149 registry: versions.registry,
6150 runtime_control: latest_control.version,
6151 state: self
6152 .state_machine
6153 .as_ref()
6154 .map(|state_machine| state_machine.generation()),
6155 },
6156 );
6157 self.finish_tool_record_after_resource_guards(resource_guards, &record)
6158 .await;
6159 return Ok(record);
6160 }
6161 let executed_arguments = final_arguments;
6162
6163 let turn_actor = current_turn_actor_context();
6164 let actor = ToolActorContext {
6165 actor_id: turn_actor
6166 .as_ref()
6167 .and_then(|context| context.effective_actor_id().map(str::to_string))
6168 .or_else(|| self.actor_id()),
6169 origin_actor_id: turn_actor
6170 .as_ref()
6171 .and_then(|context| context.origin_actor_id.clone()),
6172 sender_agent_id: turn_actor
6173 .as_ref()
6174 .and_then(|context| context.sender_agent_id.clone()),
6175 };
6176 let tool_context = ToolExecutionContext {
6177 requested_name: request.requested_name.clone(),
6178 canonical_id: canonical_id.clone(),
6179 display_name: resolved.identity.display_name.clone(),
6180 provider_id: resolved.identity.provider_id.clone(),
6181 registry_version: versions.registry,
6182 policy_version: versions.policy,
6183 runtime_control_version: versions.runtime_control,
6184 call_id: request.call_id.clone(),
6185 source: request.source.clone(),
6186 actor,
6187 cancellation: ToolCancellationToken::new(
6188 Arc::clone(&self.runtime_control.emergency_deny),
6189 Some("runtime control cancellation".to_string()),
6190 ),
6191 started_at,
6192 deadline: None,
6193 permission: ToolPolicyDecisionRecord::allow(),
6194 approval: approval_record.clone(),
6195 classification: classification.clone(),
6196 safety,
6197 limits: limits.clone(),
6198 policy_snapshot,
6199 custom_config: security_engine.custom_config(&canonical_id),
6200 };
6201 let (mut result, timed_out, cancelled, invoked) = self
6202 .run_tool_with_retries(
6203 &canonical_id,
6204 resolved.tool.clone(),
6205 executed_arguments.clone(),
6206 tool_context,
6207 timeout,
6208 tool_config.max_retries,
6209 )
6210 .await?;
6211
6212 let fallback_tool = if !result.success && !cancelled {
6216 match &tool_config.on_failure {
6217 ToolFailureAction::Skip => {
6218 result = ToolResult::ok(format!(
6219 "{{\"skipped\": true, \"reason\": \"Tool '{}' was skipped after failure\"}}",
6220 canonical_id
6221 ));
6222 None
6223 }
6224 ToolFailureAction::Fallback { fallback_tool } => Some(fallback_tool.clone()),
6225 ToolFailureAction::ReportError => None,
6226 }
6227 } else {
6228 None
6229 };
6230
6231 let output_cap = limits.max_output_chars;
6232 let (output, output_truncated) =
6233 Self::truncate_tool_output(result.output.clone(), output_cap);
6234 if let Some(result_metadata) = result.metadata {
6235 metadata.extend(result_metadata);
6236 }
6237 let mut record = self.record_from_parts_at(
6238 &request,
6239 canonical_id,
6240 executed_arguments,
6241 started_at,
6242 start,
6243 invoked,
6244 result.success,
6245 output,
6246 metadata,
6247 ToolPolicyDecisionRecord::allow(),
6248 approval_record,
6249 timed_out,
6250 output_truncated,
6251 versions,
6252 );
6253 record.cancelled = cancelled;
6254 if cancelled {
6255 record.cancellation_reason = Some("runtime control cancellation".to_string());
6256 }
6257 if let Some(fallback_tool) = fallback_tool {
6258 let fallback_arguments = record.executed_arguments.clone();
6259 let original_tool = record.canonical_id.clone();
6260 self.finish_tool_record_after_resource_guards(resource_guards, &record)
6264 .await;
6265 let fallback_request = ToolExecutionRequest::new(
6266 request.call_id.clone(),
6267 fallback_tool,
6268 fallback_arguments,
6269 ToolCallSource::Fallback { original_tool },
6270 );
6271 return Box::pin(self.execute_tool_record_inner(fallback_request, fallback_state))
6272 .await;
6273 }
6274 self.finish_tool_record_after_resource_guards(resource_guards, &record)
6275 .await;
6276 Ok(record)
6277 }
6278
6279 #[instrument(skip(self, tool_call), fields(tool = %tool_call.name))]
6280 async fn execute_tool_smart(&self, tool_call: &ToolCall) -> Result<String> {
6281 let record = self
6282 .execute_tool_record(ToolExecutionRequest::new(
6283 tool_call.id.clone(),
6284 tool_call.name.clone(),
6285 tool_call.arguments.clone(),
6286 ToolCallSource::Model,
6287 ))
6288 .await?;
6289 if record.success {
6290 Ok(record.model_output_string())
6291 } else if matches!(record.policy.outcome, PermissionOutcome::RequiresApproval) {
6292 Err(AgentError::HITLRejected(record.model_output_string()))
6293 } else {
6294 Err(AgentError::Tool(record.model_output_string()))
6295 }
6296 }
6297
6298 async fn select_skill_candidate(&self, input: &str) -> Result<Option<SkillCandidate>> {
6304 let Some(ref router) = self.skill_router else {
6305 return Ok(None);
6306 };
6307 let available_skills = self.get_available_skills();
6308 if available_skills.is_empty() {
6309 return Ok(None);
6310 }
6311 let skill_ids: Vec<&str> = available_skills.iter().map(|s| s.id.as_str()).collect();
6312 let Some(skill_id) = self
6313 .observe_purpose(
6314 ObservationPurpose::SkillRouting,
6315 router.select_skill_filtered(input, &skill_ids),
6316 )
6317 .await?
6318 else {
6319 return Ok(None);
6320 };
6321 let skill = router
6322 .get_skill(&skill_id)
6323 .cloned()
6324 .ok_or_else(|| AgentError::Skill(format!("Skill not found: {}", skill_id)))?;
6325 info!(skill_id = %skill_id, "Skill selected");
6326 Ok(Some(SkillCandidate::new(skill_id, skill)))
6327 }
6328
6329 async fn commit_skill_candidate_route_result(
6334 &self,
6335 candidate: SkillCandidate,
6336 input: &str,
6337 ) -> Result<SkillRouteResult> {
6338 let skill_id = candidate.skill_id;
6339 let skill = candidate.skill;
6340 let expected_state_generation = self
6341 .state_machine
6342 .as_ref()
6343 .map(|state_machine| state_machine.generation());
6344 let expected_disambiguation_epoch = self.disambiguation_epoch.load(Ordering::SeqCst);
6345 if let Some(ref skill_disambig) = skill.disambiguation
6346 && skill_disambig.enabled.unwrap_or(false)
6347 && let Some(ref disambiguator) = self.disambiguation_manager
6348 {
6349 let context = self.build_disambiguation_context().await?;
6350 let state_override = self
6351 .state_machine
6352 .as_ref()
6353 .and_then(|sm| sm.current_definition())
6354 .and_then(|def| def.disambiguation.clone());
6355
6356 let disambiguation_result = self
6357 .observe_purpose(
6358 ObservationPurpose::DisambiguationDetection,
6359 disambiguator.process_input_with_override(
6360 input,
6361 &context,
6362 state_override.as_ref(),
6363 Some(skill_disambig),
6364 ),
6365 )
6366 .await?;
6367 let current_state_generation = self
6368 .state_machine
6369 .as_ref()
6370 .map(|state_machine| state_machine.generation());
6371 if current_state_generation != expected_state_generation
6372 || self.disambiguation_epoch.load(Ordering::SeqCst) != expected_disambiguation_epoch
6373 {
6374 disambiguator.clear_pending().await;
6375 *self.pending_skill_id.write() = None;
6376 return Err(AgentError::Other(
6377 "State or reset ownership changed during skill disambiguation".to_string(),
6378 ));
6379 }
6380 match disambiguation_result {
6381 DisambiguationResult::Clear => {
6382 debug!(skill_id = %skill_id, "Skill disambiguation: clear");
6383 }
6384 DisambiguationResult::NeedsClarification {
6385 question,
6386 detection,
6387 } => {
6388 let admission = self
6389 .admit_disambiguation_redispatch(
6390 expected_disambiguation_epoch,
6391 expected_state_generation,
6392 )
6393 .await?;
6394 let awaiting_confirmation = disambiguator.has_pending_confirmation().await;
6395 info!(
6396 skill_id = %skill_id,
6397 ambiguity_type = ?detection.ambiguity_type,
6398 confidence = detection.confidence,
6399 "Skill requires clarification before execution"
6400 );
6401 *self.pending_skill_id.write() = Some(skill_id.clone());
6402 let response = AgentResponse::new(&question.question).with_metadata(
6403 "disambiguation",
6404 serde_json::json!({
6405 "status": if awaiting_confirmation { "awaiting_confirmation" } else { "awaiting_clarification" },
6406 "skill_id": skill_id,
6407 "options": question.options,
6408 "clarifying": question.clarifying,
6409 "detection": {
6410 "type": detection.ambiguity_type,
6411 "confidence": detection.confidence,
6412 "what_is_unclear": detection.what_is_unclear,
6413 }
6414 }),
6415 );
6416 drop(admission);
6417 return Ok(SkillRouteResult::NeedsClarification {
6418 response,
6419 ownership: Some(DisambiguationOwnership {
6420 epoch: expected_disambiguation_epoch,
6421 state_generation: expected_state_generation,
6422 }),
6423 });
6424 }
6425 DisambiguationResult::Clarified { enriched_input, .. } => {
6426 info!(skill_id = %skill_id, enriched = %enriched_input, "Skill disambiguation clarified");
6427 let admission = self
6428 .admit_disambiguation_redispatch(
6429 expected_disambiguation_epoch,
6430 expected_state_generation,
6431 )
6432 .await?;
6433 drop(admission);
6434 let content = self.execute_skill(&skill, &enriched_input).await?;
6435 return Ok(SkillRouteResult::Response { skill_id, content });
6436 }
6437 DisambiguationResult::ProceedWithBestGuess { enriched_input } => {
6438 info!(skill_id = %skill_id, "Skill disambiguation best guess");
6439 let admission = self
6440 .admit_disambiguation_redispatch(
6441 expected_disambiguation_epoch,
6442 expected_state_generation,
6443 )
6444 .await?;
6445 drop(admission);
6446 let content = self.execute_skill(&skill, &enriched_input).await?;
6447 return Ok(SkillRouteResult::Response { skill_id, content });
6448 }
6449 DisambiguationResult::GiveUp { reason } => {
6450 warn!(skill_id = %skill_id, reason = %reason, "Skill disambiguation gave up");
6451 let apology = self
6452 .generate_localized_apology(
6453 "Generate a brief, polite apology saying you couldn't understand the request. Be concise.",
6454 &reason,
6455 )
6456 .await
6457 .unwrap_or_else(|_| {
6458 format!("I'm sorry, I couldn't understand your request: {}", reason)
6459 });
6460 return Ok(SkillRouteResult::NeedsClarification {
6461 response: AgentResponse::new(&apology),
6462 ownership: None,
6463 });
6464 }
6465 DisambiguationResult::Escalate { reason } => {
6466 info!(skill_id = %skill_id, reason = %reason, "Skill disambiguation escalating");
6467 let apology = self
6468 .generate_localized_apology(
6469 "Explain briefly that you're transferring the user to a human agent for help.",
6470 &reason,
6471 )
6472 .await
6473 .unwrap_or_else(|_| {
6474 format!("I need human assistance to help with your request: {}", reason)
6475 });
6476 return Ok(SkillRouteResult::NeedsClarification {
6477 response: AgentResponse::new(&apology),
6478 ownership: None,
6479 });
6480 }
6481 DisambiguationResult::Abandoned { .. } => {
6482 debug!(skill_id = %skill_id, "Skill disambiguation abandoned");
6483 return Ok(SkillRouteResult::NoMatch);
6484 }
6485 }
6486 }
6487 let admission = self
6488 .admit_disambiguation_redispatch(
6489 expected_disambiguation_epoch,
6490 expected_state_generation,
6491 )
6492 .await?;
6493 drop(admission);
6494 let content = self.execute_skill(&skill, input).await?;
6495 Ok(SkillRouteResult::Response { skill_id, content })
6496 }
6497
6498 async fn try_skill_route(&self, input: &str) -> Result<SkillRouteResult> {
6500 if let Some(candidate) = self.select_skill_candidate(input).await? {
6501 self.commit_skill_candidate_route_result(candidate, input)
6502 .await
6503 } else {
6504 Ok(SkillRouteResult::NoMatch)
6505 }
6506 }
6507
6508 async fn execute_skill(&self, skill: &SkillDefinition, input: &str) -> Result<String> {
6510 if let Some(ref executor) = self.skill_executor {
6511 let skill_reasoning = self.get_skill_reasoning_config(skill);
6512 let skill_reflection = self.get_skill_reflection_config(skill);
6513
6514 debug!(
6515 skill_id = %skill.id,
6516 reasoning_mode = ?skill_reasoning.mode,
6517 reflection_enabled = ?skill_reflection.enabled,
6518 "Skill reasoning/reflection config"
6519 );
6520
6521 let response = self
6522 .observe_purpose(
6523 ObservationPurpose::SkillPrompt,
6524 executor.execute_with_invoker(skill, input, serde_json::json!({}), self),
6525 )
6526 .await?;
6527
6528 if skill_reflection.requires_evaluation() && skill_reflection.is_enabled() {
6529 let should_reflect = self
6530 .should_reflect_with_config(input, &response, &skill_reflection)
6531 .await?;
6532 if should_reflect {
6533 let evaluated = self
6534 .evaluate_and_retry_with_config(input, response, &skill_reflection)
6535 .await?;
6536 return Ok(evaluated);
6537 }
6538 }
6539
6540 return Ok(response);
6541 }
6542 Err(AgentError::Skill(
6543 "No skill executor configured".to_string(),
6544 ))
6545 }
6546
6547 async fn execute_skill_by_id(&self, skill_id: &str, input: &str) -> Result<String> {
6550 let skill = self
6551 .skill_router
6552 .as_ref()
6553 .and_then(|r| r.get_skill(skill_id).cloned())
6554 .ok_or_else(|| AgentError::Skill(format!("Skill not found: {}", skill_id)))?;
6555 self.execute_skill(&skill, input).await
6556 }
6557
6558 async fn should_reflect_with_config(
6559 &self,
6560 input: &str,
6561 response: &str,
6562 config: &ReflectionConfig,
6563 ) -> Result<bool> {
6564 if !config.requires_evaluation() {
6565 return Ok(false);
6566 }
6567
6568 if config.is_enabled() {
6569 return Ok(true);
6570 }
6571
6572 let evaluator_llm = config
6573 .evaluator_llm
6574 .as_ref()
6575 .and_then(|alias| self.llm_registry.get(alias).ok())
6576 .or_else(|| self.llm_registry.router().ok())
6577 .or_else(|| self.llm_registry.default().ok());
6578
6579 let Some(llm) = evaluator_llm else {
6580 return Ok(false);
6581 };
6582
6583 let response_preview: String = response.chars().take(500).collect();
6584 let prompt = format!(
6585 r#"Should this response be evaluated for quality? Consider if it's a complex or important response.
6586
6587User query: "{}"
6588Response: "{}"
6589
6590Answer YES or NO only."#,
6591 input, response_preview
6592 );
6593
6594 let messages = vec![ChatMessage::user(&prompt)];
6595 let result = self
6596 .observe_purpose(
6597 ObservationPurpose::ReflectionDecision,
6598 llm.complete(&messages, None),
6599 )
6600 .await;
6601
6602 match result {
6603 Ok(resp) => Ok(resp.content.trim().to_uppercase().contains("YES")),
6604 Err(_) => Ok(false),
6605 }
6606 }
6607
6608 async fn evaluate_and_retry_with_config(
6609 &self,
6610 input: &str,
6611 mut response: String,
6612 config: &ReflectionConfig,
6613 ) -> Result<String> {
6614 let llm = self.get_state_llm()?;
6615 let mut attempts = 0u32;
6616 let max_retries = config.max_retries;
6617
6618 loop {
6619 let evaluation = self
6620 .evaluate_response_with_config(input, &response, config)
6621 .await?;
6622
6623 if evaluation.passed || attempts >= max_retries {
6624 info!(
6625 passed = evaluation.passed,
6626 confidence = evaluation.confidence,
6627 attempts = attempts + 1,
6628 "Skill reflection evaluation complete"
6629 );
6630 return Ok(response);
6631 }
6632
6633 debug!(
6634 attempt = attempts + 1,
6635 failed_criteria = evaluation.failed_criteria().count(),
6636 "Skill response did not meet criteria, retrying"
6637 );
6638
6639 let feedback: Vec<String> = evaluation
6640 .failed_criteria()
6641 .map(|c| format!("- {}", c.criterion))
6642 .collect();
6643
6644 let retry_prompt = format!(
6645 "Your previous response did not meet these criteria:\n{}\n\nPlease provide an improved response to: {}",
6646 feedback.join("\n"),
6647 input
6648 );
6649
6650 let messages = vec![ChatMessage::user(&retry_prompt)];
6651 let retry_response = self
6652 .observe_purpose(
6653 ObservationPurpose::ReflectionEvaluation,
6654 llm.complete(&messages, None),
6655 )
6656 .await
6657 .map_err(|e| AgentError::LLM(e.to_string()))?;
6658
6659 response = retry_response.content.trim().to_string();
6660 attempts += 1;
6661 }
6662 }
6663
6664 async fn evaluate_response_with_config(
6665 &self,
6666 input: &str,
6667 response: &str,
6668 config: &ReflectionConfig,
6669 ) -> Result<EvaluationResult> {
6670 let evaluator_llm = config
6671 .evaluator_llm
6672 .as_ref()
6673 .and_then(|alias| self.llm_registry.get(alias).ok())
6674 .or_else(|| self.llm_registry.router().ok())
6675 .or_else(|| self.llm_registry.default().ok())
6676 .ok_or_else(|| AgentError::Config("No LLM available for evaluation".into()))?;
6677
6678 let criteria = &config.criteria;
6679 let criteria_list = criteria
6680 .iter()
6681 .enumerate()
6682 .map(|(i, c)| format!("{}. {}", i + 1, c))
6683 .collect::<Vec<_>>()
6684 .join("\n");
6685
6686 let prompt = format!(
6687 r#"Evaluate this response against the criteria.
6688
6689User query: "{}"
6690
6691Response to evaluate: "{}"
6692
6693Criteria:
6694{}
6695
6696For each criterion, respond with:
6697- criterion number
6698- PASS or FAIL
6699- brief reason
6700
6701Then provide overall confidence (0.0 to 1.0) and whether it passes overall.
6702
6703Format:
67041. PASS/FAIL - reason
67052. PASS/FAIL - reason
6706...
6707CONFIDENCE: 0.X
6708OVERALL: PASS/FAIL"#,
6709 input, response, criteria_list
6710 );
6711
6712 let messages = vec![ChatMessage::user(&prompt)];
6713 let eval_response = self
6714 .observe_purpose(
6715 ObservationPurpose::ReflectionEvaluation,
6716 evaluator_llm.complete(&messages, None),
6717 )
6718 .await
6719 .map_err(|e| AgentError::LLM(format!("Evaluation failed: {}", e)))?;
6720
6721 let content = eval_response.content.to_uppercase();
6722 let llm_pass = content.contains("OVERALL: PASS");
6723
6724 let confidence = content
6725 .lines()
6726 .find(|l| l.contains("CONFIDENCE:"))
6727 .and_then(|l| {
6728 l.split(':')
6729 .nth(1)
6730 .and_then(|v| v.trim().parse::<f32>().ok())
6731 })
6732 .unwrap_or(if llm_pass { 0.8 } else { 0.4 });
6733
6734 let overall_pass = llm_pass && confidence >= config.pass_threshold;
6737
6738 let mut criteria_results = Vec::new();
6739 for (i, criterion) in criteria.iter().enumerate() {
6740 let line_marker = format!("{}.", i + 1);
6741 let passed = eval_response
6742 .content
6743 .lines()
6744 .find(|l| l.contains(&line_marker))
6745 .map(|l| l.to_uppercase().contains("PASS"))
6746 .unwrap_or(overall_pass);
6747
6748 if passed {
6749 criteria_results.push(CriterionResult::pass(criterion));
6750 } else {
6751 criteria_results.push(CriterionResult::fail(criterion, "Did not meet criterion"));
6752 }
6753 }
6754
6755 Ok(EvaluationResult::new(overall_pass, confidence).with_criteria(criteria_results))
6756 }
6757
6758 async fn process_input(&self, input: &str) -> Result<ProcessData> {
6760 if let Some(processor) = self.get_state_process_processor() {
6761 let purpose = observation_purpose_for_process(processor.input_purpose_hint());
6762 return self
6763 .observe_purpose(purpose, processor.process_input(input))
6764 .await;
6765 }
6766 if let Some(ref processor) = self.process_processor {
6767 let purpose = observation_purpose_for_process(processor.input_purpose_hint());
6768 self.observe_purpose(purpose, processor.process_input(input))
6769 .await
6770 } else {
6771 Ok(ProcessData::new(input))
6772 }
6773 }
6774
6775 async fn process_output(
6777 &self,
6778 output: &str,
6779 input_context: &std::collections::HashMap<String, serde_json::Value>,
6780 ) -> Result<ProcessData> {
6781 if let Some(processor) = self.get_state_process_processor() {
6782 let purpose = observation_purpose_for_process(processor.output_purpose_hint());
6783 return self
6784 .observe_purpose(purpose, processor.process_output(output, input_context))
6785 .await;
6786 }
6787 if let Some(ref processor) = self.process_processor {
6788 let purpose = observation_purpose_for_process(processor.output_purpose_hint());
6789 self.observe_purpose(purpose, processor.process_output(output, input_context))
6790 .await
6791 } else {
6792 Ok(ProcessData::new(output))
6793 }
6794 }
6795
6796 fn get_state_process_processor(&self) -> Option<ProcessProcessor> {
6798 let sm = self.state_machine.as_ref()?;
6799 let def = sm.current_definition()?;
6800 let config = def.process.as_ref()?;
6801 let mut processor = ProcessProcessor::new(config.clone());
6802 if let Some(ref registry) = Some(self.llm_registry.clone()) {
6803 processor = processor.with_llm_registry(registry.clone());
6804 }
6805 processor = processor.with_stage_observer(Arc::new(ObservabilityProcessStageObserver));
6806 Some(processor)
6807 }
6808
6809 async fn check_turn_timeout(&self) -> Result<()> {
6811 let Some(ref sm) = self.state_machine else {
6812 return Ok(());
6813 };
6814 let Some(timeout_state) = sm.check_timeout() else {
6815 return Ok(());
6816 };
6817 let claim_admission = self.disambiguation_admission.write().await;
6818 if sm.check_timeout().as_deref() != Some(timeout_state.as_str()) {
6819 return Ok(());
6820 }
6821 let Some(reservation) = self.reserve_state_transition() else {
6822 return Ok(());
6823 };
6824 let from_state = sm.current();
6825 let expected_state_generation = sm.generation();
6826 let expected_disambiguation_epoch = self.disambiguation_epoch.load(Ordering::SeqCst);
6827 let history_before = sm.history();
6828 drop(claim_admission);
6829
6830 self.execute_state_exit_actions(&from_state).await;
6831
6832 let admission = self.disambiguation_admission.write().await;
6833 if sm.current() != from_state
6834 || sm.generation() != expected_state_generation
6835 || self.disambiguation_epoch.load(Ordering::SeqCst) != expected_disambiguation_epoch
6836 || sm.check_timeout().as_deref() != Some(timeout_state.as_str())
6837 {
6838 return Ok(());
6839 }
6840 sm.transition_to(&timeout_state, "max_turns exceeded")?;
6841 self.invalidate_pending_confirmation("state_timeout").await;
6842 let entered = sm.current();
6843 let is_reentry = Self::state_was_previously_entered(&entered, &from_state, &history_before);
6844 drop(admission);
6845
6846 self.execute_state_enter_actions(&entered, is_reentry).await;
6847 drop(reservation);
6848 info!(to = %entered, "Timeout transition");
6849 Ok(())
6850 }
6851
6852 fn increment_turn(&self) {
6853 if let Some(ref sm) = self.state_machine {
6854 sm.increment_turn();
6855 }
6856 }
6857
6858 fn transitions_available_for_commit(&self) -> Option<(Vec<Transition>, String)> {
6859 let sm = self.state_machine.as_ref()?;
6860 let current = sm.current();
6861 let transitions: Vec<_> = sm
6862 .auto_transitions()
6863 .into_iter()
6864 .filter(|t| match t.cooldown_turns {
6865 Some(cd) if cd > 0 => {
6866 let resolved = sm.config().resolve_full_path(¤t, &t.to);
6867 !sm.is_on_cooldown(&resolved, cd)
6868 }
6869 _ => true,
6870 })
6871 .collect();
6872 Some((transitions, current))
6873 }
6874
6875 fn transition_reason(transition: &Transition) -> String {
6876 if transition.when.is_empty() {
6877 "guard condition met".to_string()
6878 } else {
6879 transition.when.clone()
6880 }
6881 }
6882
6883 fn build_transition_context(
6885 &self,
6886 user_message: &str,
6887 response: &str,
6888 current_state: &str,
6889 staged: Option<&HashMap<String, Value>>,
6890 ) -> TransitionContext {
6891 let context_map = staged
6892 .map(|writes| self.build_context_with_staged(writes))
6893 .unwrap_or_else(|| self.build_context_with_overlays());
6894 TransitionContext::new(user_message, response, current_state).with_context(context_map)
6895 }
6896
6897 async fn select_transition_candidate(
6899 &self,
6900 user_message: &str,
6901 response: &str,
6902 ) -> Result<Option<TransitionCandidate>> {
6903 let Some((transitions, current_state)) = self.transitions_available_for_commit() else {
6904 return Ok(None);
6905 };
6906 let transitions: Vec<Transition> = transitions
6907 .into_iter()
6908 .filter(|transition| matches!(transition.timing, TransitionTiming::PostResponse))
6909 .collect();
6910 if transitions.is_empty() {
6911 return Ok(None);
6912 }
6913 let Some(evaluator) = self.transition_evaluator.as_ref() else {
6914 return Ok(None);
6915 };
6916 let context = self.build_transition_context(user_message, response, ¤t_state, None);
6917 let selected = self
6918 .observe_purpose(
6919 ObservationPurpose::StateTransitionEvaluation,
6920 evaluator.select_transition(&transitions, &context),
6921 )
6922 .await?;
6923 Ok(selected.map(|index| {
6924 let transition = transitions[index].clone();
6925 TransitionCandidate::new(
6926 current_state,
6927 transition.clone(),
6928 Self::transition_reason(&transition),
6929 )
6930 }))
6931 }
6932
6933 fn select_deterministic_transition_candidate(
6935 &self,
6936 user_message: &str,
6937 current_state: &str,
6938 transitions: &[Transition],
6939 staged: &HashMap<String, Value>,
6940 ) -> Option<TransitionCandidate> {
6941 let context = self.build_transition_context(user_message, "", current_state, Some(staged));
6942
6943 for transition in transitions {
6944 if let Some(guard) = transition.guard.as_ref()
6945 && evaluate_guard(guard, &context)
6946 {
6947 return Some(TransitionCandidate::new(
6948 current_state,
6949 transition.clone(),
6950 Self::transition_reason(transition),
6951 ));
6952 }
6953 }
6954
6955 let resolved_intent = context
6956 .context
6957 .get("resolved_intent")
6958 .and_then(Value::as_str)
6959 .filter(|value| !value.is_empty());
6960 if let Some(resolved_intent) = resolved_intent {
6961 for transition in transitions {
6962 if transition.intent.as_deref() == Some(resolved_intent) {
6963 return Some(TransitionCandidate::new(
6964 current_state,
6965 transition.clone(),
6966 Self::transition_reason(transition),
6967 ));
6968 }
6969 }
6970 }
6971
6972 None
6973 }
6974
6975 async fn commit_transition_candidate(&self, candidate: &TransitionCandidate) -> Result<bool> {
6977 self.commit_transition_target(&candidate.from_state, candidate.target(), &candidate.reason)
6978 .await
6979 }
6980
6981 async fn approve_transition_target(&self, from_state: &str, target: &str) -> Result<bool> {
6983 let approved = self.check_state_hitl(Some(from_state), target).await?;
6984 if !approved {
6985 info!(to = %target, "State transition rejected by HITL");
6986 }
6987 Ok(approved)
6988 }
6989
6990 async fn apply_transition_target(
6992 &self,
6993 from_state: &str,
6994 target: &str,
6995 reason: &str,
6996 staged: Option<&HashMap<String, Value>>,
6997 ) -> Result<bool> {
6998 let Some(ref sm) = self.state_machine else {
6999 return Ok(false);
7000 };
7001 let claim_admission = self.disambiguation_admission.write().await;
7002 if sm.current() != from_state {
7003 return Ok(false);
7004 }
7005 let Some(reservation) = self.reserve_state_transition() else {
7006 return Ok(false);
7007 };
7008 let expected_state_generation = sm.generation();
7009 let expected_disambiguation_epoch = self.disambiguation_epoch.load(Ordering::SeqCst);
7010 let history_before = sm.history();
7011 drop(claim_admission);
7012
7013 self.execute_state_exit_actions(from_state).await;
7014
7015 let admission = self.disambiguation_admission.write().await;
7016 if sm.current() != from_state
7017 || sm.generation() != expected_state_generation
7018 || self.disambiguation_epoch.load(Ordering::SeqCst) != expected_disambiguation_epoch
7019 {
7020 return Ok(false);
7021 }
7022 sm.transition_to(target, reason)?;
7023 self.invalidate_pending_confirmation("state_transition")
7024 .await;
7025 sm.reset_no_transition();
7026 if let Some(staged) = staged {
7027 self.commit_staged_context_writes(staged);
7028 }
7029 let entered = sm.current();
7030 let is_reentry = Self::state_was_previously_entered(&entered, from_state, &history_before);
7031 drop(admission);
7032
7033 self.execute_state_enter_actions(&entered, is_reentry).await;
7034 drop(reservation);
7035 self.hooks
7036 .on_state_transition(Some(from_state), &entered, reason)
7037 .await;
7038 info!(from = %from_state, to = %entered, "State transition");
7039 Ok(true)
7040 }
7041
7042 async fn commit_transition_target(
7044 &self,
7045 from_state: &str,
7046 target: &str,
7047 reason: &str,
7048 ) -> Result<bool> {
7049 if !self.approve_transition_target(from_state, target).await? {
7050 return Ok(false);
7051 }
7052 self.apply_transition_target(from_state, target, reason, None)
7053 .await
7054 }
7055
7056 async fn apply_pre_response_transition_candidate(
7058 &self,
7059 candidate: &TransitionCandidate,
7060 staged: &HashMap<String, Value>,
7061 processed_input: &str,
7062 ) -> Result<bool> {
7063 self.commit_root_user_message(processed_input).await?;
7064 self.apply_transition_target(
7065 &candidate.from_state,
7066 candidate.target(),
7067 &candidate.reason,
7068 Some(staged),
7069 )
7070 .await
7071 }
7072
7073 async fn commit_pre_response_transition_candidate(
7075 &self,
7076 candidate: &TransitionCandidate,
7077 staged: &HashMap<String, Value>,
7078 processed_input: &str,
7079 ) -> Result<bool> {
7080 if !self
7081 .approve_transition_target(&candidate.from_state, candidate.target())
7082 .await?
7083 {
7084 return Ok(false);
7085 }
7086 self.apply_pre_response_transition_candidate(candidate, staged, processed_input)
7087 .await
7088 }
7089
7090 async fn handle_transition_miss(&self, current_state: &str) -> Result<bool> {
7092 let Some(ref sm) = self.state_machine else {
7093 return Ok(false);
7094 };
7095 sm.increment_no_transition();
7096 let Some(fallback) = sm.check_fallback() else {
7097 return Ok(false);
7098 };
7099 self.commit_transition_target(current_state, &fallback, "fallback after no transitions")
7100 .await
7101 }
7102
7103 async fn evaluate_transitions(&self, user_message: &str, response: &str) -> Result<bool> {
7105 let Some((transitions, current_state)) = self.transitions_available_for_commit() else {
7106 return Ok(false);
7107 };
7108 if transitions.is_empty() {
7109 return Ok(false);
7110 }
7111 if let Some(candidate) = self
7112 .select_transition_candidate(user_message, response)
7113 .await?
7114 {
7115 return self.commit_transition_candidate(&candidate).await;
7116 }
7117 self.handle_transition_miss(¤t_state).await
7118 }
7119
7120 async fn try_pre_response_transition(
7122 &self,
7123 processed_input: &str,
7124 ) -> Result<Option<AgentResponse>> {
7125 let optimization = &self.runtime_config.optimization;
7126 if !optimization.enabled || !optimization.pre_response_deterministic_transitions {
7127 return Ok(None);
7128 }
7129 let Some((transitions, current_state)) = self.transitions_available_for_commit() else {
7130 return Ok(None);
7131 };
7132 let eligible: Vec<Transition> = transitions
7133 .into_iter()
7134 .filter(|transition| !transition.requires_response)
7135 .filter(|transition| matches!(transition.timing, TransitionTiming::PreResponse))
7136 .collect();
7137 if eligible.is_empty() {
7138 return Ok(None);
7139 }
7140
7141 let empty_staged = HashMap::new();
7142 let mut extracted_staged: Option<HashMap<String, Value>> = None;
7143 let mut selected: Option<(TransitionCandidate, HashMap<String, Value>)> = None;
7144
7145 for transition in &eligible {
7146 let use_extractors = optimization.pre_response_extractors || transition.run_extractors;
7147 let staged_for_eval = if use_extractors {
7148 if extracted_staged.is_none() {
7149 extracted_staged =
7150 Some(self.run_context_extractors_staged(processed_input).await);
7151 }
7152 extracted_staged.as_ref().unwrap_or(&empty_staged)
7153 } else {
7154 &empty_staged
7155 };
7156
7157 if let Some(candidate) = self.select_deterministic_transition_candidate(
7158 processed_input,
7159 ¤t_state,
7160 std::slice::from_ref(transition),
7161 staged_for_eval,
7162 ) {
7163 let staged_for_commit = if use_extractors {
7164 staged_for_eval.clone()
7165 } else {
7166 HashMap::new()
7167 };
7168 selected = Some((candidate, staged_for_commit));
7169 break;
7170 }
7171 }
7172
7173 let Some((candidate, staged)) = selected else {
7174 return Ok(None);
7175 };
7176
7177 if !self
7178 .commit_pre_response_transition_candidate(&candidate, &staged, processed_input)
7179 .await?
7180 {
7181 return Ok(None);
7182 }
7183 self.redispatch_current_state(processed_input)
7184 .await
7185 .map(Some)
7186 }
7187
7188 async fn try_speculative_branches(
7193 &self,
7194 processed_input: &str,
7195 input_context: &HashMap<String, Value>,
7196 ) -> Result<Option<AgentResponse>> {
7197 let optimization = &self.runtime_config.optimization;
7198 if !optimization.enabled {
7199 return Ok(None);
7200 }
7201
7202 let effective_reasoning_mode = self.get_effective_reasoning_config().mode.clone();
7203 if !matches!(
7204 effective_reasoning_mode,
7205 ReasoningMode::None | ReasoningMode::Auto
7206 ) {
7207 return Ok(None);
7208 }
7209
7210 let mut transition_enabled =
7211 optimization.speculative_state_transitions && self.has_parallel_transition_candidates();
7212 let mut skill_enabled = optimization.speculative_skill_routing
7213 && self.skill_router.is_some()
7214 && self.pending_skill_id.read().is_none();
7215 let mut reasoning_enabled = optimization.speculative_reasoning_auto
7216 && matches!(effective_reasoning_mode, ReasoningMode::Auto);
7217
7218 if matches!(effective_reasoning_mode, ReasoningMode::Auto)
7219 && (!reasoning_enabled || optimization.max_speculative_llm_calls_per_turn < 2)
7220 {
7221 return Ok(None);
7222 }
7223
7224 if !transition_enabled && !skill_enabled && !reasoning_enabled {
7225 return Ok(None);
7226 }
7227
7228 let mut optional_slots = optimization.max_parallel_runtime_tasks.saturating_sub(1);
7229 let mut speculative_call_slots = optimization
7230 .max_speculative_llm_calls_per_turn
7231 .saturating_sub(1);
7232 if reasoning_enabled {
7233 if optional_slots == 0 || speculative_call_slots == 0 {
7234 return Ok(None);
7235 }
7236 optional_slots -= 1;
7237 speculative_call_slots -= 1;
7238 }
7239 if transition_enabled {
7240 if optional_slots == 0 {
7241 transition_enabled = false;
7242 } else {
7243 optional_slots -= 1;
7244 }
7245 }
7246 if skill_enabled && (optional_slots == 0 || speculative_call_slots == 0) {
7247 skill_enabled = false;
7248 }
7249
7250 if !transition_enabled && !skill_enabled && !reasoning_enabled {
7251 return Ok(None);
7252 }
7253
7254 let main_kind = if transition_enabled {
7255 RuntimeOptimizationKind::ParallelStateTransition
7256 } else if skill_enabled {
7257 RuntimeOptimizationKind::SpeculativeSkillRouting
7258 } else {
7259 RuntimeOptimizationKind::SpeculativeReasoningAuto
7260 };
7261 if !self.reserve_active_speculative_llm_call(main_kind) {
7262 return Ok(None);
7263 }
7264
7265 let mut branch_set = ScheduledBranchSet::new(optimization.max_parallel_runtime_tasks)?;
7266 let main_branch = RuntimeBranch::new(
7267 RuntimeTaskPurpose::MainResponse,
7268 main_kind,
7269 RuntimeTaskPriority::Normal,
7270 RuntimeCommitBehavior::FinalResponse,
7271 );
7272 let transition_branch = RuntimeBranch::new(
7273 RuntimeTaskPurpose::StateTransition,
7274 RuntimeOptimizationKind::ParallelStateTransition,
7275 RuntimeTaskPriority::Critical,
7276 RuntimeCommitBehavior::TransitionDecision,
7277 );
7278 let skill_branch = RuntimeBranch::new(
7279 RuntimeTaskPurpose::SkillRouting,
7280 RuntimeOptimizationKind::SpeculativeSkillRouting,
7281 RuntimeTaskPriority::High,
7282 RuntimeCommitBehavior::SkillSelection,
7283 );
7284 let reasoning_branch = RuntimeBranch::new(
7285 RuntimeTaskPurpose::ReasoningJudge,
7286 RuntimeOptimizationKind::SpeculativeReasoningAuto,
7287 RuntimeTaskPriority::Normal,
7288 RuntimeCommitBehavior::ReasoningDecision,
7289 );
7290 let main_id = main_branch.branch_id();
7291 let transition_id = transition_branch.branch_id();
7292 let skill_id = skill_branch.branch_id();
7293 let reasoning_id = reasoning_branch.branch_id();
7294
7295 let main_id_for_future = main_id.clone();
7296 if !branch_set.schedule(
7297 main_branch,
7298 Box::pin(async move {
7299 match crate::optimization::observability::with_branch_observation(
7300 &main_id_for_future,
7301 main_kind,
7302 RuntimeCommitBehavior::FinalResponse,
7303 self.generate_main_response_draft(processed_input, &ReasoningMode::None),
7304 )
7305 .await
7306 {
7307 Ok(draft) => RuntimeBranchResult::MainDraft(draft),
7308 Err(error) => RuntimeBranchResult::Failed(error),
7309 }
7310 }),
7311 ) {
7312 return Ok(None);
7313 }
7314
7315 if transition_enabled {
7316 let transition_id_for_future = transition_id.clone();
7317 if !branch_set.schedule(
7318 transition_branch,
7319 Box::pin(async move {
7320 match crate::optimization::observability::with_branch_observation(
7321 &transition_id_for_future,
7322 RuntimeOptimizationKind::ParallelStateTransition,
7323 RuntimeCommitBehavior::TransitionDecision,
7324 self.select_parallel_transition_candidate(processed_input),
7325 )
7326 .await
7327 {
7328 Ok(ParallelTransitionSelection::Candidate(candidate)) => {
7329 RuntimeBranchResult::Transition(Some(candidate))
7330 }
7331 Ok(ParallelTransitionSelection::NoMatch) => {
7332 RuntimeBranchResult::Transition(None)
7333 }
7334 Ok(ParallelTransitionSelection::ReservationExhausted) => {
7335 RuntimeBranchResult::Cancelled
7336 }
7337 Err(error) => RuntimeBranchResult::Failed(error),
7338 }
7339 }),
7340 ) {
7341 transition_enabled = false;
7342 }
7343 }
7344
7345 if skill_enabled {
7346 let skill_id_for_future = skill_id.clone();
7347 if !branch_set.schedule(
7348 skill_branch,
7349 Box::pin(async move {
7350 if !self.reserve_active_speculative_llm_call(
7351 RuntimeOptimizationKind::SpeculativeSkillRouting,
7352 ) {
7353 return RuntimeBranchResult::Cancelled;
7354 }
7355 match crate::optimization::observability::with_branch_observation(
7356 &skill_id_for_future,
7357 RuntimeOptimizationKind::SpeculativeSkillRouting,
7358 RuntimeCommitBehavior::SkillSelection,
7359 self.select_skill_candidate(processed_input),
7360 )
7361 .await
7362 {
7363 Ok(candidate) => RuntimeBranchResult::Skill(candidate),
7364 Err(error) => RuntimeBranchResult::Failed(error),
7365 }
7366 }),
7367 ) {
7368 skill_enabled = false;
7369 }
7370 }
7371
7372 if reasoning_enabled {
7373 let reasoning_id_for_future = reasoning_id.clone();
7374 if !branch_set.schedule(
7375 reasoning_branch,
7376 Box::pin(async move {
7377 if !self.reserve_active_speculative_llm_call(
7378 RuntimeOptimizationKind::SpeculativeReasoningAuto,
7379 ) {
7380 return RuntimeBranchResult::Cancelled;
7381 }
7382 match crate::optimization::observability::with_branch_observation(
7383 &reasoning_id_for_future,
7384 RuntimeOptimizationKind::SpeculativeReasoningAuto,
7385 RuntimeCommitBehavior::ReasoningDecision,
7386 self.determine_reasoning_mode_strict(processed_input),
7387 )
7388 .await
7389 {
7390 Ok(mode) => RuntimeBranchResult::Reasoning(mode),
7391 Err(error) => RuntimeBranchResult::Failed(error),
7392 }
7393 }),
7394 ) {
7395 reasoning_enabled = false;
7396 }
7397 }
7398
7399 if matches!(effective_reasoning_mode, ReasoningMode::Auto) && !reasoning_enabled {
7400 self.finalize_pending_branches(branch_set.cancel_pending());
7401 return Ok(None);
7402 }
7403
7404 if !transition_enabled && !skill_enabled && !reasoning_enabled {
7405 self.finalize_pending_branches(branch_set.cancel_pending());
7406 return Ok(None);
7407 }
7408
7409 let mut main_pending = true;
7410 let mut skill_pending = skill_enabled;
7411 let mut reasoning_pending = reasoning_enabled;
7412 let mut transition_finalized = !transition_enabled;
7413 let mut skill_finalized = !skill_enabled;
7414 let mut reasoning_finalized = !reasoning_enabled;
7415 let mut main_result: Option<Result<MainResponseDraft>> = None;
7416 let mut transition_candidate: Option<TransitionCandidate> = None;
7417 let mut skill_candidate: Option<SkillCandidate> = None;
7418 let mut reasoning_decision: Option<ReasoningMode> = None;
7419 let mut transition_fallback_required = false;
7420 let mut skill_fallback_required = false;
7421 let mut reasoning_fallback_required = false;
7422
7423 loop {
7424 if let Some(candidate) = transition_candidate.take() {
7425 if self
7426 .approve_transition_target(&candidate.from_state, candidate.target())
7427 .await?
7428 {
7429 self.finalize_pending_branches(branch_set.cancel_pending());
7431 if !main_pending {
7432 self.finalize_branch_loss(
7433 &main_id,
7434 main_kind,
7435 RuntimeCommitBehavior::FinalResponse,
7436 false,
7437 main_result.as_ref().map(|result| result.is_err()),
7438 );
7439 }
7440 if skill_enabled && !skill_pending {
7441 self.finalize_branch_loss(
7442 &skill_id,
7443 RuntimeOptimizationKind::SpeculativeSkillRouting,
7444 RuntimeCommitBehavior::SkillSelection,
7445 false,
7446 Some(false),
7447 );
7448 }
7449 if reasoning_enabled && !reasoning_pending {
7450 self.finalize_branch_loss(
7451 &reasoning_id,
7452 RuntimeOptimizationKind::SpeculativeReasoningAuto,
7453 RuntimeCommitBehavior::ReasoningDecision,
7454 false,
7455 Some(false),
7456 );
7457 }
7458 if !self
7459 .apply_pre_response_transition_candidate(
7460 &candidate,
7461 &HashMap::new(),
7462 processed_input,
7463 )
7464 .await?
7465 {
7466 self.finalize_optional_branch(
7467 &transition_id,
7468 RuntimeOptimizationKind::ParallelStateTransition,
7469 RuntimeCommitBehavior::TransitionDecision,
7470 "discarded",
7471 false,
7472 );
7473 return Ok(None);
7474 }
7475 self.finalize_optional_branch(
7476 &transition_id,
7477 RuntimeOptimizationKind::ParallelStateTransition,
7478 RuntimeCommitBehavior::TransitionDecision,
7479 "committed",
7480 true,
7481 );
7482 return self
7483 .redispatch_current_state(processed_input)
7484 .await
7485 .map(Some);
7486 }
7487 self.finalize_optional_branch(
7488 &transition_id,
7489 RuntimeOptimizationKind::ParallelStateTransition,
7490 RuntimeCommitBehavior::TransitionDecision,
7491 "discarded",
7492 false,
7493 );
7494 transition_finalized = true;
7495 }
7496
7497 if transition_finalized && skill_candidate.is_some() {
7498 let candidate = skill_candidate.take().unwrap();
7499 self.finalize_optional_branch(
7500 &skill_id,
7501 RuntimeOptimizationKind::SpeculativeSkillRouting,
7502 RuntimeCommitBehavior::SkillSelection,
7503 "committed",
7504 true,
7505 );
7506 if !main_pending {
7507 self.finalize_branch_loss(
7508 &main_id,
7509 main_kind,
7510 RuntimeCommitBehavior::FinalResponse,
7511 false,
7512 main_result.as_ref().map(|result| result.is_err()),
7513 );
7514 }
7515 if reasoning_enabled && !reasoning_pending {
7516 self.finalize_branch_loss(
7517 &reasoning_id,
7518 RuntimeOptimizationKind::SpeculativeReasoningAuto,
7519 RuntimeCommitBehavior::ReasoningDecision,
7520 false,
7521 Some(false),
7522 );
7523 }
7524 self.finalize_pending_branches(branch_set.cancel_pending());
7525 self.commit_root_user_message(processed_input).await?;
7526 return match self
7527 .commit_skill_candidate_route_result(candidate, processed_input)
7528 .await?
7529 {
7530 SkillRouteResult::Response { skill_id, content } => self
7531 .handle_skill_response(processed_input, &skill_id, content, input_context)
7532 .await
7533 .map(Some),
7534 SkillRouteResult::NeedsClarification {
7535 response,
7536 ownership,
7537 } => {
7538 let admission = self
7539 .admit_optional_disambiguation_ownership(ownership)
7540 .await?;
7541 if response
7542 .metadata
7543 .as_ref()
7544 .and_then(|m| m.get("disambiguation"))
7545 .and_then(|d| d.get("status"))
7546 .and_then(|s| s.as_str())
7547 == Some("awaiting_clarification")
7548 {
7549 self.memory
7550 .add_message(ChatMessage::assistant(&response.content))
7551 .await?;
7552 }
7553 drop(admission);
7554 self.finish_turn_if_root(&response).await?;
7555 Ok(Some(response))
7556 }
7557 SkillRouteResult::NoMatch => Ok(None),
7558 };
7559 }
7560
7561 if transition_finalized
7562 && skill_finalized
7563 && let Some(reasoning_mode) = reasoning_decision.take()
7564 {
7565 if !matches!(reasoning_mode, ReasoningMode::None) {
7566 self.finalize_optional_branch(
7567 &reasoning_id,
7568 RuntimeOptimizationKind::SpeculativeReasoningAuto,
7569 RuntimeCommitBehavior::ReasoningDecision,
7570 "committed",
7571 true,
7572 );
7573 if !main_pending {
7574 self.finalize_branch_loss(
7575 &main_id,
7576 main_kind,
7577 RuntimeCommitBehavior::FinalResponse,
7578 false,
7579 main_result.as_ref().map(|result| result.is_err()),
7580 );
7581 }
7582 self.finalize_pending_branches(branch_set.cancel_pending());
7583 self.commit_root_user_message(processed_input).await?;
7584 return if matches!(reasoning_mode, ReasoningMode::PlanAndExecute) {
7585 self.handle_plan_and_execute(processed_input, input_context, true)
7586 .await
7587 .map(Some)
7588 } else {
7589 self.run_committed_response_loop_with_reasoning(
7590 processed_input,
7591 input_context,
7592 reasoning_mode,
7593 true,
7594 )
7595 .await
7596 .map(Some)
7597 };
7598 }
7599 self.finalize_optional_branch(
7600 &reasoning_id,
7601 RuntimeOptimizationKind::SpeculativeReasoningAuto,
7602 RuntimeCommitBehavior::ReasoningDecision,
7603 "committed",
7604 true,
7605 );
7606 reasoning_finalized = true;
7607 }
7608
7609 if transition_finalized && skill_finalized && reasoning_finalized {
7610 if transition_fallback_required
7611 || skill_fallback_required
7612 || reasoning_fallback_required
7613 {
7614 if !main_pending {
7615 self.finalize_branch_loss(
7616 &main_id,
7617 main_kind,
7618 RuntimeCommitBehavior::FinalResponse,
7619 false,
7620 main_result.as_ref().map(|result| result.is_err()),
7621 );
7622 }
7623 self.finalize_pending_branches(branch_set.cancel_pending());
7624 return Ok(None);
7625 }
7626
7627 if let Some(result) = main_result.take() {
7628 let draft = match result {
7629 Ok(draft) => draft,
7630 Err(error) => {
7631 self.finalize_optional_branch(
7632 &main_id,
7633 main_kind,
7634 RuntimeCommitBehavior::FinalResponse,
7635 "failed",
7636 false,
7637 );
7638 self.finalize_pending_branches(branch_set.cancel_pending());
7639 return Err(error);
7640 }
7641 };
7642 self.finalize_optional_branch(
7643 &main_id,
7644 main_kind,
7645 RuntimeCommitBehavior::FinalResponse,
7646 "committed",
7647 true,
7648 );
7649 self.finalize_pending_branches(branch_set.cancel_pending());
7650 return self
7651 .commit_main_response_draft(
7652 processed_input,
7653 input_context,
7654 draft,
7655 ReasoningMode::None,
7656 reasoning_enabled,
7657 )
7658 .await
7659 .map(Some);
7660 }
7661 }
7662
7663 if branch_set.is_empty() {
7664 return Ok(None);
7665 }
7666
7667 let Some(outcome) = branch_set.next_completed().await else {
7668 return Ok(None);
7669 };
7670 let branch_id = outcome.branch.branch_id();
7671 match outcome.result {
7672 RuntimeBranchResult::MainDraft(draft) => {
7673 main_pending = false;
7674 main_result = Some(Ok(draft));
7675 }
7676 RuntimeBranchResult::Transition(candidate) => {
7677 if let Some(candidate) = candidate {
7678 transition_candidate = Some(candidate);
7679 } else {
7680 self.finalize_optional_branch(
7681 &transition_id,
7682 RuntimeOptimizationKind::ParallelStateTransition,
7683 RuntimeCommitBehavior::TransitionDecision,
7684 "discarded",
7685 false,
7686 );
7687 transition_finalized = true;
7688 }
7689 }
7690 RuntimeBranchResult::Skill(candidate) => {
7691 skill_pending = false;
7692 if let Some(candidate) = candidate {
7693 skill_candidate = Some(candidate);
7694 } else {
7695 self.finalize_optional_branch(
7696 &skill_id,
7697 RuntimeOptimizationKind::SpeculativeSkillRouting,
7698 RuntimeCommitBehavior::SkillSelection,
7699 "discarded",
7700 false,
7701 );
7702 skill_finalized = true;
7703 }
7704 }
7705 RuntimeBranchResult::Reasoning(mode) => {
7706 reasoning_pending = false;
7707 reasoning_decision = Some(mode);
7708 }
7709 RuntimeBranchResult::Failed(error) => {
7710 if branch_id == main_id {
7711 main_pending = false;
7712 main_result = Some(Err(error));
7713 } else if branch_id == transition_id {
7714 self.finalize_optional_branch(
7715 &transition_id,
7716 RuntimeOptimizationKind::ParallelStateTransition,
7717 RuntimeCommitBehavior::TransitionDecision,
7718 "failed",
7719 false,
7720 );
7721 transition_finalized = true;
7722 } else if branch_id == skill_id {
7723 skill_pending = false;
7724 self.finalize_optional_branch(
7725 &skill_id,
7726 RuntimeOptimizationKind::SpeculativeSkillRouting,
7727 RuntimeCommitBehavior::SkillSelection,
7728 "failed",
7729 false,
7730 );
7731 skill_finalized = true;
7732 } else if branch_id == reasoning_id {
7733 reasoning_pending = false;
7734 self.finalize_optional_branch(
7735 &reasoning_id,
7736 RuntimeOptimizationKind::SpeculativeReasoningAuto,
7737 RuntimeCommitBehavior::ReasoningDecision,
7738 "failed",
7739 false,
7740 );
7741 reasoning_finalized = true;
7742 }
7743 }
7744 RuntimeBranchResult::Cancelled => {
7745 self.finalize_optional_branch(
7746 &branch_id,
7747 outcome.branch.optimization,
7748 outcome.branch.commit_behavior,
7749 "cancelled",
7750 false,
7751 );
7752 if branch_id == main_id {
7753 main_pending = false;
7754 main_result =
7755 Some(Err(AgentError::Other("main branch cancelled".to_string())));
7756 } else if branch_id == transition_id {
7757 transition_finalized = true;
7758 transition_fallback_required = true;
7759 } else if branch_id == skill_id {
7760 skill_pending = false;
7761 skill_finalized = true;
7762 skill_fallback_required = true;
7763 } else if branch_id == reasoning_id {
7764 reasoning_pending = false;
7765 reasoning_finalized = true;
7766 reasoning_fallback_required = true;
7767 }
7768 }
7769 }
7770 }
7771 }
7772
7773 fn finalize_pending_branches(&self, branches: Vec<RuntimeBranch>) {
7774 for branch in branches {
7775 self.finalize_optional_branch(
7776 &branch.branch_id(),
7777 branch.optimization,
7778 branch.commit_behavior,
7779 "cancelled",
7780 false,
7781 );
7782 }
7783 }
7784
7785 fn finalize_branch_loss(
7790 &self,
7791 branch_id: &str,
7792 optimization: RuntimeOptimizationKind,
7793 commit_behavior: RuntimeCommitBehavior,
7794 pending: bool,
7795 completed_failed: Option<bool>,
7796 ) {
7797 let status = if pending {
7798 "cancelled"
7799 } else if completed_failed.unwrap_or(false) {
7800 "failed"
7801 } else {
7802 "discarded"
7803 };
7804 self.finalize_optional_branch(branch_id, optimization, commit_behavior, status, false);
7805 }
7806
7807 fn finalize_optional_branch(
7812 &self,
7813 branch_id: &str,
7814 optimization: RuntimeOptimizationKind,
7815 commit_behavior: RuntimeCommitBehavior,
7816 status: &str,
7817 winner: bool,
7818 ) {
7819 crate::optimization::observability::finalize_branch(
7820 self.observability_manager.as_ref(),
7821 branch_id,
7822 status,
7823 winner,
7824 optimization,
7825 commit_behavior,
7826 );
7827 }
7828
7829 fn has_parallel_transition_candidates(&self) -> bool {
7834 self.transitions_available_for_commit()
7835 .map(|(transitions, _)| {
7836 transitions
7837 .iter()
7838 .any(|transition| matches!(transition.timing, TransitionTiming::Parallel))
7839 })
7840 .unwrap_or(false)
7841 }
7842
7843 async fn select_parallel_transition_candidate(
7848 &self,
7849 processed_input: &str,
7850 ) -> Result<ParallelTransitionSelection> {
7851 let Some((transitions, current_state)) = self.transitions_available_for_commit() else {
7852 return Ok(ParallelTransitionSelection::NoMatch);
7853 };
7854 let parallel: Vec<Transition> = transitions
7855 .into_iter()
7856 .filter(|transition| matches!(transition.timing, TransitionTiming::Parallel))
7857 .filter(|transition| !transition.requires_response)
7858 .collect();
7859 if parallel.is_empty() {
7860 return Ok(ParallelTransitionSelection::NoMatch);
7861 }
7862 let empty_staged = HashMap::new();
7863 if let Some(candidate) = self.select_deterministic_transition_candidate(
7864 processed_input,
7865 ¤t_state,
7866 ¶llel,
7867 &empty_staged,
7868 ) {
7869 return Ok(ParallelTransitionSelection::Candidate(candidate));
7870 }
7871 let when_transitions: Vec<(usize, &Transition)> = parallel
7872 .iter()
7873 .enumerate()
7874 .filter(|(_, transition)| !transition.when.trim().is_empty())
7875 .collect();
7876 if when_transitions.is_empty() {
7877 return Ok(ParallelTransitionSelection::NoMatch);
7878 }
7879 let llm = self
7880 .llm_registry
7881 .router()
7882 .or_else(|_| self.llm_registry.default())
7883 .map_err(|e| AgentError::Config(e.to_string()))?;
7884 let conditions = when_transitions
7885 .iter()
7886 .enumerate()
7887 .map(|(display_idx, (_, transition))| {
7888 format!("{}. {}", display_idx + 1, transition.when)
7889 })
7890 .collect::<Vec<_>>()
7891 .join("\n");
7892 if !self
7893 .reserve_active_speculative_llm_call(RuntimeOptimizationKind::ParallelStateTransition)
7894 {
7895 return Ok(ParallelTransitionSelection::ReservationExhausted);
7896 }
7897 let context_preview = self.branch_context_preview();
7898 let prompt = format!(
7899 "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-{}).",
7900 current_state,
7901 processed_input,
7902 context_preview,
7903 conditions,
7904 when_transitions.len()
7905 );
7906 let response = self
7907 .observe_purpose(
7908 ObservationPurpose::StateTransitionEvaluation,
7909 llm.complete(&[ChatMessage::user(prompt)], None),
7910 )
7911 .await
7912 .map_err(|e| AgentError::LLM(e.to_string()))?;
7913 let choice = response.content.trim().parse::<usize>().unwrap_or(0);
7914 if choice == 0 || choice > when_transitions.len() {
7915 return Ok(ParallelTransitionSelection::NoMatch);
7916 }
7917 let transition = when_transitions[choice - 1].1.clone();
7918 Ok(ParallelTransitionSelection::Candidate(
7919 TransitionCandidate::new(
7920 current_state,
7921 transition.clone(),
7922 Self::transition_reason(&transition),
7923 ),
7924 ))
7925 }
7926
7927 async fn redispatch_current_state(&self, processed_input: &str) -> Result<AgentResponse> {
7929 const MAX_REDISPATCH_DEPTH: u32 = 3;
7930 let current_depth = *self.redispatch_depth.read();
7931 if current_depth >= MAX_REDISPATCH_DEPTH {
7932 warn!(depth = current_depth, "Re-dispatch depth limit reached");
7933 let response = AgentResponse::new("");
7934 self.finish_turn_if_root(&response).await?;
7935 return Ok(response);
7936 }
7937 *self.redispatch_depth.write() += 1;
7938 if let Some(context) = self.active_turn_context.write().as_mut() {
7939 context.enter_redispatch();
7940 }
7941 let result = Box::pin(self.run_loop_internal(processed_input)).await;
7942 *self.redispatch_depth.write() -= 1;
7943 if let Some(context) = self.active_turn_context.write().as_mut() {
7944 context.exit_redispatch();
7945 }
7946 let response = result?;
7947 self.finish_turn_if_root(&response).await?;
7948 Ok(response)
7949 }
7950
7951 async fn finish_turn_if_root(&self, response: &AgentResponse) -> Result<()> {
7953 if *self.redispatch_depth.read() == 0 {
7954 self.post_turn_session_lifecycle().await?;
7955 if let Some(context) = self.active_turn_context.write().as_mut() {
7956 context.mark_post_turn_lifecycle_completed();
7957 }
7958 self.hooks.on_response(response).await;
7959 self.end_root_turn();
7960 }
7961 Ok(())
7962 }
7963
7964 async fn execute_state_exit_actions(&self, state_path: &str) {
7966 if let Some(ref sm) = self.state_machine
7967 && let Some(def) = sm.get_definition(state_path)
7968 && !def.on_exit.is_empty()
7969 {
7970 debug!(state = %state_path, count = def.on_exit.len(), "Executing on_exit actions");
7971 self.execute_state_actions(&def.on_exit).await;
7972 }
7973 }
7974
7975 fn state_was_previously_entered(
7977 state_path: &str,
7978 from_state: &str,
7979 history_before: &[StateTransitionEvent],
7980 ) -> bool {
7981 state_path == from_state
7982 || history_before
7983 .iter()
7984 .any(|event| event.from == state_path || event.to == state_path)
7985 }
7986
7987 async fn execute_state_enter_actions(&self, state_path: &str, is_reentry: bool) {
7989 if let Some(ref sm) = self.state_machine
7990 && let Some(def) = sm.get_definition(state_path)
7991 {
7992 if is_reentry && !def.on_reenter.is_empty() {
7993 debug!(state = %state_path, count = def.on_reenter.len(), "Executing on_reenter actions");
7994 self.execute_state_actions(&def.on_reenter).await;
7995 } else if !def.on_enter.is_empty() {
7996 debug!(state = %state_path, count = def.on_enter.len(), "Executing on_enter actions");
7997 self.execute_state_actions(&def.on_enter).await;
7998 }
7999 }
8000 }
8001
8002 async fn execute_state_actions(&self, actions: &[StateAction]) {
8004 for (action_index, action) in actions.iter().enumerate() {
8005 match action {
8006 StateAction::Tool { tool, args } => {
8007 let raw_args = args.clone().unwrap_or(Value::Object(Default::default()));
8008 let args_value = self.render_action_args(&raw_args);
8009 let state = self.state_machine.as_ref().map(|sm| sm.current());
8010 let request = ToolExecutionRequest::new(
8011 uuid::Uuid::new_v4().to_string(),
8012 tool.clone(),
8013 args_value,
8014 ToolCallSource::StateAction {
8015 state,
8016 action_index,
8017 },
8018 );
8019 match self.execute_tool_record(request).await {
8020 Ok(record) if record.success => {
8021 debug!(tool = %record.canonical_id, "State action: tool executed");
8022 let _ = self.context_manager.set(
8023 "last_tool_result",
8024 serde_json::Value::String(record.model_output_string()),
8025 );
8026 let _ = self.context_manager.set(
8027 "last_tool_record",
8028 serde_json::to_value(record).unwrap_or(Value::Null),
8029 );
8030 }
8031 Ok(record) => {
8032 warn!(tool = %record.canonical_id, error = %record.output, "State action: tool failed");
8033 }
8034 Err(e) => {
8035 warn!(tool = %tool, error = %e, "State action: tool failed")
8036 }
8037 }
8038 }
8039 StateAction::Skill { skill } => {
8040 if let Some(ref executor) = self.skill_executor {
8041 if let Some(def) = self.skills.iter().find(|s| s.id == *skill) {
8042 match executor
8043 .execute_with_invoker(def, "", serde_json::json!({}), self)
8044 .await
8045 {
8046 Ok(_) => debug!(skill = %skill, "State action: skill executed"),
8047 Err(e) => {
8048 warn!(skill = %skill, error = %e, "State action: skill failed")
8049 }
8050 }
8051 } else {
8052 warn!(skill = %skill, "State action: skill not found");
8053 }
8054 }
8055 }
8056 StateAction::SetContext { set_context } => {
8057 for (key, value) in set_context {
8058 if let Err(e) = self.context_manager.set(key, value.clone()) {
8059 warn!(key = %key, error = %e, "State action: set_context failed");
8060 } else {
8061 debug!(key = %key, "State action: context set");
8062 }
8063 }
8064 }
8065 StateAction::Prompt {
8066 prompt,
8067 llm,
8068 store_as,
8069 } => {
8070 let llm_result = if let Some(alias) = llm {
8071 self.llm_registry.get(alias)
8072 } else {
8073 self.llm_registry.default()
8074 };
8075 match llm_result {
8076 Ok(llm_provider) => {
8077 let context = self.build_context_with_overlays();
8079 let rendered_prompt = self
8080 .template_renderer
8081 .render(prompt, &context)
8082 .unwrap_or_else(|_| prompt.clone());
8083 let recent =
8084 self.memory.get_messages(Some(5)).await.unwrap_or_default();
8085 let mut messages: Vec<ChatMessage> = recent;
8086 messages.push(ChatMessage::user(&rendered_prompt));
8087 match self
8088 .observe_purpose(
8089 ObservationPurpose::StateAction,
8090 llm_provider.complete(&messages, None),
8091 )
8092 .await
8093 {
8094 Ok(response) => {
8095 if let Some(key) = store_as {
8096 let _ = self
8097 .context_manager
8098 .set(key, Value::String(response.content));
8099 debug!(key = %key, "State action: prompt result stored");
8100 }
8101 }
8102 Err(e) => {
8103 warn!(error = %e, "State action: prompt LLM call failed");
8104 }
8105 }
8106 }
8107 Err(e) => {
8108 warn!(error = %e, "State action: LLM not found for prompt");
8109 }
8110 }
8111 }
8112 }
8113 }
8114 }
8115
8116 async fn run_context_extractors_staged(&self, user_message: &str) -> HashMap<String, Value> {
8117 let extractors = match &self.state_machine {
8118 Some(sm) => match sm.current_definition() {
8119 Some(def) if !def.extract.is_empty() => def.extract.clone(),
8120 _ => return HashMap::new(),
8121 },
8122 None => return HashMap::new(),
8123 };
8124
8125 let mut staged = HashMap::new();
8126 for extractor in &extractors {
8127 let prompt = if let Some(ref custom) = extractor.llm_extract {
8128 format!(
8129 "User message:\n\"{}\"\n\nInstruction:\n{}",
8130 user_message, custom
8131 )
8132 } else if let Some(ref desc) = extractor.description {
8133 format!(
8134 "From the following message, extract: {}\n\n\
8135 Message: \"{}\"\n\n\
8136 If the information is present, return ONLY the extracted value.\n\
8137 If NOT present, return exactly: __NONE__",
8138 desc, user_message
8139 )
8140 } else {
8141 continue;
8142 };
8143
8144 let llm = match self
8145 .llm_registry
8146 .get(&extractor.llm)
8147 .or_else(|_| self.llm_registry.get("router"))
8148 .or_else(|_| self.llm_registry.get("default"))
8149 {
8150 Ok(llm) => llm,
8151 Err(e) => {
8152 warn!(key = %extractor.key, error = %e, "Extractor LLM not found");
8153 continue;
8154 }
8155 };
8156
8157 let messages = vec![ChatMessage::user(&prompt)];
8158 match self
8159 .observe_purpose(
8160 ObservationPurpose::ContextExtraction,
8161 llm.complete(&messages, None),
8162 )
8163 .await
8164 {
8165 Ok(response) => {
8166 let value = response.content.trim().to_string();
8167 if value != "__NONE__" && !value.is_empty() {
8168 staged.insert(
8169 extractor.key.clone(),
8170 serde_json::Value::String(value.clone()),
8171 );
8172 debug!(key = %extractor.key, value = %value, "Context extracted");
8173 } else if extractor.required {
8174 warn!(key = %extractor.key, "Required extraction returned no value");
8175 }
8176 }
8177 Err(e) => {
8178 warn!(key = %extractor.key, error = %e, "Context extraction LLM call failed");
8179 }
8180 }
8181 }
8182 staged
8183 }
8184
8185 fn commit_staged_context_writes(&self, staged: &HashMap<String, Value>) {
8186 for (key, value) in staged {
8187 if let Err(error) = self.context_manager.update(key, value.clone()) {
8188 warn!(key = %key, error = %error, "staged context write failed");
8189 }
8190 }
8191 }
8192
8193 async fn run_context_extractors(&self, user_message: &str) {
8195 let staged = self.run_context_extractors_staged(user_message).await;
8196 self.commit_staged_context_writes(&staged);
8197 }
8198
8199 async fn check_memory_compression(&self) -> Result<()> {
8200 if self.memory.needs_compression() {
8201 let result = self.memory.compress(None).await?;
8202 if let CompressResult::Compressed {
8203 messages_summarized,
8204 new_summary_length,
8205 tokens_saved,
8206 } = result
8207 {
8208 let event = MemoryCompressEvent::new(
8209 messages_summarized,
8210 tokens_saved,
8211 new_summary_length as u32,
8212 );
8213 self.hooks.on_memory_compress(&event).await;
8214 debug!(
8215 messages = messages_summarized,
8216 tokens_saved = tokens_saved,
8217 "Memory compressed"
8218 );
8219 }
8220 }
8221
8222 self.handle_memory_overflow().await?;
8224 self.check_memory_budget().await;
8225
8226 Ok(())
8227 }
8228
8229 async fn check_memory_budget(&self) {
8230 let Some(ref budget) = self.memory_token_budget else {
8231 return;
8232 };
8233
8234 let context = match self.memory.get_context().await {
8235 Ok(ctx) => ctx,
8236 Err(_) => return,
8237 };
8238
8239 let used_tokens = context.estimated_tokens();
8241 if budget.is_over_warn_threshold(used_tokens) {
8242 let event = MemoryBudgetEvent::new("memory", used_tokens, budget.total);
8243 self.hooks.on_memory_budget_warning(&event).await;
8244 debug!(
8245 used = used_tokens,
8246 total = budget.total,
8247 percent = event.usage_percent,
8248 "Memory budget warning"
8249 );
8250 }
8251
8252 if let Some(ref summary) = context.summary {
8254 let summary_tokens = ai_agents_memory::estimate_tokens(summary);
8255 let summary_budget = budget.allocation.summary;
8256 if summary_budget > 0 {
8257 let warn_threshold =
8258 (summary_budget as f64 * budget.warn_at_percent as f64 / 100.0) as u32;
8259 if summary_tokens >= warn_threshold {
8260 let event = MemoryBudgetEvent::new("summary", summary_tokens, summary_budget);
8261 self.hooks.on_memory_budget_warning(&event).await;
8262 }
8263 }
8264 }
8265
8266 let recent_tokens: u32 = context
8268 .messages
8269 .iter()
8270 .map(ai_agents_memory::estimate_message_tokens)
8271 .sum();
8272 let recent_budget = budget.allocation.recent_messages;
8273 if recent_budget > 0 {
8274 let warn_threshold =
8275 (recent_budget as f64 * budget.warn_at_percent as f64 / 100.0) as u32;
8276 if recent_tokens >= warn_threshold {
8277 let event = MemoryBudgetEvent::new("recent_messages", recent_tokens, recent_budget);
8278 self.hooks.on_memory_budget_warning(&event).await;
8279 }
8280 }
8281
8282 let relationship_budget = budget.allocation.relationships;
8283 if relationship_budget > 0 {
8284 let relationship_tokens = self
8285 .relationship_memory_text()
8286 .map(|text| ai_agents_memory::estimate_tokens(&text))
8287 .unwrap_or(0);
8288 let warn_threshold =
8289 (relationship_budget as f64 * budget.warn_at_percent as f64 / 100.0) as u32;
8290 if relationship_tokens >= warn_threshold {
8291 let event = MemoryBudgetEvent::new(
8292 "relationships",
8293 relationship_tokens,
8294 relationship_budget,
8295 );
8296 self.hooks.on_memory_budget_warning(&event).await;
8297 }
8298 }
8299 }
8300
8301 async fn handle_memory_overflow(&self) -> Result<()> {
8302 let Some(ref budget) = self.memory_token_budget else {
8303 return Ok(());
8304 };
8305
8306 let context = self.memory.get_context().await?;
8307 let used_tokens = context.estimated_tokens();
8308
8309 if used_tokens <= budget.total {
8310 return Ok(());
8311 }
8312
8313 match budget.overflow_strategy {
8314 OverflowStrategy::TruncateOldest => {
8315 let tokens_to_free = used_tokens - budget.total;
8316 let messages_to_evict = self.calculate_eviction_count(tokens_to_free);
8317 if messages_to_evict > 0 {
8318 self.evict_messages(messages_to_evict, EvictionReason::TokenBudgetExceeded)
8319 .await?;
8320 }
8321 }
8322 OverflowStrategy::SummarizeMore => {
8323 let max_attempts = context.total_messages.max(1);
8324 for _ in 0..max_attempts {
8325 match self.memory.compress(None).await? {
8326 CompressResult::Compressed {
8327 messages_summarized,
8328 ..
8329 } if messages_summarized > 0 => {
8330 let context = self.memory.get_context().await?;
8331 if context.estimated_tokens() <= budget.total {
8332 return Ok(());
8333 }
8334 }
8335 _ => break,
8336 }
8337 }
8338 let context = self.memory.get_context().await?;
8339 let used_tokens = context.estimated_tokens();
8340 if used_tokens > budget.total {
8341 return Err(AgentError::MemoryBudgetExceeded {
8342 used: used_tokens,
8343 budget: budget.total,
8344 });
8345 }
8346 }
8347 OverflowStrategy::Error => {
8348 return Err(AgentError::MemoryBudgetExceeded {
8349 used: used_tokens,
8350 budget: budget.total,
8351 });
8352 }
8353 }
8354 Ok(())
8355 }
8356
8357 fn calculate_eviction_count(&self, tokens_to_free: u32) -> usize {
8358 ((tokens_to_free as f64 / 50.0).ceil() as usize).max(1)
8360 }
8361
8362 async fn evict_messages(&self, count: usize, reason: EvictionReason) -> Result<()> {
8363 let evicted = self.memory.evict_oldest(count).await?;
8364 if !evicted.is_empty() {
8365 let event = MemoryEvictEvent {
8366 reason,
8367 messages_evicted: evicted.len(),
8368 importance_scores: vec![],
8369 };
8370 self.hooks.on_memory_evict(&event).await;
8371 debug!(count = evicted.len(), "Messages evicted from memory");
8372 }
8373 Ok(())
8374 }
8375
8376 #[instrument(skip(self, input), fields(agent = %self.info.name))]
8377 async fn determine_reasoning_mode(&self, input: &str) -> Result<ReasoningMode> {
8378 match self.determine_reasoning_mode_strict(input).await {
8379 Ok(mode) => Ok(mode),
8380 Err(_) => Ok(ReasoningMode::None),
8381 }
8382 }
8383
8384 async fn determine_reasoning_mode_strict(&self, input: &str) -> Result<ReasoningMode> {
8385 let effective_config = self.get_effective_reasoning_config();
8386
8387 if !matches!(effective_config.mode, ReasoningMode::Auto) {
8388 return Ok(effective_config.mode.clone());
8389 }
8390
8391 let judge_llm = effective_config
8392 .judge_llm
8393 .as_ref()
8394 .and_then(|alias| self.llm_registry.get(alias).ok())
8395 .or_else(|| self.llm_registry.router().ok())
8396 .or_else(|| self.llm_registry.default().ok());
8397
8398 let Some(llm) = judge_llm else {
8399 return Ok(ReasoningMode::None);
8400 };
8401
8402 let prompt = format!(
8403 r#"Analyze this user request and determine the appropriate reasoning mode.
8404
8405User request: "{}"
8406
8407Choose ONE of these modes:
8408- none: Simple queries, greetings, direct answers (fastest)
8409- cot: Complex analysis, multi-step reasoning, math problems
8410- react: Tasks requiring multiple tool calls with observation
8411- plan_and_execute: Complex multi-step tasks requiring coordination
8412
8413Respond with ONLY the mode name (none, cot, react, or plan_and_execute)."#,
8414 input
8415 );
8416
8417 let messages = vec![ChatMessage::user(&prompt)];
8418 let response = self
8419 .observe_purpose(
8420 ObservationPurpose::ReflectionDecision,
8421 llm.complete(&messages, None),
8422 )
8423 .await
8424 .map_err(|e| AgentError::LLM(e.to_string()))?;
8425
8426 let mode_str = response.content.trim().to_lowercase();
8427 Ok(match mode_str.as_str() {
8428 "cot" => ReasoningMode::CoT,
8429 "react" => ReasoningMode::React,
8430 "plan_and_execute" => ReasoningMode::PlanAndExecute,
8431 _ => ReasoningMode::None,
8432 })
8433 }
8434
8435 async fn should_reflect(&self, input: &str, response: &str) -> Result<bool> {
8436 let effective_config = self.get_effective_reflection_config();
8437
8438 if !effective_config.requires_evaluation() {
8439 return Ok(false);
8440 }
8441
8442 if effective_config.is_enabled() {
8443 return Ok(true);
8444 }
8445
8446 let evaluator_llm = effective_config
8447 .evaluator_llm
8448 .as_ref()
8449 .and_then(|alias| self.llm_registry.get(alias).ok())
8450 .or_else(|| self.llm_registry.router().ok())
8451 .or_else(|| self.llm_registry.default().ok());
8452
8453 let Some(llm) = evaluator_llm else {
8454 return Ok(false);
8455 };
8456
8457 let response_preview: String = response.chars().take(500).collect();
8458 let prompt = format!(
8459 r#"Should this response be evaluated for quality? Consider if it's a complex or important response.
8460
8461User query: "{}"
8462Response: "{}"
8463
8464Answer YES or NO only."#,
8465 input, response_preview
8466 );
8467
8468 let messages = vec![ChatMessage::user(&prompt)];
8469 let result = self
8470 .observe_purpose(
8471 ObservationPurpose::ReflectionDecision,
8472 llm.complete(&messages, None),
8473 )
8474 .await;
8475
8476 match result {
8477 Ok(resp) => Ok(resp.content.trim().to_uppercase().contains("YES")),
8478 Err(_) => Ok(false),
8479 }
8480 }
8481
8482 fn build_cot_system_prompt(&self, base_prompt: &str) -> String {
8483 format!(
8484 "{}\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>",
8485 base_prompt
8486 )
8487 }
8488
8489 fn build_react_system_prompt(&self, base_prompt: &str) -> String {
8490 format!(
8491 "{}\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>",
8492 base_prompt
8493 )
8494 }
8495
8496 async fn generate_plan(&self, input: &str) -> Result<Plan> {
8497 let effective = self.get_effective_reasoning_config();
8498 let planning_config = effective.get_planning();
8499
8500 let planner_llm = planning_config
8501 .and_then(|c| c.planner_llm.as_ref())
8502 .and_then(|alias| self.llm_registry.get(alias).ok())
8503 .or_else(|| self.llm_registry.router().ok())
8504 .or_else(|| self.llm_registry.default().ok())
8505 .ok_or_else(|| AgentError::Config("No LLM available for planning".into()))?;
8506
8507 let mut available_tool_ids: Vec<String> = self
8508 .get_available_tool_ids()
8509 .await
8510 .unwrap_or_else(|_| self.tools.list_ids());
8511 let mut available_skills: Vec<String> = self.skills.iter().map(|s| s.id.clone()).collect();
8512
8513 if let Some(config) = planning_config {
8515 if !config.available.tools.is_all() {
8516 available_tool_ids.retain(|t| config.available.tools.allows(t));
8517 }
8518 if !config.available.skills.is_all() {
8519 available_skills.retain(|s| config.available.skills.allows(s));
8520 }
8521 }
8522
8523 let tool_descriptions: Vec<String> = available_tool_ids
8526 .iter()
8527 .filter_map(|id| {
8528 self.tools.get(id).map(|tool| {
8529 let schema = tool.input_schema();
8530 let args_desc = schema
8531 .get("properties")
8532 .and_then(|p| serde_json::to_string(p).ok())
8533 .unwrap_or_else(|| "{}".to_string());
8534 format!(
8535 "- {} ({}): {}\n Arguments: {}",
8536 id,
8537 tool.name(),
8538 tool.description(),
8539 args_desc
8540 )
8541 })
8542 })
8543 .collect();
8544
8545 let tools_section = if tool_descriptions.is_empty() {
8546 "Available tools: none".to_string()
8547 } else {
8548 format!("Available tools:\n{}", tool_descriptions.join("\n"))
8549 };
8550
8551 let skills_section = if available_skills.is_empty() {
8552 "Available skills: none".to_string()
8553 } else {
8554 format!("Available skills: {}", available_skills.join(", "))
8555 };
8556
8557 let prompt = format!(
8558 r#"Create a step-by-step plan to accomplish this goal.
8559
8560Goal: "{}"
8561
8562{}
8563
8564{}
8565
8566Create a plan with clear steps. For each step, specify:
8567- description: What this step accomplishes
8568- action_type: "tool", "skill", "think", or "respond"
8569- action_target: The tool/skill id (if applicable)
8570- args: The arguments object matching the tool's schema (if action_type is "tool")
8571- dependencies: List of step IDs this depends on (empty if none)
8572
8573Respond in JSON format:
8574{{
8575 "steps": [
8576 {{"id": "step1", "description": "...", "action_type": "tool", "action_target": "tool_id", "args": {{"required_field": "value"}}, "dependencies": []}},
8577 {{"id": "step2", "description": "...", "action_type": "think", "action_target": "...", "dependencies": ["step1"]}}
8578 ]
8579}}"#,
8580 input, tools_section, skills_section,
8581 );
8582
8583 let messages = vec![ChatMessage::user(&prompt)];
8584 let response = self
8585 .observe_purpose(
8586 ObservationPurpose::PlanGeneration,
8587 planner_llm.complete(&messages, None),
8588 )
8589 .await
8590 .map_err(|e| AgentError::LLM(format!("Planning failed: {}", e)))?;
8591
8592 let mut plan = Plan::new(input);
8593
8594 if let Some(json_start) = response.content.find('{')
8595 && let Some(json_end) = response.content.rfind('}')
8596 {
8597 let json_str = &response.content[json_start..=json_end];
8598 if let Ok(parsed) = serde_json::from_str::<serde_json::Value>(json_str)
8599 && let Some(steps) = parsed.get("steps").and_then(|s| s.as_array())
8600 {
8601 for step_value in steps {
8602 let id = step_value
8603 .get("id")
8604 .and_then(|v| v.as_str())
8605 .unwrap_or("step");
8606 let desc = step_value
8607 .get("description")
8608 .and_then(|v| v.as_str())
8609 .unwrap_or("");
8610 let action_type = step_value
8611 .get("action_type")
8612 .and_then(|v| v.as_str())
8613 .unwrap_or("think");
8614 let action_target = step_value
8615 .get("action_target")
8616 .and_then(|v| v.as_str())
8617 .unwrap_or("");
8618 let args = step_value
8619 .get("args")
8620 .cloned()
8621 .unwrap_or(serde_json::json!({}));
8622 let deps: Vec<String> = step_value
8623 .get("dependencies")
8624 .and_then(|v| v.as_array())
8625 .map(|arr| {
8626 arr.iter()
8627 .filter_map(|v| v.as_str().map(String::from))
8628 .collect()
8629 })
8630 .unwrap_or_default();
8631
8632 let action = match action_type {
8633 "tool" => PlanAction::tool(action_target, args),
8634 "skill" => PlanAction::skill(action_target),
8635 "respond" => PlanAction::respond(action_target),
8636 _ => PlanAction::think(desc),
8637 };
8638
8639 let step = PlanStep::new(desc, action)
8640 .with_id(id)
8641 .with_dependencies(deps);
8642 plan.add_step(step);
8643 }
8644 }
8645 }
8646
8647 if plan.steps.is_empty() {
8648 plan.add_step(PlanStep::new(
8649 "Process the request",
8650 PlanAction::think(input),
8651 ));
8652 plan.add_step(PlanStep::new(
8653 "Provide response",
8654 PlanAction::respond("Answer based on analysis"),
8655 ));
8656 }
8657
8658 Ok(plan)
8659 }
8660
8661 async fn execute_plan(&self, plan: &mut Plan) -> Result<String> {
8662 let llm = self.get_state_llm()?;
8663 let mut results: HashMap<String, serde_json::Value> = HashMap::new();
8664 let effective = self.get_effective_reasoning_config();
8665 let max_steps = effective.get_planning().map(|c| c.max_steps).unwrap_or(10);
8666
8667 plan.status = PlanStatus::InProgress;
8668
8669 for step_idx in 0..plan.steps.len().min(max_steps as usize) {
8670 let step = &plan.steps[step_idx];
8671
8672 let deps_satisfied = step.dependencies.iter().all(|dep| {
8673 plan.steps
8674 .iter()
8675 .find(|s| &s.id == dep)
8676 .map(|s| s.status.is_completed())
8677 .unwrap_or(false)
8678 });
8679
8680 if !deps_satisfied {
8681 continue;
8682 }
8683
8684 plan.steps[step_idx].mark_running();
8685
8686 let result = match &plan.steps[step_idx].action {
8687 PlanAction::Tool { tool, args } => {
8688 let has_dep_results = plan.steps[step_idx]
8694 .dependencies
8695 .iter()
8696 .any(|dep| results.contains_key(dep));
8697
8698 let final_args = if has_dep_results {
8699 let dep_context: String = plan.steps[step_idx]
8700 .dependencies
8701 .iter()
8702 .filter_map(|dep| results.get(dep).map(|r| format!("{}: {}", dep, r)))
8703 .collect::<Vec<_>>()
8704 .join("\n");
8705
8706 let tool_schema = self
8707 .tools
8708 .get(tool)
8709 .map(|t| {
8710 let schema = t.input_schema();
8711 let props = schema
8712 .get("properties")
8713 .and_then(|p| serde_json::to_string(p).ok())
8714 .unwrap_or_else(|| "{}".to_string());
8715 format!(
8716 "{}: {}\nArguments schema: {}",
8717 t.id(),
8718 t.description(),
8719 props
8720 )
8721 })
8722 .unwrap_or_default();
8723
8724 let step_desc = &plan.steps[step_idx].description;
8725 let arg_prompt = format!(
8726 "Generate the JSON arguments for a tool call.\n\n\
8727 Tool: {}\n\n\
8728 Task: {}\n\n\
8729 Previous step results:\n{}\n\n\
8730 Planner's draft arguments: {}\n\n\
8731 Produce ONLY a valid JSON object with the correct argument values.\n\
8732 Use actual values from the previous step results, not template references.",
8733 tool_schema,
8734 step_desc,
8735 dep_context,
8736 serde_json::to_string(args).unwrap_or_default()
8737 );
8738 let messages = vec![ChatMessage::user(&arg_prompt)];
8739 match self
8740 .observe_purpose(
8741 ObservationPurpose::PlanStep,
8742 llm.complete(&messages, None),
8743 )
8744 .await
8745 {
8746 Ok(resp) => {
8747 let content = resp.content.trim();
8748 let json_start = content.find('{');
8750 let json_end = content.rfind('}');
8751 if let (Some(start), Some(end)) = (json_start, json_end) {
8752 serde_json::from_str(&content[start..=end])
8753 .unwrap_or_else(|_| args.clone())
8754 } else {
8755 args.clone()
8756 }
8757 }
8758 Err(_) => args.clone(),
8759 }
8760 } else {
8761 args.clone()
8762 };
8763
8764 let request = ToolExecutionRequest::new(
8765 uuid::Uuid::new_v4().to_string(),
8766 tool.clone(),
8767 final_args,
8768 ToolCallSource::Plan {
8769 step_index: step_idx,
8770 },
8771 );
8772 match self.execute_tool_record(request).await {
8773 Ok(record) if record.success => {
8774 serde_json::json!({ "output": record.model_output_string() })
8775 }
8776 Ok(record) => {
8777 plan.steps[step_idx].mark_failed(record.model_output_string());
8778 continue;
8779 }
8780 Err(e) => {
8781 plan.steps[step_idx].mark_failed(e.to_string());
8782 continue;
8783 }
8784 }
8785 }
8786 PlanAction::Skill { skill } => {
8787 if let Some(skill_def) = self.skills.iter().find(|s| &s.id == skill) {
8788 if let Some(ref executor) = self.skill_executor {
8789 match executor
8790 .execute_with_invoker(skill_def, "", serde_json::json!({}), self)
8791 .await
8792 {
8793 Ok(output) => serde_json::json!({ "output": output }),
8794 Err(e) => {
8795 plan.steps[step_idx].mark_failed(e.to_string());
8796 continue;
8797 }
8798 }
8799 } else {
8800 serde_json::json!({ "output": "Skill executor not available" })
8801 }
8802 } else {
8803 plan.steps[step_idx].mark_failed("Skill not found");
8804 continue;
8805 }
8806 }
8807 PlanAction::Think { prompt } => {
8808 let context: String = results
8809 .iter()
8810 .map(|(k, v)| format!("{}: {}", k, v))
8811 .collect::<Vec<_>>()
8812 .join("\n");
8813
8814 let think_prompt = format!("Context:\n{}\n\nTask: {}", context, prompt);
8815 let messages = vec![ChatMessage::user(&think_prompt)];
8816
8817 match self
8818 .observe_purpose(
8819 ObservationPurpose::PlanStep,
8820 llm.complete(&messages, None),
8821 )
8822 .await
8823 {
8824 Ok(resp) => serde_json::json!({ "output": resp.content }),
8825 Err(e) => {
8826 plan.steps[step_idx].mark_failed(e.to_string());
8827 continue;
8828 }
8829 }
8830 }
8831 PlanAction::Respond { template } => {
8832 let context: String = results
8833 .iter()
8834 .map(|(k, v)| format!("{}: {}", k, v))
8835 .collect::<Vec<_>>()
8836 .join("\n");
8837
8838 let respond_prompt = format!(
8839 "Based on this context:\n{}\n\nGenerate a response following this template/instruction: {}",
8840 context, template
8841 );
8842 let messages = vec![ChatMessage::user(&respond_prompt)];
8843
8844 match self
8845 .observe_purpose(
8846 ObservationPurpose::PlanStep,
8847 llm.complete(&messages, None),
8848 )
8849 .await
8850 {
8851 Ok(resp) => serde_json::json!({ "output": resp.content }),
8852 Err(e) => {
8853 plan.steps[step_idx].mark_failed(e.to_string());
8854 continue;
8855 }
8856 }
8857 }
8858 };
8859
8860 results.insert(plan.steps[step_idx].id.clone(), result.clone());
8861 plan.steps[step_idx].mark_completed(Some(result));
8862 }
8863
8864 let has_failures = plan.steps.iter().any(|s| s.status.is_failed());
8866 if has_failures {
8867 let failed_ids: Vec<String> = plan
8868 .steps
8869 .iter()
8870 .filter(|s| s.status.is_failed())
8871 .map(|s| s.id.clone())
8872 .collect();
8873 plan.status = PlanStatus::Failed {
8874 error: format!("Steps failed: {}", failed_ids.join(", ")),
8875 };
8876 } else {
8877 plan.status = PlanStatus::Completed;
8878 }
8879
8880 let all_outputs: Vec<String> = plan
8882 .steps
8883 .iter()
8884 .filter(|s| s.status.is_completed())
8885 .filter_map(|s| {
8886 s.result
8887 .as_ref()
8888 .and_then(|r| r.get("output"))
8889 .and_then(|o| o.as_str())
8890 .map(|o| format!("{}: {}", s.description, o))
8891 })
8892 .collect();
8893
8894 if all_outputs.is_empty() {
8895 return Ok("Plan execution completed but produced no results.".to_string());
8896 }
8897
8898 if all_outputs.len() == 1 {
8899 return Ok(all_outputs.into_iter().next().unwrap());
8900 }
8901
8902 let context = all_outputs.join("\n\n");
8904 let prompt = format!(
8905 "You completed a multi-step plan for: \"{}\"\n\nStep results:\n{}\n\nProvide a coherent final response that synthesizes these results.",
8906 plan.goal, context
8907 );
8908 let messages = vec![ChatMessage::user(&prompt)];
8909 match self
8910 .observe_purpose(ObservationPurpose::PlanStep, llm.complete(&messages, None))
8911 .await
8912 {
8913 Ok(resp) => Ok(resp.content.trim().to_string()),
8914 Err(_) => Ok(context),
8915 }
8916 }
8917
8918 async fn evaluate_response(&self, input: &str, response: &str) -> Result<EvaluationResult> {
8919 let effective_config = self.get_effective_reflection_config();
8920 self.evaluate_response_with_config(input, response, &effective_config)
8921 .await
8922 }
8923
8924 fn extract_thinking(&self, content: &str) -> (Option<String>, String) {
8925 if let Some(start) = content.find("<thinking>")
8926 && let Some(end) = content.find("</thinking>")
8927 {
8928 let thinking = content[start + 10..end].trim().to_string();
8929 let answer = content[end + 11..].trim().to_string();
8930 return (Some(thinking), answer);
8931 }
8932 (None, content.to_string())
8933 }
8934
8935 fn format_response_with_thinking(&self, thinking: Option<&str>, answer: &str) -> String {
8936 match self.get_effective_reasoning_config().output {
8937 ReasoningOutput::Hidden => answer.to_string(),
8938 ReasoningOutput::Visible => {
8939 if let Some(t) = thinking {
8940 format!("Thinking:\n{}\n\nAnswer:\n{}", t, answer)
8941 } else {
8942 answer.to_string()
8943 }
8944 }
8945 ReasoningOutput::Tagged => {
8946 if let Some(t) = thinking {
8947 format!("<thinking>{}</thinking>\n{}", t, answer)
8948 } else {
8949 answer.to_string()
8950 }
8951 }
8952 }
8953 }
8954
8955 async fn run_loop(&self, input: &str) -> Result<AgentResponse> {
8960 self.init_storage().await?;
8964 self.begin_root_turn();
8965 let _root_cleanup = RootTurnCleanup::new(self);
8966 info!(input_len = input.len(), "Starting chat");
8967
8968 self.hooks.on_message_received(input).await;
8969
8970 if !self.context_initialized.swap(true, Ordering::SeqCst) {
8974 self.context_manager.initialize().await?;
8975 debug!("Context manager initialized (defaults, env, builtins)");
8976 }
8977
8978 self.check_turn_timeout().await?;
8979 self.context_manager.refresh_per_turn().await?;
8980
8981 self.clear_disambiguation_context();
8984
8985 if let Some(ref disambiguator) = self.disambiguation_manager {
8987 let disambiguation_context = self.build_disambiguation_context().await?;
8988
8989 let state_override = self
8991 .state_machine
8992 .as_ref()
8993 .and_then(|sm| sm.current_definition())
8994 .and_then(|def| def.disambiguation.clone());
8995
8996 let state_generation = self
8997 .state_machine
8998 .as_ref()
8999 .map(|state_machine| state_machine.generation());
9000 let disambiguation_epoch = self.disambiguation_epoch.load(Ordering::SeqCst);
9001 let mut disambiguation_result = self
9002 .observe_purpose(
9003 ObservationPurpose::DisambiguationDetection,
9004 disambiguator.process_input_with_override(
9005 input,
9006 &disambiguation_context,
9007 state_override.as_ref(),
9008 None,
9009 ),
9010 )
9011 .await?;
9012 let current_state_generation = self
9013 .state_machine
9014 .as_ref()
9015 .map(|state_machine| state_machine.generation());
9016 if current_state_generation != state_generation
9017 || self.disambiguation_epoch.load(Ordering::SeqCst) != disambiguation_epoch
9018 {
9019 disambiguator.clear_pending().await;
9020 *self.pending_skill_id.write() = None;
9021 disambiguation_result = DisambiguationResult::Abandoned { new_input: None };
9022 info!(
9023 confirmation_event = "invalidated",
9024 invalidation_reason = "state_generation_changed",
9025 "Disambiguation result invalidated before redispatch"
9026 );
9027 }
9028 match disambiguation_result {
9029 DisambiguationResult::Clear => {
9030 debug!("Input is clear, proceeding normally");
9031 }
9032 DisambiguationResult::NeedsClarification {
9033 question,
9034 detection,
9035 } => {
9036 let admission = self
9037 .admit_disambiguation_redispatch(disambiguation_epoch, state_generation)
9038 .await?;
9039 let awaiting_confirmation = disambiguator.has_pending_confirmation().await;
9040 info!(
9041 ambiguity_type = ?detection.ambiguity_type,
9042 confidence = detection.confidence,
9043 "Input requires clarification"
9044 );
9045
9046 self.commit_root_user_message(input).await?;
9049 self.memory
9050 .add_message(ChatMessage::assistant(&question.question))
9051 .await?;
9052
9053 let status = if awaiting_confirmation {
9054 "awaiting_confirmation"
9055 } else {
9056 "awaiting_clarification"
9057 };
9058 let response = AgentResponse::new(&question.question).with_metadata(
9059 "disambiguation",
9060 serde_json::json!({
9061 "status": status,
9062 "options": question.options,
9063 "clarifying": question.clarifying,
9064 "detection": {
9065 "type": detection.ambiguity_type,
9066 "confidence": detection.confidence,
9067 "what_is_unclear": detection.what_is_unclear,
9068 }
9069 }),
9070 );
9071 drop(admission);
9072 self.finish_turn_if_root(&response).await?;
9073 return Ok(response);
9074 }
9075 DisambiguationResult::Clarified {
9076 enriched_input,
9077 resolved,
9078 ..
9079 } => {
9080 let admission = match self
9081 .admit_disambiguation_redispatch(disambiguation_epoch, state_generation)
9082 .await
9083 {
9084 Ok(admission) => admission,
9085 Err(error) => {
9086 *self.pending_skill_id.write() = None;
9087 return Err(error);
9088 }
9089 };
9090 info!(
9091 resolved_count = resolved.len(),
9092 enriched = %enriched_input,
9093 "Input clarified, injecting resolved intent into context"
9094 );
9095
9096 for (key, value) in &resolved {
9099 let context_key = format!("disambiguation.{}", key);
9100 let _ = self.context_manager.set(&context_key, value.clone());
9101 }
9102
9103 if let Some(intent) = resolved.get("intent") {
9104 let _ = self.context_manager.set("resolved_intent", intent.clone());
9105 }
9106
9107 let _ = self
9108 .context_manager
9109 .set("disambiguation.resolved", serde_json::Value::Bool(true));
9110
9111 let skill_id = self.pending_skill_id.read().clone();
9115 if let Some(skill_id) = skill_id {
9116 info!(skill_id = %skill_id, "Re-checking skill disambiguation on clarified input");
9117 drop(admission);
9118 return self
9119 .recheck_skill_disambiguation(
9120 &skill_id,
9121 &enriched_input,
9122 disambiguation_epoch,
9123 state_generation,
9124 )
9125 .await;
9126 }
9127
9128 drop(admission);
9129 return self.run_loop_internal(&enriched_input).await;
9130 }
9131 DisambiguationResult::ProceedWithBestGuess { enriched_input } => {
9132 info!("Proceeding with best guess interpretation");
9133
9134 let skill_id = self.pending_skill_id.read().clone();
9136 if let Some(skill_id) = skill_id {
9137 info!(skill_id = %skill_id, "Re-checking skill disambiguation on best-guess input");
9138 return self
9139 .recheck_skill_disambiguation(
9140 &skill_id,
9141 &enriched_input,
9142 disambiguation_epoch,
9143 state_generation,
9144 )
9145 .await;
9146 }
9147
9148 return self.run_loop_internal(&enriched_input).await;
9149 }
9150 DisambiguationResult::GiveUp { reason } => {
9151 *self.pending_skill_id.write() = None;
9152 warn!(reason = %reason, "Disambiguation gave up");
9153 let apology = self
9154 .generate_localized_apology(
9155 "Generate a brief, polite apology saying you couldn't understand the request. Be concise.",
9156 &reason,
9157 )
9158 .await
9159 .unwrap_or_else(|_| {
9160 format!("I'm sorry, I couldn't understand your request: {}", reason)
9161 });
9162 let response = AgentResponse::new(&apology);
9163 self.finish_turn_if_root(&response).await?;
9164 return Ok(response);
9165 }
9166 DisambiguationResult::Escalate { reason } => {
9167 *self.pending_skill_id.write() = None;
9168 info!(reason = %reason, "Escalating to human");
9169 if let Some(ref hitl) = self.hitl_engine {
9170 let trigger =
9171 ApprovalTrigger::condition("disambiguation_escalation", reason.clone());
9172 let mut context_map = HashMap::new();
9173 context_map.insert("original_input".to_string(), serde_json::json!(input));
9174 context_map.insert("reason".to_string(), serde_json::json!(&reason));
9175 let check_result = HITLCheckResult::required(
9176 trigger,
9177 context_map,
9178 format!("User request needs human assistance: {}", reason),
9179 Some(hitl.config().default_timeout_seconds),
9180 );
9181 let result = self.request_hitl_approval(check_result).await?;
9182 if matches!(
9183 result,
9184 ApprovalResult::Approved | ApprovalResult::Modified { .. }
9185 ) {
9186 return self.run_loop_internal(input).await;
9187 }
9188 }
9189 let apology = self
9190 .generate_localized_apology(
9191 "Explain briefly that you're transferring the user to a human agent for help.",
9192 &reason,
9193 )
9194 .await
9195 .unwrap_or_else(|_| {
9196 format!("I need human assistance to help with your request: {}", reason)
9197 });
9198 let response = AgentResponse::new(&apology);
9199 self.finish_turn_if_root(&response).await?;
9200 return Ok(response);
9201 }
9202 DisambiguationResult::Abandoned { new_input } => {
9203 *self.pending_skill_id.write() = None;
9204
9205 info!(
9206 has_new_input = new_input.is_some(),
9207 "Clarification abandoned by user"
9208 );
9209
9210 self.commit_root_user_message(input).await?;
9211
9212 match new_input {
9213 Some(fresh_input) => {
9214 return self.run_loop_internal(&fresh_input).await;
9217 }
9218 None => {
9219 let ack = self
9221 .generate_localized_apology(
9222 "The user changed their mind about their previous request. \
9223 Generate a brief, friendly acknowledgment (e.g. 'OK, no problem. What else can I help with?'). \
9224 Do NOT apologize excessively. Be concise.",
9225 "User abandoned clarification",
9226 )
9227 .await
9228 .unwrap_or_else(|_| {
9229 "OK, no problem. What else can I help with?".to_string()
9230 });
9231
9232 self.memory
9233 .add_message(ChatMessage::assistant(&ack))
9234 .await?;
9235
9236 let response = AgentResponse::new(&ack);
9237 self.finish_turn_if_root(&response).await?;
9238 return Ok(response);
9239 }
9240 }
9241 }
9242 }
9243 }
9244
9245 self.run_loop_internal(input).await
9246 }
9247
9248 async fn generate_localized_apology(&self, instruction: &str, reason: &str) -> Result<String> {
9250 let llm = self.llm_registry.router().map_err(|e| {
9251 AgentError::LLM(format!(
9252 "Router LLM not available for localized response: {}",
9253 e
9254 ))
9255 })?;
9256
9257 let recent: Vec<String> = self
9258 .memory
9259 .get_messages(Some(3))
9260 .await?
9261 .iter()
9262 .map(|m| m.content.clone())
9263 .collect();
9264
9265 let context_hint = if recent.is_empty() {
9266 String::new()
9267 } else {
9268 format!(
9269 "\nRecent conversation (detect the user's language from this):\n{}\n",
9270 recent.join("\n")
9271 )
9272 };
9273
9274 let prompt = format!(
9275 "{}\nReason: {}\n{}Respond in the same language as the user. Output ONLY the message, nothing else.",
9276 instruction, reason, context_hint
9277 );
9278
9279 let messages = vec![ChatMessage::user(&prompt)];
9280 let response = self
9281 .observe_purpose(
9282 ObservationPurpose::DisambiguationClarification,
9283 llm.complete(&messages, None),
9284 )
9285 .await
9286 .map_err(|e| AgentError::LLM(format!("Localized response generation failed: {}", e)))?;
9287
9288 Ok(response.content.trim().to_string())
9289 }
9290
9291 fn render_action_args(&self, args: &Value) -> Value {
9295 let context = self.build_context_with_overlays();
9296 match args {
9297 Value::Object(map) => {
9298 let mut rendered = serde_json::Map::new();
9299 for (k, v) in map {
9300 match v {
9301 Value::String(s) if s.contains("{{") => {
9302 match self.template_renderer.render(s, &context) {
9303 Ok(rendered_str) => {
9304 rendered.insert(k.clone(), Value::String(rendered_str));
9305 }
9306 Err(_) => {
9307 rendered.insert(k.clone(), v.clone());
9308 }
9309 }
9310 }
9311 _ => {
9312 rendered.insert(k.clone(), v.clone());
9313 }
9314 }
9315 }
9316 Value::Object(rendered)
9317 }
9318 _ => args.clone(),
9319 }
9320 }
9321
9322 fn clear_disambiguation_context(&self) {
9324 let _ = self
9325 .context_manager
9326 .set("resolved_intent", serde_json::Value::Null);
9327
9328 let all = self.context_manager.get_all();
9329 for key in all.keys() {
9330 if key.starts_with("disambiguation.") {
9331 let _ = self.context_manager.set(key, serde_json::Value::Null);
9332 }
9333 }
9334 }
9335
9336 async fn recheck_skill_disambiguation(
9342 &self,
9343 skill_id: &str,
9344 enriched_input: &str,
9345 expected_disambiguation_epoch: u64,
9346 expected_state_generation: Option<u64>,
9347 ) -> Result<AgentResponse> {
9348 let skill = self
9349 .skill_router
9350 .as_ref()
9351 .and_then(|r| r.get_skill(skill_id).cloned());
9352
9353 if let Some(ref skill) = skill
9355 && let Some(ref skill_disambig) = skill.disambiguation
9356 && skill_disambig.enabled.unwrap_or(false)
9357 && let Some(ref disambiguator) = self.disambiguation_manager
9358 {
9359 let context = self.build_disambiguation_context().await?;
9360 let state_override = self
9361 .state_machine
9362 .as_ref()
9363 .and_then(|sm| sm.current_definition())
9364 .and_then(|def| def.disambiguation.clone());
9365
9366 let disambiguation_result = self
9367 .observe_purpose(
9368 ObservationPurpose::DisambiguationDetection,
9369 disambiguator.process_input_with_override(
9370 enriched_input,
9371 &context,
9372 state_override.as_ref(),
9373 Some(skill_disambig),
9374 ),
9375 )
9376 .await?;
9377 let current_state_generation = self
9378 .state_machine
9379 .as_ref()
9380 .map(|state_machine| state_machine.generation());
9381 if current_state_generation != expected_state_generation
9382 || self.disambiguation_epoch.load(Ordering::SeqCst) != expected_disambiguation_epoch
9383 {
9384 disambiguator.clear_pending().await;
9385 *self.pending_skill_id.write() = None;
9386 return Err(AgentError::Other(
9387 "State or reset ownership changed during skill disambiguation recheck"
9388 .to_string(),
9389 ));
9390 }
9391 match disambiguation_result {
9392 DisambiguationResult::Clear => {
9393 debug!(skill_id = %skill_id, "Skill re-check: all fields present");
9394 }
9395 DisambiguationResult::NeedsClarification {
9396 question,
9397 detection,
9398 } => {
9399 let admission = self
9400 .admit_disambiguation_redispatch(
9401 expected_disambiguation_epoch,
9402 expected_state_generation,
9403 )
9404 .await?;
9405 let awaiting_confirmation = disambiguator.has_pending_confirmation().await;
9406 info!(
9407 skill_id = %skill_id,
9408 ambiguity_type = ?detection.ambiguity_type,
9409 what_is_unclear = ?detection.what_is_unclear,
9410 "Skill re-check: still missing fields, asking again"
9411 );
9412 self.memory
9416 .add_message(ChatMessage::user(enriched_input))
9417 .await?;
9418 self.memory
9419 .add_message(ChatMessage::assistant(&question.question))
9420 .await?;
9421
9422 let response = AgentResponse::new(&question.question).with_metadata(
9423 "disambiguation",
9424 serde_json::json!({
9425 "status": if awaiting_confirmation { "awaiting_confirmation" } else { "awaiting_clarification" },
9426 "skill_id": skill_id,
9427 "options": question.options,
9428 "clarifying": question.clarifying,
9429 "detection": {
9430 "type": detection.ambiguity_type,
9431 "confidence": detection.confidence,
9432 "what_is_unclear": detection.what_is_unclear,
9433 }
9434 }),
9435 );
9436 drop(admission);
9437 self.finish_turn_if_root(&response).await?;
9438 return Ok(response);
9439 }
9440 DisambiguationResult::Clarified {
9441 enriched_input: re_enriched,
9442 ..
9443 } => {
9444 debug!(skill_id = %skill_id, "Skill re-check: clarified immediately, executing");
9445 let admission = self
9446 .admit_disambiguation_redispatch(
9447 expected_disambiguation_epoch,
9448 expected_state_generation,
9449 )
9450 .await?;
9451 *self.pending_skill_id.write() = None;
9452 drop(admission);
9453 let skill_response = self.execute_skill_by_id(skill_id, &re_enriched).await?;
9454 self.memory
9455 .add_message(ChatMessage::user(&re_enriched))
9456 .await?;
9457 return self
9458 .handle_skill_response(
9459 &re_enriched,
9460 skill_id,
9461 skill_response,
9462 &HashMap::new(),
9463 )
9464 .await;
9465 }
9466 DisambiguationResult::ProceedWithBestGuess {
9467 enriched_input: re_enriched,
9468 } => {
9469 debug!(skill_id = %skill_id, "Skill re-check: proceeding with best guess");
9470 let admission = self
9471 .admit_disambiguation_redispatch(
9472 expected_disambiguation_epoch,
9473 expected_state_generation,
9474 )
9475 .await?;
9476 *self.pending_skill_id.write() = None;
9477 drop(admission);
9478 let skill_response = self.execute_skill_by_id(skill_id, &re_enriched).await?;
9479 self.memory
9480 .add_message(ChatMessage::user(&re_enriched))
9481 .await?;
9482 return self
9483 .handle_skill_response(
9484 &re_enriched,
9485 skill_id,
9486 skill_response,
9487 &HashMap::new(),
9488 )
9489 .await;
9490 }
9491 DisambiguationResult::GiveUp { reason } => {
9492 *self.pending_skill_id.write() = None;
9493 let apology = self
9494 .generate_localized_apology(
9495 "Generate a brief, polite apology saying you couldn't understand the request. Be concise.",
9496 &reason,
9497 )
9498 .await
9499 .unwrap_or_else(|_| {
9500 format!("I'm sorry, I couldn't understand your request: {}", reason)
9501 });
9502 let response = AgentResponse::new(&apology);
9503 self.finish_turn_if_root(&response).await?;
9504 return Ok(response);
9505 }
9506 DisambiguationResult::Escalate { reason } => {
9507 *self.pending_skill_id.write() = None;
9508 let apology = self
9509 .generate_localized_apology(
9510 "Explain briefly that you're transferring the user to a human agent for help.",
9511 &reason,
9512 )
9513 .await
9514 .unwrap_or_else(|_| {
9515 format!("I need human assistance to help with your request: {}", reason)
9516 });
9517 let response = AgentResponse::new(&apology);
9518 self.finish_turn_if_root(&response).await?;
9519 return Ok(response);
9520 }
9521 DisambiguationResult::Abandoned { new_input } => {
9522 *self.pending_skill_id.write() = None;
9525 debug!(skill_id = %skill_id, "Skill re-check: abandoned by user");
9526 if let Some(fresh) = new_input {
9527 return self.run_loop_internal(&fresh).await;
9528 }
9529 let ack = self
9530 .generate_localized_apology(
9531 "The user changed their mind about their previous request. \
9532 Generate a brief, friendly acknowledgment (e.g. 'OK, no problem. What else can I help with?'). \
9533 Do NOT apologize excessively. Be concise.",
9534 "User abandoned clarification",
9535 )
9536 .await
9537 .unwrap_or_else(|_| {
9538 "OK, no problem. What else can I help with?".to_string()
9539 });
9540 self.memory
9541 .add_message(ChatMessage::assistant(&ack))
9542 .await?;
9543 let response = AgentResponse::new(&ack);
9544 self.finish_turn_if_root(&response).await?;
9545 return Ok(response);
9546 }
9547 }
9548 }
9549
9550 let admission = self
9552 .admit_disambiguation_redispatch(
9553 expected_disambiguation_epoch,
9554 expected_state_generation,
9555 )
9556 .await?;
9557 *self.pending_skill_id.write() = None;
9558 drop(admission);
9559 let skill_response = self.execute_skill_by_id(skill_id, enriched_input).await?;
9560 self.memory
9561 .add_message(ChatMessage::user(enriched_input))
9562 .await?;
9563 self.handle_skill_response(enriched_input, skill_id, skill_response, &HashMap::new())
9564 .await
9565 }
9566
9567 async fn handle_skill_response(
9570 &self,
9571 processed_input: &str,
9572 skill_id: &str,
9573 skill_response: String,
9574 input_context: &HashMap<String, Value>,
9575 ) -> Result<AgentResponse> {
9576 let output_data = self.process_output(&skill_response, input_context).await?;
9577 let final_response = output_data.content;
9578
9579 self.memory
9580 .add_message(ChatMessage::assistant(&final_response))
9581 .await?;
9582
9583 self.check_memory_compression().await?;
9584
9585 self.increment_turn();
9586 self.evaluate_transitions(processed_input, &final_response)
9587 .await?;
9588
9589 let response = AgentResponse::new(final_response)
9590 .with_metadata("skill_id", serde_json::json!(skill_id));
9591 self.finish_turn_if_root(&response).await?;
9592 Ok(response)
9593 }
9594
9595 async fn handle_plan_and_execute(
9598 &self,
9599 processed_input: &str,
9600 input_context: &HashMap<String, Value>,
9601 auto_detected: bool,
9602 ) -> Result<AgentResponse> {
9603 let effective = self.get_effective_reasoning_config();
9604 let plan_reflection = effective
9605 .get_planning()
9606 .map(|c| c.reflection.clone())
9607 .unwrap_or_default();
9608
9609 let max_attempts = if plan_reflection.enabled {
9610 1 + plan_reflection.max_replans
9611 } else {
9612 1
9613 };
9614
9615 let mut plan = self.generate_plan(processed_input).await?;
9616 info!(
9617 plan_id = %plan.id,
9618 steps = plan.steps.len(),
9619 "Plan generated"
9620 );
9621
9622 let mut plan_result = String::new();
9623
9624 for attempt in 0..max_attempts {
9625 *self.current_plan.write() = Some(plan.clone());
9626 plan_result = self.execute_plan(&mut plan).await?;
9627
9628 info!(
9629 plan_status = ?plan.status,
9630 completed_steps = plan.completed_steps().count(),
9631 attempt = attempt + 1,
9632 "Plan execution completed"
9633 );
9634
9635 if !plan_reflection.enabled {
9636 break;
9637 }
9638
9639 let has_failures = plan.steps.iter().any(|s| s.status.is_failed());
9640 if !has_failures {
9641 break;
9642 }
9643
9644 if attempt + 1 >= max_attempts {
9645 break;
9646 }
9647
9648 match plan_reflection.on_step_failure {
9649 StepFailureAction::Replan => {
9650 info!(attempt = attempt + 1, "Plan had failures, replanning");
9651 plan = self.generate_plan(processed_input).await?;
9652 }
9653 StepFailureAction::Abort => {
9654 warn!("Plan step failed, aborting");
9655 break;
9656 }
9657 StepFailureAction::Skip | StepFailureAction::Continue => {
9658 break;
9659 }
9660 }
9661 }
9662
9663 *self.current_plan.write() = Some(plan);
9664
9665 let output_data = self.process_output(&plan_result, input_context).await?;
9666 let final_content = output_data.content;
9667
9668 self.memory
9669 .add_message(ChatMessage::assistant(&final_content))
9670 .await?;
9671
9672 self.check_memory_compression().await?;
9673 self.increment_turn();
9674 self.evaluate_transitions(processed_input, &final_content)
9675 .await?;
9676
9677 let reasoning_metadata =
9678 ReasoningMetadata::new(ReasoningMode::PlanAndExecute).with_auto_detected(auto_detected);
9679
9680 let response = AgentResponse::new(&final_content).with_metadata(
9681 "reasoning",
9682 serde_json::to_value(&reasoning_metadata).unwrap_or_default(),
9683 );
9684
9685 self.finish_turn_if_root(&response).await?;
9686 Ok(response)
9687 }
9688
9689 fn inject_reasoning_prompt(
9691 &self,
9692 messages: &mut [ChatMessage],
9693 reasoning_mode: &ReasoningMode,
9694 is_first_iteration: bool,
9695 ) {
9696 if !is_first_iteration {
9697 return;
9698 }
9699 match reasoning_mode {
9700 ReasoningMode::CoT => {
9701 if let Some(msg) = messages.first_mut()
9702 && matches!(msg.role, ai_agents_core::Role::System)
9703 {
9704 msg.content = self.build_cot_system_prompt(&msg.content);
9705 debug!("Applied Chain-of-Thought system prompt");
9706 }
9707 }
9708 ReasoningMode::React => {
9709 if let Some(msg) = messages.first_mut()
9710 && matches!(msg.role, ai_agents_core::Role::System)
9711 {
9712 msg.content = self.build_react_system_prompt(&msg.content);
9713 debug!("Applied ReAct system prompt");
9714 }
9715 }
9716 _ => {}
9717 }
9718 }
9719
9720 async fn generate_main_response_draft(
9725 &self,
9726 processed_input: &str,
9727 reasoning_mode: &ReasoningMode,
9728 ) -> Result<MainResponseDraft> {
9729 let llm = self.get_state_llm()?;
9730 let protocol = self.main_tool_protocol(llm.as_ref(), true).await?;
9731 let mut messages = self
9732 .build_messages_internal(false, Some(processed_input), protocol.choice.is_none())
9733 .await?;
9734 self.inject_reasoning_prompt(&mut messages, reasoning_mode, true);
9735 let response = self
9736 .complete_main_llm_with_recovery(llm, &messages, &protocol)
9737 .await?;
9738 let content = response.content.trim().to_string();
9739 let (thinking, answer) = self.extract_thinking(&content);
9740 if let Some(calls) = self.parse_main_tool_calls(&content, &protocol)? {
9741 return Ok(MainResponseDraft::ToolCalls {
9742 raw_content: content,
9743 calls,
9744 thinking,
9745 });
9746 }
9747 Ok(MainResponseDraft::Text {
9748 raw_content: answer,
9749 thinking,
9750 })
9751 }
9752
9753 async fn commit_main_response_draft(
9758 &self,
9759 processed_input: &str,
9760 input_context: &HashMap<String, Value>,
9761 draft: MainResponseDraft,
9762 reasoning_mode: ReasoningMode,
9763 auto_detected: bool,
9764 ) -> Result<AgentResponse> {
9765 self.commit_root_user_message(processed_input).await?;
9766 match draft {
9767 MainResponseDraft::Text {
9768 raw_content,
9769 thinking,
9770 } => {
9771 self.finish_text_response_from_model(CommittedTextResponse {
9772 processed_input,
9773 input_context,
9774 answer: raw_content,
9775 reasoning_mode,
9776 auto_detected,
9777 iterations: 1,
9778 thinking_content: thinking,
9779 all_tool_calls: Vec::new(),
9780 })
9781 .await
9782 }
9783 MainResponseDraft::ToolCalls {
9784 raw_content,
9785 calls,
9786 thinking: _,
9787 } => {
9788 let mut all_tool_calls = Vec::new();
9789 match self
9790 .handle_tool_calls(processed_input, &raw_content, calls, &mut all_tool_calls)
9791 .await?
9792 {
9793 ToolCallOutcome::Rejected(response) => {
9794 self.finish_turn_if_root(&response).await?;
9795 Ok(response)
9796 }
9797 ToolCallOutcome::Continue | ToolCallOutcome::TransitionFired => {
9798 self.continue_after_committed_tool_draft(processed_input)
9799 .await
9800 }
9801 }
9802 }
9803 }
9804 }
9805
9806 async fn continue_after_committed_tool_draft(
9811 &self,
9812 processed_input: &str,
9813 ) -> Result<AgentResponse> {
9814 *self.redispatch_depth.write() += 1;
9815 if let Some(context) = self.active_turn_context.write().as_mut() {
9816 context.enter_redispatch();
9817 }
9818 let result = Box::pin(self.run_loop_internal(processed_input)).await;
9819 *self.redispatch_depth.write() -= 1;
9820 if let Some(context) = self.active_turn_context.write().as_mut() {
9821 context.exit_redispatch();
9822 }
9823 let response = result?;
9824 self.finish_turn_if_root(&response).await?;
9825 Ok(response)
9826 }
9827
9828 async fn finish_text_response_from_model(
9833 &self,
9834 response: CommittedTextResponse<'_>,
9835 ) -> Result<AgentResponse> {
9836 let CommittedTextResponse {
9837 processed_input,
9838 input_context,
9839 answer,
9840 reasoning_mode,
9841 auto_detected,
9842 iterations,
9843 thinking_content,
9844 all_tool_calls,
9845 } = response;
9846 let output_data = self.process_output(&answer, input_context).await?;
9847 let mut final_content = if output_data.metadata.rejected {
9848 output_data
9849 .metadata
9850 .rejection_reason
9851 .unwrap_or_else(|| answer.to_string())
9852 } else {
9853 output_data.content
9854 };
9855 let llm = self.get_state_llm()?;
9856 let reflection_metadata;
9857 (final_content, reflection_metadata) = self
9858 .run_reflection(&*llm, processed_input, final_content)
9859 .await?;
9860 final_content =
9861 self.format_response_with_thinking(thinking_content.as_deref(), &final_content);
9862 let final_content = {
9863 let result = self
9864 .post_loop_processing(processed_input, final_content)
9865 .await?;
9866 self.apply_post_loop_result(processed_input, result).await?
9867 };
9868 let response = self.build_agent_response(AgentResponseParts {
9869 content: final_content,
9870 all_tool_calls,
9871 reasoning_mode,
9872 auto_detected,
9873 iterations,
9874 thinking: thinking_content,
9875 reflection_metadata,
9876 });
9877 self.finish_turn_if_root(&response).await?;
9878 Ok(response)
9879 }
9880
9881 async fn run_committed_response_loop_with_reasoning(
9886 &self,
9887 processed_input: &str,
9888 input_context: &HashMap<String, Value>,
9889 reasoning_mode: ReasoningMode,
9890 auto_detected: bool,
9891 ) -> Result<AgentResponse> {
9892 self.commit_root_user_message(processed_input).await?;
9893 let llm = self.get_state_llm()?;
9894 let mut iterations = 0u32;
9895 let mut all_tool_calls = Vec::new();
9896 let mut thinking_content = None;
9897 loop {
9898 let effective_max = if reasoning_mode != ReasoningMode::None {
9899 let rc = self.get_effective_reasoning_config();
9900 self.max_iterations.min(rc.max_iterations)
9901 } else {
9902 self.max_iterations
9903 };
9904 if iterations >= effective_max {
9905 return Err(AgentError::Other(format!(
9906 "Max iterations ({}) exceeded",
9907 effective_max
9908 )));
9909 }
9910 iterations += 1;
9911 *self.iteration_count.write() = iterations;
9912 let protocol = self.main_tool_protocol(llm.as_ref(), false).await?;
9913 let mut messages = self
9914 .build_messages_internal(true, None, protocol.choice.is_none())
9915 .await?;
9916 self.inject_reasoning_prompt(&mut messages, &reasoning_mode, iterations == 1);
9917 self.hooks.on_llm_start(&messages).await;
9918 let llm_start = Instant::now();
9919 let response = self
9920 .complete_main_llm_with_recovery(Arc::clone(&llm), &messages, &protocol)
9921 .await?;
9922 let llm_duration_ms = llm_start.elapsed().as_millis() as u64;
9923 self.hooks.on_llm_complete(&response, llm_duration_ms).await;
9924 let content = response.content.trim();
9925 if let Some(tool_calls) = self.parse_main_tool_calls(content, &protocol)? {
9926 match self
9927 .handle_tool_calls(processed_input, content, tool_calls, &mut all_tool_calls)
9928 .await?
9929 {
9930 ToolCallOutcome::Continue | ToolCallOutcome::TransitionFired => continue,
9931 ToolCallOutcome::Rejected(resp) => {
9932 self.finish_turn_if_root(&resp).await?;
9933 return Ok(resp);
9934 }
9935 }
9936 }
9937 let (extracted_thinking, answer) = self.extract_thinking(content);
9938 if extracted_thinking.is_some() {
9939 thinking_content = extracted_thinking;
9940 }
9941 return self
9942 .finish_text_response_from_model(CommittedTextResponse {
9943 processed_input,
9944 input_context,
9945 answer,
9946 reasoning_mode,
9947 auto_detected,
9948 iterations,
9949 thinking_content,
9950 all_tool_calls,
9951 })
9952 .await;
9953 }
9954 }
9955
9956 async fn handle_tool_calls(
9958 &self,
9959 processed_input: &str,
9960 content: &str,
9961 tool_calls: Vec<ToolCall>,
9962 all_tool_calls: &mut Vec<ToolCall>,
9963 ) -> Result<ToolCallOutcome> {
9964 let transition_content = native_readable_projection(content)
9968 .map_err(|error| AgentError::LLM(error.to_string()))?;
9969 let transition_fired = self
9970 .evaluate_transitions(processed_input, &transition_content)
9971 .await?;
9972 if transition_fired {
9973 self.memory
9974 .add_message(ChatMessage::assistant(
9975 "(Transitioned to new state — tool call handled by workflow)",
9976 ))
9977 .await?;
9978 return Ok(ToolCallOutcome::TransitionFired);
9979 }
9980
9981 self.memory
9983 .add_message(ChatMessage::assistant(content))
9984 .await?;
9985 self.remember_committed_native_exchange(content).await?;
9986 let native_tool_call = Self::is_native_tool_call_content(content)?;
9987
9988 let results = self.execute_tools_parallel(&tool_calls).await;
9989 let mut rejection = None;
9990
9991 for ((_id, result), tool_call) in results.into_iter().zip(tool_calls.iter()) {
9992 match result {
9993 Ok(output) => {
9994 self.memory
9995 .add_message(Self::tool_result_message(
9996 tool_call,
9997 &output,
9998 native_tool_call,
9999 )?)
10000 .await?;
10001 }
10002 Err(e) => {
10003 if matches!(e, AgentError::HITLRejected(_)) {
10004 if !native_tool_call {
10005 self.memory
10006 .add_message(ChatMessage::assistant(format!(
10007 "The operation was rejected by the approver: {e}"
10008 )))
10009 .await?;
10010 return Ok(ToolCallOutcome::Rejected(AgentResponse {
10011 content: format!("Operation cancelled: {e}"),
10012 metadata: None,
10013 tool_calls: Some(all_tool_calls.clone()),
10014 }));
10015 }
10016 if rejection.is_none() {
10017 rejection = Some(e.to_string());
10018 }
10019 }
10020 self.memory
10021 .add_message(Self::tool_result_message(
10022 tool_call,
10023 &format!("Error: {}", e),
10024 native_tool_call,
10025 )?)
10026 .await?;
10027 }
10028 }
10029 all_tool_calls.push(tool_call.clone());
10030 }
10031 if let Some(rejection) = rejection {
10032 self.memory
10033 .add_message(ChatMessage::assistant(format!(
10034 "The operation was rejected by the approver: {rejection}"
10035 )))
10036 .await?;
10037 return Ok(ToolCallOutcome::Rejected(AgentResponse {
10038 content: format!("Operation cancelled: {rejection}"),
10039 metadata: None,
10040 tool_calls: Some(all_tool_calls.clone()),
10041 }));
10042 }
10043 Ok(ToolCallOutcome::Continue)
10044 }
10045
10046 async fn run_reflection(
10048 &self,
10049 llm: &dyn LLMProvider,
10050 processed_input: &str,
10051 mut content: String,
10052 ) -> Result<(String, Option<ReflectionMetadata>)> {
10053 let should_reflect = self.should_reflect(processed_input, &content).await?;
10054 if !should_reflect {
10055 return Ok((content, None));
10056 }
10057
10058 info!("Starting response reflection evaluation");
10059 let mut attempts = 0u32;
10060 let max_retries = self.reflection_config.max_retries;
10061 let mut history: Vec<ReflectionAttempt> = Vec::new();
10062
10063 loop {
10064 let evaluation = self.evaluate_response(processed_input, &content).await?;
10065
10066 if evaluation.passed || attempts >= max_retries {
10067 info!(
10068 passed = evaluation.passed,
10069 confidence = evaluation.confidence,
10070 attempts = attempts + 1,
10071 "Reflection evaluation complete"
10072 );
10073 let reflection_metadata = Some(
10074 ReflectionMetadata::new(evaluation)
10075 .with_attempts(attempts + 1)
10076 .with_history(history),
10077 );
10078 return Ok((content, reflection_metadata));
10079 }
10080
10081 debug!(
10082 attempt = attempts + 1,
10083 failed_criteria = evaluation.failed_criteria().count(),
10084 "Response did not meet criteria, retrying"
10085 );
10086
10087 history.push(
10088 ReflectionAttempt::new(&content, evaluation.clone())
10089 .with_feedback("Response did not meet quality criteria"),
10090 );
10091
10092 let feedback: Vec<String> = evaluation
10093 .failed_criteria()
10094 .map(|c| format!("- {}", c.criterion))
10095 .collect();
10096
10097 let retry_prompt = format!(
10098 "Your previous response did not meet these criteria:\n{}\n\nPlease provide an improved response.",
10099 feedback.join("\n")
10100 );
10101
10102 self.memory
10103 .add_message(ChatMessage::user(&retry_prompt))
10104 .await?;
10105
10106 let retry_messages = self.build_messages().await?;
10107 let retry_response = self
10108 .observe_purpose(
10109 ObservationPurpose::ReflectionEvaluation,
10110 llm.complete(&retry_messages, None),
10111 )
10112 .await
10113 .map_err(|e| AgentError::LLM(e.to_string()))?;
10114
10115 content = retry_response.content.trim().to_string();
10116 attempts += 1;
10117 }
10118 }
10119
10120 async fn post_loop_processing(
10123 &self,
10124 processed_input: &str,
10125 content: String,
10126 ) -> Result<PostLoopResult> {
10127 self.increment_turn();
10132
10133 self.run_context_extractors(processed_input).await;
10135
10136 let transitioned = self.evaluate_transitions(processed_input, &content).await?;
10137
10138 if !transitioned {
10139 self.memory
10140 .add_message(ChatMessage::assistant(&content))
10141 .await?;
10142 self.check_memory_compression().await?;
10143 return Ok(PostLoopResult::NoTransition(content));
10144 }
10145
10146 if !self.should_regenerate_after_transition() {
10148 self.memory
10149 .add_message(ChatMessage::assistant(&content))
10150 .await?;
10151 self.check_memory_compression().await?;
10152 return Ok(PostLoopResult::Transitioned(content));
10153 }
10154
10155 if self.needs_redispatch_for_new_state() {
10159 info!("Post-transition NeedsRedispatch: new state requires full dispatch");
10160 return Ok(PostLoopResult::NeedsRedispatch);
10163 }
10164
10165 self.memory
10168 .add_message(ChatMessage::assistant(&content))
10169 .await?;
10170 self.check_memory_compression().await?;
10171
10172 let new_llm = self.get_state_llm()?;
10178 let mut final_content;
10179
10180 for post_iter in 0..self.max_iterations {
10181 let protocol = self.main_tool_protocol(new_llm.as_ref(), false).await?;
10182 let new_messages = self
10183 .build_messages_internal(true, None, protocol.choice.is_none())
10184 .await?;
10185 if post_iter == 0
10186 && let Some(system_msg) = new_messages.first()
10187 && system_msg.role == ai_agents_core::Role::System
10188 {
10189 debug!(
10190 prompt_preview =
10191 &system_msg.content[system_msg.content.len().saturating_sub(200)..],
10192 "Post-transition system prompt (last 200 chars)"
10193 );
10194 }
10195
10196 let new_response = self
10197 .complete_main_llm_with_recovery(Arc::clone(&new_llm), &new_messages, &protocol)
10198 .await?;
10199 final_content = new_response.content.trim().to_string();
10200
10201 if let Some(tool_calls) = self.parse_main_tool_calls(&final_content, &protocol)? {
10204 let native_tool_call = Self::is_native_tool_call_content(&final_content)?;
10205 debug!(
10206 post_iter = post_iter,
10207 tools = tool_calls.len(),
10208 "Post-transition tool call detected, executing"
10209 );
10210
10211 self.memory
10212 .add_message(ChatMessage::assistant(&final_content))
10213 .await?;
10214 self.remember_committed_native_exchange(&final_content)
10215 .await?;
10216
10217 let results = self.execute_tools_parallel(&tool_calls).await;
10218 let mut rejection = None;
10219 for ((_id, result), tool_call) in results.into_iter().zip(tool_calls.iter()) {
10220 match result {
10221 Ok(output) => {
10222 self.memory
10223 .add_message(Self::tool_result_message(
10224 tool_call,
10225 &output,
10226 native_tool_call,
10227 )?)
10228 .await?;
10229 }
10230 Err(e) => {
10231 if native_tool_call
10232 && rejection.is_none()
10233 && matches!(e, AgentError::HITLRejected(_))
10234 {
10235 rejection = Some(e.to_string());
10236 }
10237 self.memory
10238 .add_message(Self::tool_result_message(
10239 tool_call,
10240 &format!("Error: {}", e),
10241 native_tool_call,
10242 )?)
10243 .await?;
10244 }
10245 }
10246 }
10247 if let Some(rejection) = rejection {
10248 self.memory
10249 .add_message(ChatMessage::assistant(format!(
10250 "The operation was rejected by the approver: {rejection}"
10251 )))
10252 .await?;
10253 return Err(AgentError::HITLRejected(rejection));
10254 }
10255 continue;
10257 }
10258
10259 self.memory
10261 .add_message(ChatMessage::assistant(&final_content))
10262 .await?;
10263 return Ok(PostLoopResult::Transitioned(final_content));
10264 }
10265
10266 final_content = "Post-transition processing completed.".to_string();
10268 self.memory
10269 .add_message(ChatMessage::assistant(&final_content))
10270 .await?;
10271
10272 Ok(PostLoopResult::Transitioned(final_content))
10273 }
10274
10275 fn should_regenerate_after_transition(&self) -> bool {
10278 if let Some(ref sm) = self.state_machine {
10279 if !sm.config().regenerate_on_transition {
10281 return false;
10282 }
10283 if let Some(def) = sm.current_definition()
10285 && let Some(regen) = def.regenerate_on_enter
10286 {
10287 return regen;
10288 }
10289 }
10290 true
10291 }
10292
10293 fn needs_redispatch_for_new_state(&self) -> bool {
10296 if let Some(ref sm) = self.state_machine
10297 && let Some(def) = sm.current_definition()
10298 {
10299 if def.concurrent.is_some()
10300 || def.group_chat.is_some()
10301 || def.pipeline.is_some()
10302 || def.handoff.is_some()
10303 || def.delegate.is_some()
10304 {
10305 return true;
10306 }
10307 let effective = self.get_effective_reasoning_config();
10309 if !matches!(effective.mode, ReasoningMode::None) {
10310 return true;
10311 }
10312 }
10313 false
10314 }
10315
10316 async fn apply_post_loop_result(
10319 &self,
10320 processed_input: &str,
10321 result: PostLoopResult,
10322 ) -> Result<String> {
10323 match result {
10324 PostLoopResult::NoTransition(content) | PostLoopResult::Transitioned(content) => {
10325 Ok(content)
10326 }
10327 PostLoopResult::NeedsRedispatch => {
10328 const MAX_REDISPATCH_DEPTH: u32 = 3;
10329 let current_depth = *self.redispatch_depth.read();
10330 if current_depth >= MAX_REDISPATCH_DEPTH {
10331 warn!(
10332 depth = current_depth,
10333 "Post-transition re-dispatch depth limit reached, returning empty response"
10334 );
10335 let content = String::new();
10336 self.memory
10337 .add_message(ChatMessage::assistant(&content))
10338 .await?;
10339 return Ok(content);
10340 }
10341 *self.redispatch_depth.write() += 1;
10342 if let Some(context) = self.active_turn_context.write().as_mut() {
10343 context.enter_redispatch();
10344 }
10345 info!(
10346 depth = current_depth + 1,
10347 "Re-dispatching for new state after transition"
10348 );
10349 let resp = Box::pin(self.run_loop_internal(processed_input)).await;
10350 *self.redispatch_depth.write() -= 1;
10351 if let Some(context) = self.active_turn_context.write().as_mut() {
10352 context.exit_redispatch();
10353 }
10354 resp.map(|r| r.content)
10355 }
10356 }
10357 }
10358
10359 fn build_agent_response(&self, parts: AgentResponseParts) -> AgentResponse {
10361 let AgentResponseParts {
10362 content,
10363 all_tool_calls,
10364 reasoning_mode,
10365 auto_detected,
10366 iterations,
10367 thinking,
10368 reflection_metadata,
10369 } = parts;
10370 let reasoning_metadata = ReasoningMetadata::new(reasoning_mode.clone())
10371 .with_thinking(thinking.clone().unwrap_or_default())
10372 .with_iterations(iterations)
10373 .with_auto_detected(auto_detected);
10374
10375 let mut response = AgentResponse::new(&content);
10376 if !all_tool_calls.is_empty() {
10377 response = response.with_tool_calls(all_tool_calls);
10378 }
10379
10380 if let Some(state) = self.current_state() {
10381 response = response.with_metadata("current_state", serde_json::json!(state));
10382 }
10383
10384 response = response.with_metadata(
10385 "reasoning",
10386 serde_json::to_value(&reasoning_metadata).unwrap_or_default(),
10387 );
10388
10389 if let Some(ref refl_meta) = reflection_metadata {
10390 response = response.with_metadata(
10391 "reflection",
10392 serde_json::to_value(refl_meta).unwrap_or_default(),
10393 );
10394 }
10395
10396 response
10397 }
10398
10399 async fn handle_delegated_state(
10401 &self,
10402 input: &str,
10403 delegate_id: &str,
10404 state_def: &ai_agents_state::StateDefinition,
10405 ) -> Result<AgentResponse> {
10406 use std::time::Instant;
10407
10408 let registry = self.spawner_registry.as_ref().ok_or_else(|| {
10409 AgentError::Config(format!(
10410 "State delegates to '{}' but no agent registry is configured. \
10411 Add a spawner section with auto_spawn to your YAML.",
10412 delegate_id
10413 ))
10414 })?;
10415
10416 let state_name = self
10417 .state_machine
10418 .as_ref()
10419 .map(|sm| sm.current())
10420 .unwrap_or_else(|| "unknown".to_string());
10421
10422 self.hooks.on_delegate_start(delegate_id, &state_name).await;
10423 let start = Instant::now();
10424
10425 let delegate = registry.get(delegate_id).ok_or_else(|| {
10426 AgentError::Other(format!(
10427 "State '{}' delegates to '{}' but no agent with that ID exists in the registry.",
10428 state_name, delegate_id
10429 ))
10430 })?;
10431
10432 let context_mode = state_def.delegate_context.clone().unwrap_or_default();
10434 let effective_input = self
10435 .observe_purpose(
10436 ObservationPurpose::OrchestrationRouting,
10437 crate::orchestration::context::prepare_delegate_input(
10438 input,
10439 &context_mode,
10440 &*self.memory,
10441 self.llm_registry.get("router").ok().as_deref(),
10442 ),
10443 )
10444 .await?;
10445
10446 let response = delegate
10447 .chat_with_actor_context(&effective_input, self.outbound_actor_context())
10448 .await?;
10449
10450 let duration_ms = start.elapsed().as_millis() as u64;
10451 self.hooks
10452 .on_delegate_complete(delegate_id, &state_name, duration_ms)
10453 .await;
10454
10455 let ctx_key = format!("delegation.{}.last_response", delegate_id);
10457 let _ = self.context_manager.set(
10458 &ctx_key,
10459 serde_json::Value::String(response.content.clone()),
10460 );
10461
10462 let _ = self.context_manager.set(
10464 "orchestration",
10465 serde_json::json!({
10466 "type": "delegate",
10467 "agent": delegate_id,
10468 "state": state_name,
10469 "response": response.content,
10470 "duration_ms": duration_ms,
10471 }),
10472 );
10473
10474 self.commit_root_user_message(input).await?;
10475
10476 let post_result = self
10479 .post_loop_processing(
10480 input,
10481 format!("[Delegated to {}]: {}", delegate_id, response.content),
10482 )
10483 .await?;
10484 let final_content = self.apply_post_loop_result(input, post_result).await?;
10485
10486 let mut result = AgentResponse::new(final_content);
10487
10488 let metadata = serde_json::json!({
10489 "orchestration": {
10490 "type": "delegate",
10491 "agent": delegate_id,
10492 "state": state_name,
10493 "response": response.content,
10494 "duration_ms": duration_ms,
10495 }
10496 });
10497 result.metadata = Some(
10498 serde_json::from_value::<std::collections::HashMap<String, serde_json::Value>>(
10499 metadata,
10500 )
10501 .unwrap_or_default(),
10502 );
10503
10504 self.finish_turn_if_root(&result).await?;
10505 Ok(result)
10506 }
10507
10508 async fn handle_concurrent_state(
10510 &self,
10511 input: &str,
10512 config: &ai_agents_state::ConcurrentStateConfig,
10513 ) -> Result<AgentResponse> {
10514 use std::time::Instant;
10515
10516 let registry = self.spawner_registry.as_ref().ok_or_else(|| {
10517 AgentError::Config(
10518 "Concurrent state requires an agent registry. Add a spawner section.".into(),
10519 )
10520 })?;
10521
10522 let context_mode = config.context_mode.clone().unwrap_or_default();
10527 let context_input = self
10528 .observe_purpose(
10529 ObservationPurpose::OrchestrationRouting,
10530 crate::orchestration::context::prepare_delegate_input(
10531 input,
10532 &context_mode,
10533 &*self.memory,
10534 self.llm_registry.get("router").ok().as_deref(),
10535 ),
10536 )
10537 .await?;
10538
10539 let effective_input = if let Some(ref tmpl) = config.input {
10540 render_concurrent_template(tmpl, &context_input, &self.build_context_with_overlays())
10541 .unwrap_or_else(|_| context_input.clone())
10542 } else {
10543 context_input
10544 };
10545
10546 let start = Instant::now();
10547
10548 let llm_name = config
10549 .aggregation
10550 .synthesizer_llm
10551 .as_deref()
10552 .unwrap_or("router");
10553 let llm_provider = self.llm_registry.get(llm_name).ok();
10554
10555 let vote_parallelism = if self.runtime_config.optimization.enabled
10556 && self
10557 .runtime_config
10558 .optimization
10559 .parallel_orchestration_vote_extraction
10560 {
10561 Some(self.runtime_config.optimization.max_parallel_runtime_tasks)
10562 } else {
10563 None
10564 };
10565
10566 let result = self
10567 .observe_purpose(
10568 ObservationPurpose::OrchestrationAggregation,
10569 scope_actor_context(
10570 self.outbound_actor_context(),
10571 crate::orchestration::concurrent(
10572 registry,
10573 &effective_input,
10574 &config.agents,
10575 &config.aggregation,
10576 llm_provider.as_deref(),
10577 config.min_required,
10578 config.timeout_ms,
10579 config.on_partial_failure.clone(),
10580 vote_parallelism,
10581 ),
10582 ),
10583 )
10584 .await?;
10585
10586 let duration_ms = start.elapsed().as_millis() as u64;
10587 let agent_ids: Vec<String> = config.agents.iter().map(|a| a.id().to_string()).collect();
10588 let strategy = format!("{:?}", config.aggregation.strategy);
10589 self.hooks
10590 .on_concurrent_complete(&agent_ids, &strategy, duration_ms)
10591 .await;
10592
10593 let _ = self.context_manager.set(
10595 "concurrent.result",
10596 serde_json::Value::String(result.response.content.clone()),
10597 );
10598
10599 let agents_json: Vec<serde_json::Value> = result
10601 .agent_results
10602 .iter()
10603 .map(|ar| {
10604 serde_json::json!({
10605 "id": ar.agent_id,
10606 "response": ar.response.as_ref().map(|r| r.content.as_str()),
10607 "success": ar.success,
10608 "error": ar.error,
10609 "duration_ms": ar.duration_ms,
10610 })
10611 })
10612 .collect();
10613
10614 let _ = self.context_manager.set(
10616 "orchestration",
10617 serde_json::json!({
10618 "type": "concurrent",
10619 "result": result.response.content,
10620 "strategy": strategy,
10621 "agents": agents_json,
10622 "duration_ms": duration_ms,
10623 }),
10624 );
10625
10626 self.commit_root_user_message(input).await?;
10627
10628 let post_result = self
10629 .post_loop_processing(input, result.response.content.clone())
10630 .await?;
10631 let final_content = self.apply_post_loop_result(input, post_result).await?;
10632
10633 let mut response = AgentResponse::new(final_content);
10634 let metadata = serde_json::json!({
10635 "orchestration": {
10636 "type": "concurrent",
10637 "result": result.response.content,
10638 "strategy": strategy,
10639 "agents": agents_json,
10640 "duration_ms": duration_ms,
10641 }
10642 });
10643 response.metadata = Some(
10644 serde_json::from_value::<std::collections::HashMap<String, serde_json::Value>>(
10645 metadata,
10646 )
10647 .unwrap_or_default(),
10648 );
10649
10650 self.finish_turn_if_root(&response).await?;
10651 Ok(response)
10652 }
10653
10654 async fn handle_group_chat_state(
10656 &self,
10657 input: &str,
10658 config: &ai_agents_state::GroupChatStateConfig,
10659 ) -> Result<AgentResponse> {
10660 use std::time::Instant;
10661
10662 let registry = self.spawner_registry.as_ref().ok_or_else(|| {
10663 AgentError::Config(
10664 "Group chat state requires an agent registry. Add a spawner section.".into(),
10665 )
10666 })?;
10667
10668 let start = Instant::now();
10669
10670 let llm_provider = self.llm_registry.get("router").ok();
10671
10672 let context_mode = config.context_mode.clone().unwrap_or_default();
10674 let context_input = self
10675 .observe_purpose(
10676 ObservationPurpose::OrchestrationRouting,
10677 crate::orchestration::context::prepare_delegate_input(
10678 input,
10679 &context_mode,
10680 &*self.memory,
10681 self.llm_registry.get("router").ok().as_deref(),
10682 ),
10683 )
10684 .await?;
10685
10686 let effective_topic = if let Some(ref tmpl) = config.input {
10688 render_concurrent_template(tmpl, &context_input, &self.build_context_with_overlays())
10689 .unwrap_or_else(|_| context_input.clone())
10690 } else {
10691 context_input
10692 };
10693
10694 let result = self
10695 .observe_purpose(
10696 ObservationPurpose::OrchestrationConversation,
10697 scope_actor_context(
10698 self.outbound_actor_context(),
10699 crate::orchestration::group_chat(
10700 registry,
10701 &effective_topic,
10702 config,
10703 llm_provider.as_deref(),
10704 Some(&*self.hooks),
10705 ),
10706 ),
10707 )
10708 .await?;
10709
10710 let duration_ms = start.elapsed().as_millis() as u64;
10711
10712 let _ = self.context_manager.set(
10714 "group_chat.conclusion",
10715 serde_json::Value::String(result.response.content.clone()),
10716 );
10717
10718 let transcript_json: Vec<serde_json::Value> = result
10720 .transcript
10721 .iter()
10722 .map(|t| {
10723 serde_json::json!({
10724 "speaker": t.speaker,
10725 "round": t.round,
10726 "content": t.content,
10727 })
10728 })
10729 .collect();
10730
10731 let _ = self.context_manager.set(
10733 "orchestration",
10734 serde_json::json!({
10735 "type": "group_chat",
10736 "conclusion": result.response.content,
10737 "transcript": transcript_json,
10738 "rounds": result.rounds_completed,
10739 "termination": result.termination_reason,
10740 "duration_ms": duration_ms,
10741 }),
10742 );
10743
10744 self.commit_root_user_message(input).await?;
10745
10746 let post_result = self
10747 .post_loop_processing(input, result.response.content.clone())
10748 .await?;
10749 let final_content = self.apply_post_loop_result(input, post_result).await?;
10750
10751 let mut response = AgentResponse::new(final_content);
10752 let metadata = serde_json::json!({
10753 "orchestration": {
10754 "type": "group_chat",
10755 "conclusion": result.response.content,
10756 "transcript": transcript_json,
10757 "rounds": result.rounds_completed,
10758 "termination": result.termination_reason,
10759 "duration_ms": duration_ms,
10760 }
10761 });
10762 response.metadata = Some(
10763 serde_json::from_value::<std::collections::HashMap<String, serde_json::Value>>(
10764 metadata,
10765 )
10766 .unwrap_or_default(),
10767 );
10768
10769 self.finish_turn_if_root(&response).await?;
10770 Ok(response)
10771 }
10772
10773 async fn handle_pipeline_state(
10775 &self,
10776 input: &str,
10777 config: &ai_agents_state::PipelineStateConfig,
10778 ) -> Result<AgentResponse> {
10779 use std::time::Instant;
10780
10781 let registry = self.spawner_registry.as_ref().ok_or_else(|| {
10782 AgentError::Config(
10783 "Pipeline state requires an agent registry. Add a spawner section.".into(),
10784 )
10785 })?;
10786
10787 let start = Instant::now();
10788
10789 let stages: Vec<crate::orchestration::PipelineStage> = config
10790 .stages
10791 .iter()
10792 .map(|entry| {
10793 let mut stage = crate::orchestration::PipelineStage::id(entry.id());
10794 if let Some(tmpl) = entry.input() {
10795 stage = stage.with_input(tmpl);
10796 }
10797 stage
10798 })
10799 .collect();
10800
10801 let context_mode = config.context_mode.clone().unwrap_or_default();
10803 let context_input = self
10804 .observe_purpose(
10805 ObservationPurpose::OrchestrationRouting,
10806 crate::orchestration::context::prepare_delegate_input(
10807 input,
10808 &context_mode,
10809 &*self.memory,
10810 self.llm_registry.get("router").ok().as_deref(),
10811 ),
10812 )
10813 .await?;
10814
10815 let context_values = self.build_context_with_overlays();
10816 let result = self
10817 .observe_purpose(
10818 ObservationPurpose::OrchestrationRouting,
10819 scope_actor_context(
10820 self.outbound_actor_context(),
10821 crate::orchestration::pipeline(
10822 registry,
10823 &context_input,
10824 &stages,
10825 config.timeout_ms,
10826 Some(&*self.hooks),
10827 Some(&context_values),
10828 ),
10829 ),
10830 )
10831 .await?;
10832
10833 let duration_ms = start.elapsed().as_millis() as u64;
10834
10835 let _ = self.context_manager.set(
10837 "pipeline.result",
10838 serde_json::Value::String(result.response.content.clone()),
10839 );
10840
10841 let stages_json: Vec<serde_json::Value> = result
10843 .stage_outputs
10844 .iter()
10845 .map(|s| {
10846 serde_json::json!({
10847 "agent_id": s.agent_id,
10848 "output": s.output,
10849 "duration_ms": s.duration_ms,
10850 "skipped": s.skipped,
10851 })
10852 })
10853 .collect();
10854
10855 let _ = self.context_manager.set(
10857 "orchestration",
10858 serde_json::json!({
10859 "type": "pipeline",
10860 "result": result.response.content,
10861 "stages": stages_json,
10862 "duration_ms": duration_ms,
10863 }),
10864 );
10865
10866 self.commit_root_user_message(input).await?;
10867
10868 let post_result = self
10869 .post_loop_processing(input, result.response.content.clone())
10870 .await?;
10871 let final_content = self.apply_post_loop_result(input, post_result).await?;
10872
10873 let mut response = AgentResponse::new(final_content);
10874 let metadata = serde_json::json!({
10875 "orchestration": {
10876 "type": "pipeline",
10877 "result": result.response.content,
10878 "stages": stages_json,
10879 "duration_ms": duration_ms,
10880 }
10881 });
10882 response.metadata = Some(
10883 serde_json::from_value::<std::collections::HashMap<String, serde_json::Value>>(
10884 metadata,
10885 )
10886 .unwrap_or_default(),
10887 );
10888
10889 self.finish_turn_if_root(&response).await?;
10890 Ok(response)
10891 }
10892
10893 async fn handle_handoff_state(
10895 &self,
10896 input: &str,
10897 config: &ai_agents_state::HandoffStateConfig,
10898 ) -> Result<AgentResponse> {
10899 use std::time::Instant;
10900
10901 let registry = self.spawner_registry.as_ref().ok_or_else(|| {
10902 AgentError::Config(
10903 "Handoff state requires an agent registry. Add a spawner section.".into(),
10904 )
10905 })?;
10906
10907 let llm = self
10908 .llm_registry
10909 .get("router")
10910 .map_err(|_| AgentError::Config("Handoff state requires a router LLM.".into()))?;
10911
10912 let start = Instant::now();
10913
10914 let context_mode = config.context_mode.clone().unwrap_or_default();
10916 let context_input = self
10917 .observe_purpose(
10918 ObservationPurpose::OrchestrationRouting,
10919 crate::orchestration::context::prepare_delegate_input(
10920 input,
10921 &context_mode,
10922 &*self.memory,
10923 self.llm_registry.get("router").ok().as_deref(),
10924 ),
10925 )
10926 .await?;
10927
10928 let effective_input = if let Some(ref tmpl) = config.input {
10930 render_concurrent_template(tmpl, &context_input, &self.build_context_with_overlays())
10931 .unwrap_or_else(|_| context_input.clone())
10932 } else {
10933 context_input
10934 };
10935
10936 let result = self
10937 .observe_purpose(
10938 ObservationPurpose::OrchestrationRouting,
10939 scope_actor_context(
10940 self.outbound_actor_context(),
10941 crate::orchestration::handoff(
10942 registry,
10943 &effective_input,
10944 &config.initial_agent,
10945 &config.available_agents,
10946 config.max_handoffs,
10947 llm.as_ref(),
10948 Some(&*self.hooks),
10949 ),
10950 ),
10951 )
10952 .await?;
10953
10954 let duration_ms = start.elapsed().as_millis() as u64;
10955
10956 let _ = self.context_manager.set(
10958 "handoff.result",
10959 serde_json::Value::String(result.response.content.clone()),
10960 );
10961
10962 let chain_json: Vec<serde_json::Value> = result
10964 .handoff_chain
10965 .iter()
10966 .map(|h| {
10967 serde_json::json!({
10968 "from": h.from_agent,
10969 "to": h.to_agent,
10970 "reason": h.reason,
10971 })
10972 })
10973 .collect();
10974
10975 let _ = self.context_manager.set(
10977 "orchestration",
10978 serde_json::json!({
10979 "type": "handoff",
10980 "result": result.response.content,
10981 "final_agent": result.final_agent,
10982 "handoff_chain": chain_json,
10983 "duration_ms": duration_ms,
10984 }),
10985 );
10986
10987 self.commit_root_user_message(input).await?;
10988
10989 let post_result = self
10990 .post_loop_processing(input, result.response.content.clone())
10991 .await?;
10992 let final_content = self.apply_post_loop_result(input, post_result).await?;
10993
10994 let mut response = AgentResponse::new(final_content);
10995 let metadata = serde_json::json!({
10996 "orchestration": {
10997 "type": "handoff",
10998 "result": result.response.content,
10999 "final_agent": result.final_agent,
11000 "handoff_chain": chain_json,
11001 "duration_ms": duration_ms,
11002 }
11003 });
11004 response.metadata = Some(
11005 serde_json::from_value::<std::collections::HashMap<String, serde_json::Value>>(
11006 metadata,
11007 )
11008 .unwrap_or_default(),
11009 );
11010
11011 self.finish_turn_if_root(&response).await?;
11012 Ok(response)
11013 }
11014
11015 async fn run_loop_internal(&self, input: &str) -> Result<AgentResponse> {
11017 self.begin_root_turn();
11018 self.pre_turn_session_lifecycle().await;
11020
11021 let input_data = self.process_input(input).await?;
11022 self.update_active_turn_context(&input_data.content, input_data.context.clone());
11023
11024 for (key, value) in &input_data.context {
11027 let _ = self.context_manager.set(key, value.clone());
11028 }
11029
11030 if input_data.metadata.rejected {
11031 let reason = input_data
11032 .metadata
11033 .rejection_reason
11034 .unwrap_or_else(|| "Input rejected".to_string());
11035 warn!(reason = %reason, "Input rejected");
11036 let response = AgentResponse::new(reason);
11037 self.finish_turn_if_root(&response).await?;
11038 return Ok(response);
11039 }
11040
11041 let processed_input = &input_data.content;
11042
11043 if let Some(response) = self.try_pre_response_transition(processed_input).await? {
11044 return Ok(response);
11045 }
11046
11047 if let Some(ref sm) = self.state_machine
11049 && let Some(def) = sm.current_definition()
11050 {
11051 if let Some(ref delegate_id) = def.delegate {
11052 return self
11053 .handle_delegated_state(processed_input, delegate_id, &def)
11054 .await;
11055 }
11056 if let Some(ref concurrent_config) = def.concurrent {
11057 return self
11058 .handle_concurrent_state(processed_input, concurrent_config)
11059 .await;
11060 }
11061 if let Some(ref group_chat_config) = def.group_chat {
11062 return self
11063 .handle_group_chat_state(processed_input, group_chat_config)
11064 .await;
11065 }
11066 if let Some(ref pipeline_config) = def.pipeline {
11067 return self
11068 .handle_pipeline_state(processed_input, pipeline_config)
11069 .await;
11070 }
11071 if let Some(ref handoff_config) = def.handoff {
11072 return self
11073 .handle_handoff_state(processed_input, handoff_config)
11074 .await;
11075 }
11076 }
11077
11078 if let Some(response) =
11083 Box::pin(self.try_speculative_branches(processed_input, &input_data.context)).await?
11084 {
11085 return Ok(response);
11086 }
11087
11088 match self.try_skill_route(processed_input).await? {
11089 SkillRouteResult::Response { skill_id, content } => {
11090 self.commit_root_user_message(processed_input).await?;
11091 return self
11092 .handle_skill_response(processed_input, &skill_id, content, &input_data.context)
11093 .await;
11094 }
11095 SkillRouteResult::NeedsClarification {
11096 response,
11097 ownership,
11098 } => {
11099 let admission = self
11100 .admit_optional_disambiguation_ownership(ownership)
11101 .await?;
11102 self.commit_root_user_message(processed_input).await?;
11103 if let Some(q) = response
11104 .metadata
11105 .as_ref()
11106 .and_then(|m| m.get("disambiguation"))
11107 .and_then(|d| d.get("status"))
11108 .and_then(|s| s.as_str())
11109 && q == "awaiting_clarification"
11110 {
11111 self.memory
11114 .add_message(ChatMessage::assistant(&response.content))
11115 .await?;
11116 }
11117 drop(admission);
11118 self.finish_turn_if_root(&response).await?;
11119 return Ok(response);
11120 }
11121 SkillRouteResult::NoMatch => {} }
11123
11124 let effective_reasoning = self.get_effective_reasoning_config();
11125 let reasoning_mode = self.determine_reasoning_mode(processed_input).await?;
11126 let auto_detected = matches!(effective_reasoning.mode, ReasoningMode::Auto);
11127
11128 info!(
11129 reasoning_mode = ?reasoning_mode,
11130 auto_detected = auto_detected,
11131 reflection_enabled = ?self.reflection_config.enabled,
11132 "Reasoning mode determined"
11133 );
11134
11135 if matches!(reasoning_mode, ReasoningMode::PlanAndExecute) {
11136 self.commit_root_user_message(processed_input).await?;
11137 return self
11138 .handle_plan_and_execute(processed_input, &input_data.context, auto_detected)
11139 .await;
11140 }
11141
11142 self.commit_root_user_message(processed_input).await?;
11143
11144 let mut iterations = 0u32;
11145 let mut all_tool_calls: Vec<ToolCall> = Vec::new();
11146 let mut thinking_content: Option<String> = None;
11147
11148 let llm = self.get_state_llm()?;
11149
11150 loop {
11151 let effective_max = if reasoning_mode != ReasoningMode::None {
11153 let rc = self.get_effective_reasoning_config();
11154 self.max_iterations.min(rc.max_iterations)
11155 } else {
11156 self.max_iterations
11157 };
11158
11159 if iterations >= effective_max {
11160 let err = AgentError::Other(format!("Max iterations ({}) exceeded", effective_max));
11161 self.hooks.on_error(&err).await;
11162 error!(iterations = iterations, "Max iterations exceeded");
11163 return Err(err);
11164 }
11165 iterations += 1;
11166 *self.iteration_count.write() = iterations;
11167
11168 debug!(iteration = iterations, max = effective_max, "LLM call");
11169
11170 let protocol = self.main_tool_protocol(llm.as_ref(), false).await?;
11171 let mut messages = self
11172 .build_messages_internal(true, None, protocol.choice.is_none())
11173 .await?;
11174 self.inject_reasoning_prompt(&mut messages, &reasoning_mode, iterations == 1);
11175
11176 self.hooks.on_llm_start(&messages).await;
11177 let llm_start = Instant::now();
11178 let response = self
11179 .complete_main_llm_with_recovery(Arc::clone(&llm), &messages, &protocol)
11180 .await?;
11181
11182 let llm_duration_ms = llm_start.elapsed().as_millis() as u64;
11183 self.hooks.on_llm_complete(&response, llm_duration_ms).await;
11184
11185 let content = response.content.trim();
11186
11187 if let Some(tool_calls) = self.parse_main_tool_calls(content, &protocol)? {
11188 match self
11189 .handle_tool_calls(processed_input, content, tool_calls, &mut all_tool_calls)
11190 .await?
11191 {
11192 ToolCallOutcome::Continue | ToolCallOutcome::TransitionFired => continue,
11193 ToolCallOutcome::Rejected(resp) => {
11194 self.finish_turn_if_root(&resp).await?;
11195 return Ok(resp);
11196 }
11197 }
11198 }
11199
11200 let (extracted_thinking, answer) = self.extract_thinking(content);
11201 if extracted_thinking.is_some() {
11202 thinking_content = extracted_thinking;
11203 }
11204
11205 let output_data = self.process_output(&answer, &input_data.context).await?;
11206
11207 let mut final_content = if output_data.metadata.rejected {
11208 output_data
11209 .metadata
11210 .rejection_reason
11211 .unwrap_or_else(|| answer.to_string())
11212 } else {
11213 output_data.content
11214 };
11215
11216 let reflection_metadata;
11218 (final_content, reflection_metadata) = self
11219 .run_reflection(&*llm, processed_input, final_content)
11220 .await?;
11221
11222 final_content =
11223 self.format_response_with_thinking(thinking_content.as_deref(), &final_content);
11224
11225 let final_content = {
11229 let result = self
11230 .post_loop_processing(processed_input, final_content)
11231 .await?;
11232 self.apply_post_loop_result(processed_input, result).await?
11233 };
11234
11235 let reflected = reflection_metadata.is_some();
11236 let reasoning_mode_debug = format!("{:?}", reasoning_mode);
11237
11238 let response = self.build_agent_response(AgentResponseParts {
11239 content: final_content,
11240 all_tool_calls,
11241 reasoning_mode,
11242 auto_detected,
11243 iterations,
11244 thinking: thinking_content,
11245 reflection_metadata,
11246 });
11247
11248 self.finish_turn_if_root(&response).await?;
11249
11250 let tool_call_count = response.tool_calls.as_ref().map(|tc| tc.len()).unwrap_or(0);
11251 info!(
11252 tool_calls = tool_call_count,
11253 response_len = response.content.len(),
11254 reasoning_mode = %reasoning_mode_debug,
11255 reflected = reflected,
11256 "Chat completed"
11257 );
11258 return Ok(response);
11259 }
11260 }
11261
11262 async fn generate_buffered_streaming_draft(
11263 &self,
11264 processed_input: &str,
11265 routing_resolved: Arc<AtomicBool>,
11266 ) -> Result<StreamingDraftResult> {
11267 let llm = self.get_state_llm()?;
11268 if llm.configured_tool_choice().is_some() {
11269 let draft = self
11270 .generate_main_response_draft(processed_input, &ReasoningMode::None)
11271 .await?;
11272 return Ok(StreamingDraftResult::new(draft, Vec::new()));
11273 }
11274 let messages = self.build_messages_for_draft(processed_input).await?;
11275 let mut stream = self
11276 .observe_purpose(
11277 ObservationPurpose::MainResponse,
11278 llm.complete_stream(&messages, None),
11279 )
11280 .await
11281 .map_err(|e| AgentError::LLM(e.to_string()))?;
11282 let mut buffer = crate::optimization::StreamBranchBuffer::new(self.streaming.buffer_size)?;
11283 let mut chunks = Vec::new();
11284 let mut accumulated = String::new();
11285 while let Some(chunk_result) = stream.next().await {
11286 let chunk = chunk_result.map_err(|e| AgentError::LLM(e.to_string()))?;
11287 accumulated.push_str(&chunk.delta);
11288 let stream_chunk = StreamChunk::content(chunk.delta);
11289 if routing_resolved.load(Ordering::SeqCst) {
11290 chunks.push(stream_chunk);
11291 } else {
11292 buffer.push(stream_chunk)?;
11293 }
11294 }
11295 chunks.splice(0..0, buffer.drain());
11296 let content = accumulated.trim().to_string();
11297 let draft = if let Some(calls) = self.parse_tool_calls(&content)? {
11298 MainResponseDraft::ToolCalls {
11299 raw_content: content,
11300 calls,
11301 thinking: None,
11302 }
11303 } else {
11304 MainResponseDraft::Text {
11305 raw_content: content,
11306 thinking: None,
11307 }
11308 };
11309 Ok(StreamingDraftResult::new(draft, chunks))
11310 }
11311
11312 async fn try_buffered_streaming_branches(
11313 &self,
11314 processed_input: &str,
11315 input_context: &HashMap<String, Value>,
11316 ) -> Result<Option<(AgentResponse, Vec<StreamChunk>)>> {
11317 let optimization = &self.runtime_config.optimization;
11318 if !optimization.enabled {
11319 return Ok(None);
11320 }
11321 let transition_enabled =
11322 optimization.speculative_state_transitions && self.has_parallel_transition_candidates();
11323 if !transition_enabled {
11324 return Ok(None);
11325 }
11326 let mut branch_scheduler =
11327 TurnBranchScheduler::new(optimization.max_parallel_runtime_tasks)?;
11328 if !branch_scheduler.reserve_task() {
11329 return Ok(None);
11330 }
11331 if !self
11332 .reserve_active_speculative_llm_call(RuntimeOptimizationKind::BufferedStreamingRouting)
11333 {
11334 branch_scheduler.release_task();
11335 return Ok(None);
11336 }
11337 if !branch_scheduler.reserve_task() {
11338 branch_scheduler.release_task();
11339 return Ok(None);
11340 }
11341 let mut main_branch = RuntimeBranch::new(
11342 RuntimeTaskPurpose::MainResponse,
11343 RuntimeOptimizationKind::BufferedStreamingRouting,
11344 RuntimeTaskPriority::Normal,
11345 RuntimeCommitBehavior::FinalResponse,
11346 );
11347 let mut transition_branch = RuntimeBranch::new(
11348 RuntimeTaskPurpose::StateTransition,
11349 RuntimeOptimizationKind::ParallelStateTransition,
11350 RuntimeTaskPriority::Critical,
11351 RuntimeCommitBehavior::TransitionDecision,
11352 );
11353 let main_id = main_branch.branch_id();
11354 let transition_id = transition_branch.branch_id();
11355 let routing_resolved = Arc::new(AtomicBool::new(false));
11356 let mut main_future =
11357 Box::pin(crate::optimization::observability::with_branch_observation(
11358 &main_id,
11359 RuntimeOptimizationKind::BufferedStreamingRouting,
11360 RuntimeCommitBehavior::FinalResponse,
11361 self.generate_buffered_streaming_draft(
11362 processed_input,
11363 Arc::clone(&routing_resolved),
11364 ),
11365 ));
11366 let mut transition_future =
11367 Box::pin(crate::optimization::observability::with_branch_observation(
11368 &transition_id,
11369 RuntimeOptimizationKind::ParallelStateTransition,
11370 RuntimeCommitBehavior::TransitionDecision,
11371 self.select_parallel_transition_candidate(processed_input),
11372 ));
11373 let mut main_pending = true;
11374 let mut transition_pending = true;
11375 let mut main_result: Option<Result<StreamingDraftResult>> = None;
11376 let mut transition_finalized = false;
11377 let mut transition_candidate: Option<TransitionCandidate> = None;
11378 loop {
11379 if let Some(candidate) = transition_candidate.take() {
11380 if self
11381 .approve_transition_target(&candidate.from_state, candidate.target())
11382 .await?
11383 {
11384 drop(main_future);
11386 drop(transition_future);
11387 self.finalize_branch_loss(
11388 &main_id,
11389 RuntimeOptimizationKind::BufferedStreamingRouting,
11390 RuntimeCommitBehavior::FinalResponse,
11391 main_pending,
11392 main_result.as_ref().map(|result| result.is_err()),
11393 );
11394 if !self
11395 .apply_pre_response_transition_candidate(
11396 &candidate,
11397 &HashMap::new(),
11398 processed_input,
11399 )
11400 .await?
11401 {
11402 self.finalize_optional_branch(
11403 &transition_id,
11404 RuntimeOptimizationKind::ParallelStateTransition,
11405 RuntimeCommitBehavior::TransitionDecision,
11406 "discarded",
11407 false,
11408 );
11409 return Ok(None);
11410 }
11411 self.finalize_optional_branch(
11412 &transition_id,
11413 RuntimeOptimizationKind::ParallelStateTransition,
11414 RuntimeCommitBehavior::TransitionDecision,
11415 "committed",
11416 true,
11417 );
11418 let response = self.redispatch_current_state(processed_input).await?;
11419 return Ok(Some((
11420 response.clone(),
11421 vec![StreamChunk::content(response.content)],
11422 )));
11423 }
11424 self.finalize_optional_branch(
11425 &transition_id,
11426 RuntimeOptimizationKind::ParallelStateTransition,
11427 RuntimeCommitBehavior::TransitionDecision,
11428 "discarded",
11429 false,
11430 );
11431 routing_resolved.store(true, Ordering::SeqCst);
11432 transition_finalized = true;
11433 }
11434 if transition_finalized && let Some(result) = main_result.take() {
11435 let stream_draft = match result {
11436 Ok(stream_draft) => stream_draft,
11437 Err(error) => {
11438 self.finalize_optional_branch(
11439 &main_id,
11440 RuntimeOptimizationKind::BufferedStreamingRouting,
11441 RuntimeCommitBehavior::FinalResponse,
11442 "failed",
11443 false,
11444 );
11445 return Err(error);
11446 }
11447 };
11448 let raw_draft_content = stream_draft.draft.raw_content().to_string();
11449 let buffered_chunks = stream_draft.chunks;
11450 self.finalize_optional_branch(
11451 &main_id,
11452 RuntimeOptimizationKind::BufferedStreamingRouting,
11453 RuntimeCommitBehavior::FinalResponse,
11454 "committed",
11455 true,
11456 );
11457 let response = self
11458 .commit_main_response_draft(
11459 processed_input,
11460 input_context,
11461 stream_draft.draft,
11462 ReasoningMode::None,
11463 false,
11464 )
11465 .await?;
11466 let chunks = if response.content == raw_draft_content {
11467 buffered_chunks
11468 } else {
11469 vec![StreamChunk::content(response.content.clone())]
11470 };
11471 return Ok(Some((response, chunks)));
11472 }
11473 tokio::select! {
11474 result = &mut main_future, if main_pending => {
11475 main_pending = false;
11476 main_branch.transition_to(RuntimeBranchStatus::Completed)?;
11477 main_result = Some(result);
11478 }
11479 result = &mut transition_future, if transition_pending => {
11480 transition_pending = false;
11481 transition_branch.transition_to(RuntimeBranchStatus::Completed)?;
11482 match result {
11483 Ok(ParallelTransitionSelection::Candidate(candidate)) => {
11484 transition_candidate = Some(candidate)
11485 }
11486 Ok(ParallelTransitionSelection::NoMatch) => {
11487 self.finalize_optional_branch(
11488 &transition_id,
11489 RuntimeOptimizationKind::ParallelStateTransition,
11490 RuntimeCommitBehavior::TransitionDecision,
11491 "discarded",
11492 false,
11493 );
11494 routing_resolved.store(true, Ordering::SeqCst);
11495 transition_finalized = true;
11496 }
11497 Ok(ParallelTransitionSelection::ReservationExhausted) => {
11498 self.finalize_optional_branch(
11499 &transition_id,
11500 RuntimeOptimizationKind::ParallelStateTransition,
11501 RuntimeCommitBehavior::TransitionDecision,
11502 "cancelled",
11503 false,
11504 );
11505 routing_resolved.store(true, Ordering::SeqCst);
11506 self.finalize_branch_loss(
11507 &main_id,
11508 RuntimeOptimizationKind::BufferedStreamingRouting,
11509 RuntimeCommitBehavior::FinalResponse,
11510 main_pending,
11511 main_result.as_ref().map(|result| result.is_err()),
11512 );
11513 return Ok(None);
11514 }
11515 Err(_) => {
11516 self.finalize_optional_branch(
11517 &transition_id,
11518 RuntimeOptimizationKind::ParallelStateTransition,
11519 RuntimeCommitBehavior::TransitionDecision,
11520 "failed",
11521 false,
11522 );
11523 routing_resolved.store(true, Ordering::SeqCst);
11524 transition_finalized = true;
11525 }
11526 }
11527 }
11528 }
11529 }
11530 }
11531
11532 fn run_loop_internal_stream<'a>(
11536 &'a self,
11537 input: &'a str,
11538 terminal: RuntimeStreamTerminalSlot,
11539 ) -> Pin<Box<dyn Stream<Item = StreamChunk> + Send + 'a>> {
11540 let include_tool_events = self.streaming.include_tool_events;
11541 let include_state_events = self.streaming.include_state_events;
11542
11543 Box::pin(async_stream::stream! {
11544 self.begin_root_turn();
11545 self.pre_turn_session_lifecycle().await;
11547
11548 let input_data = match self.process_input(input).await {
11549 Ok(data) => data,
11550 Err(e) => {
11551 yield StreamChunk::error(e.to_string());
11552 return;
11553 }
11554 };
11555 self.update_active_turn_context(&input_data.content, input_data.context.clone());
11556
11557 for (key, value) in &input_data.context {
11559 let _ = self.context_manager.set(key, value.clone());
11560 }
11561
11562 if input_data.metadata.rejected {
11563 let reason = input_data
11564 .metadata
11565 .rejection_reason
11566 .unwrap_or_else(|| "Input rejected".to_string());
11567 warn!(reason = %reason, "Input rejected (stream)");
11568 yield StreamChunk::error(reason);
11569 return;
11570 }
11571
11572 let processed_input = &input_data.content;
11573
11574 if self.runtime_config.optimization.enabled
11575 && matches!(
11576 self.runtime_config.optimization.streaming_policy,
11577 crate::optimization::StreamingOptimizationPolicy::BufferUntilRoutingDone
11578 )
11579 {
11580 match Box::pin(self.try_buffered_streaming_branches(processed_input, &input_data.context)).await {
11585 Ok(Some((response, chunks))) => {
11586 for chunk in chunks {
11587 yield chunk;
11588 }
11589 record_runtime_stream_final(&terminal, response);
11590 yield StreamChunk::Done {};
11591 return;
11592 }
11593 Ok(None) => {}
11594 Err(e) => {
11595 yield StreamChunk::error(e.to_string());
11596 return;
11597 }
11598 }
11599 }
11600
11601 if self.runtime_config.optimization.enabled
11602 && matches!(
11603 self.runtime_config.optimization.streaming_policy,
11604 crate::optimization::StreamingOptimizationPolicy::PreflightOnly
11605 )
11606 {
11607 match self.try_pre_response_transition(processed_input).await {
11608 Ok(Some(response)) => {
11609 yield StreamChunk::content(&response.content);
11610 record_runtime_stream_final(&terminal, response);
11611 yield StreamChunk::Done {};
11612 return;
11613 }
11614 Ok(None) => {}
11615 Err(e) => {
11616 yield StreamChunk::error(e.to_string());
11617 return;
11618 }
11619 }
11620 }
11621
11622 if let Some(ref sm) = self.state_machine
11624 && let Some(def) = sm.current_definition()
11625 {
11626 let orchestration_result = if let Some(ref delegate_id) = def.delegate {
11627 Some(self.handle_delegated_state(processed_input, delegate_id, &def).await)
11628 } else if let Some(ref concurrent_config) = def.concurrent {
11629 Some(self.handle_concurrent_state(processed_input, concurrent_config).await)
11630 } else if let Some(ref group_chat_config) = def.group_chat {
11631 Some(self.handle_group_chat_state(processed_input, group_chat_config).await)
11632 } else if let Some(ref pipeline_config) = def.pipeline {
11633 Some(self.handle_pipeline_state(processed_input, pipeline_config).await)
11634 } else if let Some(ref handoff_config) = def.handoff {
11635 Some(self.handle_handoff_state(processed_input, handoff_config).await)
11636 } else {
11637 None
11638 };
11639
11640 if let Some(result) = orchestration_result {
11641 match result {
11642 Ok(response) => {
11643 yield StreamChunk::content(&response.content);
11644 record_runtime_stream_final(&terminal, response);
11645 yield StreamChunk::Done {};
11646 }
11647 Err(e) => {
11648 yield StreamChunk::error(e.to_string());
11649 }
11650 }
11651 return;
11652 }
11653 }
11654
11655 match self.try_skill_route(processed_input).await {
11657 Ok(SkillRouteResult::Response { skill_id, content }) => {
11658 if let Err(e) = self.commit_root_user_message(processed_input).await {
11659 yield StreamChunk::error(e.to_string());
11660 return;
11661 }
11662 match self.handle_skill_response(processed_input, &skill_id, content, &input_data.context).await {
11663 Ok(resp) => {
11664 yield StreamChunk::content(&resp.content);
11665 record_runtime_stream_final(&terminal, resp);
11666 yield StreamChunk::Done {};
11667 return;
11668 }
11669 Err(e) => {
11670 yield StreamChunk::error(e.to_string());
11671 return;
11672 }
11673 }
11674 }
11675 Ok(SkillRouteResult::NeedsClarification {
11676 response,
11677 ownership,
11678 }) => {
11679 let admission = match self
11680 .admit_optional_disambiguation_ownership(ownership)
11681 .await
11682 {
11683 Ok(admission) => admission,
11684 Err(e) => {
11685 yield StreamChunk::error(e.to_string());
11686 return;
11687 }
11688 };
11689 if let Err(e) = self.commit_root_user_message(processed_input).await {
11690 yield StreamChunk::error(e.to_string());
11691 return;
11692 }
11693 let _ = self.memory.add_message(ChatMessage::assistant(&response.content)).await;
11694 drop(admission);
11695 if let Err(e) = self.finish_turn_if_root(&response).await {
11696 yield StreamChunk::error(e.to_string());
11697 return;
11698 }
11699 yield StreamChunk::content(&response.content);
11700 record_runtime_stream_final(&terminal, response);
11701 yield StreamChunk::Done {};
11702 return;
11703 }
11704 Ok(SkillRouteResult::NoMatch) => {} Err(e) => {
11706 yield StreamChunk::error(e.to_string());
11707 return;
11708 }
11709 }
11710
11711 let effective_reasoning = self.get_effective_reasoning_config();
11713 let reasoning_mode = match self.determine_reasoning_mode(processed_input).await {
11714 Ok(mode) => mode,
11715 Err(e) => {
11716 yield StreamChunk::error(e.to_string());
11717 return;
11718 }
11719 };
11720 let auto_detected = matches!(effective_reasoning.mode, ReasoningMode::Auto);
11721
11722 info!(
11723 reasoning_mode = ?reasoning_mode,
11724 auto_detected = auto_detected,
11725 "Reasoning mode determined (stream)"
11726 );
11727
11728 if matches!(reasoning_mode, ReasoningMode::PlanAndExecute) {
11730 if let Err(e) = self.commit_root_user_message(processed_input).await {
11731 yield StreamChunk::error(e.to_string());
11732 return;
11733 }
11734 match self.handle_plan_and_execute(processed_input, &input_data.context, auto_detected).await {
11735 Ok(resp) => {
11736 yield StreamChunk::content(&resp.content);
11737 record_runtime_stream_final(&terminal, resp);
11738 yield StreamChunk::Done {};
11739 return;
11740 }
11741 Err(e) => {
11742 yield StreamChunk::error(e.to_string());
11743 return;
11744 }
11745 }
11746 }
11747
11748 if let Err(e) = self.commit_root_user_message(processed_input).await {
11749 yield StreamChunk::error(e.to_string());
11750 return;
11751 }
11752
11753 let llm = match self.get_state_llm() {
11754 Ok(llm) => llm,
11755 Err(e) => {
11756 yield StreamChunk::error(e.to_string());
11757 return;
11758 }
11759 };
11760
11761 let mut iterations = 0u32;
11762 let mut all_tool_calls: Vec<ToolCall> = Vec::new();
11763 let mut thinking_content: Option<String> = None;
11764
11765 loop {
11766 let effective_max = if reasoning_mode != ReasoningMode::None {
11768 let rc = self.get_effective_reasoning_config();
11769 self.max_iterations.min(rc.max_iterations)
11770 } else {
11771 self.max_iterations
11772 };
11773
11774 if iterations >= effective_max {
11775 let err_msg = format!("Max iterations ({}) exceeded", effective_max);
11776 let err = AgentError::Other(err_msg.clone());
11777 self.hooks.on_error(&err).await;
11778 error!(iterations = iterations, "Max iterations exceeded (stream)");
11779 yield StreamChunk::error(err_msg);
11780 return;
11781 }
11782 iterations += 1;
11783 *self.iteration_count.write() = iterations;
11784
11785 debug!(iteration = iterations, max = effective_max, "LLM call (stream)");
11786
11787 let protocol = match self.main_tool_protocol(llm.as_ref(), false).await {
11788 Ok(protocol) => protocol,
11789 Err(e) => {
11790 yield StreamChunk::error(e.to_string());
11791 return;
11792 }
11793 };
11794 let mut messages = match self
11795 .build_messages_internal(true, None, protocol.choice.is_none())
11796 .await
11797 {
11798 Ok(m) => m,
11799 Err(e) => {
11800 yield StreamChunk::error(e.to_string());
11801 return;
11802 }
11803 };
11804 self.inject_reasoning_prompt(&mut messages, &reasoning_mode, iterations == 1);
11805
11806 self.hooks.on_llm_start(&messages).await;
11807 let llm_start = Instant::now();
11808
11809 let reflection_active = self
11812 .should_reflect(processed_input, "")
11813 .await
11814 .unwrap_or_default();
11815
11816 let buffered_decision = reflection_active || protocol.choice.is_some();
11817 let content = if buffered_decision {
11818 let response = match self
11822 .complete_main_llm_with_recovery(
11823 Arc::clone(&llm),
11824 &messages,
11825 &protocol,
11826 )
11827 .await
11828 {
11829 Ok(r) => r,
11830 Err(e) => {
11831 yield StreamChunk::error(e.to_string());
11832 return;
11833 }
11834 };
11835 let llm_duration_ms = llm_start.elapsed().as_millis() as u64;
11836 self.hooks.on_llm_complete(&response, llm_duration_ms).await;
11837 response.content.trim().to_string()
11838 } else {
11839 let llm_stream = match self
11841 .observe_purpose(
11842 ObservationPurpose::MainResponse,
11843 llm.complete_stream(&messages, None),
11844 )
11845 .await
11846 {
11847 Ok(s) => s,
11848 Err(e) => {
11849 yield StreamChunk::error(e.to_string());
11850 return;
11851 }
11852 };
11853 let mut accumulated = String::new();
11854 let mut stream_inner = llm_stream;
11855 while let Some(chunk_result) = stream_inner.next().await {
11856 match chunk_result {
11857 Ok(chunk) => {
11858 accumulated.push_str(&chunk.delta);
11859 yield StreamChunk::content(chunk.delta);
11860 }
11861 Err(e) => {
11862 yield StreamChunk::error(e.to_string());
11863 return;
11864 }
11865 }
11866 }
11867 let llm_duration_ms = llm_start.elapsed().as_millis() as u64;
11868 let llm_response = ai_agents_core::LLMResponse::new(
11870 accumulated.trim(),
11871 ai_agents_core::FinishReason::Stop,
11872 );
11873 self.hooks.on_llm_complete(&llm_response, llm_duration_ms).await;
11874 accumulated.trim().to_string()
11875 };
11876
11877 let parsed_tool_calls = match self.parse_main_tool_calls(&content, &protocol) {
11879 Ok(calls) => calls,
11880 Err(error) => {
11881 yield StreamChunk::error(error.to_string());
11882 return;
11883 }
11884 };
11885 if let Some(tool_calls) = parsed_tool_calls {
11886 let native_tool_call = match Self::is_native_tool_call_content(&content) {
11887 Ok(native) => native,
11888 Err(error) => {
11889 yield StreamChunk::error(error.to_string());
11890 return;
11891 }
11892 };
11893 let transition_content = match native_readable_projection(&content) {
11896 Ok(content) => content,
11897 Err(error) => {
11898 yield StreamChunk::error(error.to_string());
11899 return;
11900 }
11901 };
11902 let transition_fired = match self.evaluate_transitions(processed_input, &transition_content).await {
11903 Ok(v) => v,
11904 Err(e) => {
11905 yield StreamChunk::error(e.to_string());
11906 return;
11907 }
11908 };
11909 if transition_fired {
11910 let _ = self.memory.add_message(ChatMessage::assistant(
11911 "(Transitioned to new state — tool call handled by workflow)",
11912 )).await;
11913
11914 if include_state_events
11915 && let Some(state) = self.current_state()
11916 {
11917 yield StreamChunk::state_transition(None, state);
11918 }
11919 continue;
11920 }
11921
11922 if let Err(error) = self.memory.add_message(ChatMessage::assistant(&content)).await {
11924 yield StreamChunk::error(error.to_string());
11925 return;
11926 }
11927 if let Err(error) = self.remember_committed_native_exchange(&content).await {
11928 yield StreamChunk::error(error.to_string());
11929 return;
11930 }
11931
11932 let results = self.execute_tools_parallel(&tool_calls).await;
11934 let mut rejection = None;
11935
11936 for ((_id, result), tool_call) in results.into_iter().zip(tool_calls.iter()) {
11937 if include_tool_events {
11938 yield StreamChunk::tool_start(&tool_call.id, &tool_call.name);
11939 }
11940
11941 match result {
11942 Ok(output) => {
11943 if include_tool_events {
11944 yield StreamChunk::tool_result(
11945 &tool_call.id,
11946 &tool_call.name,
11947 &output,
11948 true,
11949 );
11950 }
11951 let result_message = match Self::tool_result_message(
11952 tool_call,
11953 &output,
11954 native_tool_call,
11955 ) {
11956 Ok(message) => message,
11957 Err(error) => {
11958 yield StreamChunk::error(error.to_string());
11959 return;
11960 }
11961 };
11962 if let Err(error) = self.memory.add_message(result_message).await {
11963 yield StreamChunk::error(error.to_string());
11964 return;
11965 }
11966 }
11967 Err(e) => {
11968 if matches!(e, AgentError::HITLRejected(_)) && !native_tool_call {
11969 let response = AgentResponse {
11970 content: format!("Operation cancelled: {e}"),
11971 metadata: None,
11972 tool_calls: Some(all_tool_calls.clone()),
11973 };
11974 if let Err(error) = self.memory.add_message(ChatMessage::assistant(
11975 format!("The operation was rejected by the approver: {e}"),
11976 )).await {
11977 yield StreamChunk::error(error.to_string());
11978 return;
11979 }
11980 if let Err(finalize_error) = self.finish_turn_if_root(&response).await {
11981 yield StreamChunk::error(finalize_error.to_string());
11982 return;
11983 }
11984 let legacy_error = response.content.clone();
11985 record_runtime_stream_final(&terminal, response);
11986 yield StreamChunk::error(legacy_error);
11987 yield StreamChunk::Done {};
11988 return;
11989 }
11990 if rejection.is_none() && matches!(e, AgentError::HITLRejected(_)) {
11991 rejection = Some(e.to_string());
11992 }
11993 if include_tool_events {
11994 yield StreamChunk::tool_result(
11995 &tool_call.id,
11996 &tool_call.name,
11997 e.to_string(),
11998 false,
11999 );
12000 }
12001 let result_message = match Self::tool_result_message(
12002 tool_call,
12003 &format!("Error: {}", e),
12004 native_tool_call,
12005 ) {
12006 Ok(message) => message,
12007 Err(error) => {
12008 yield StreamChunk::error(error.to_string());
12009 return;
12010 }
12011 };
12012 if let Err(error) = self.memory.add_message(result_message).await {
12013 yield StreamChunk::error(error.to_string());
12014 return;
12015 }
12016 }
12017 }
12018 all_tool_calls.push(tool_call.clone());
12019
12020 if include_tool_events {
12021 yield StreamChunk::tool_end(&tool_call.id);
12022 }
12023 }
12024 if let Some(rejection) = rejection {
12025 if let Err(error) = self.memory.add_message(ChatMessage::assistant(
12026 format!("The operation was rejected by the approver: {rejection}"),
12027 )).await {
12028 yield StreamChunk::error(error.to_string());
12029 return;
12030 }
12031 let response = AgentResponse {
12032 content: format!("Operation cancelled: {rejection}"),
12033 metadata: None,
12034 tool_calls: Some(all_tool_calls.clone()),
12035 };
12036 if let Err(finalize_error) = self.finish_turn_if_root(&response).await {
12037 yield StreamChunk::error(finalize_error.to_string());
12038 return;
12039 }
12040 let legacy_error = response.content.clone();
12041 record_runtime_stream_final(&terminal, response);
12042 yield StreamChunk::error(legacy_error);
12043 yield StreamChunk::Done {};
12044 return;
12045 }
12046 continue;
12047 }
12048
12049 let (extracted_thinking, answer) = self.extract_thinking(&content);
12051 if extracted_thinking.is_some() {
12052 thinking_content = extracted_thinking;
12053 }
12054
12055 let output_data = match self.process_output(&answer, &input_data.context).await {
12056 Ok(d) => d,
12057 Err(e) => {
12058 yield StreamChunk::error(e.to_string());
12059 return;
12060 }
12061 };
12062
12063 let final_content = if output_data.metadata.rejected {
12064 output_data
12065 .metadata
12066 .rejection_reason
12067 .unwrap_or_else(|| answer.to_string())
12068 } else {
12069 output_data.content
12070 };
12071
12072 let (final_content, reflection_metadata) = match self
12074 .run_reflection(&*llm, processed_input, final_content)
12075 .await
12076 {
12077 Ok(r) => r,
12078 Err(e) => {
12079 yield StreamChunk::error(e.to_string());
12080 return;
12081 }
12082 };
12083
12084 let final_content = self.format_response_with_thinking(
12085 thinking_content.as_deref(),
12086 &final_content,
12087 );
12088
12089 if buffered_decision {
12091 yield StreamChunk::content(&final_content);
12092 }
12093
12094 let post_result = match self
12098 .post_loop_processing(processed_input, final_content)
12099 .await
12100 {
12101 Ok(r) => r,
12102 Err(e) => {
12103 yield StreamChunk::error(e.to_string());
12104 return;
12105 }
12106 };
12107
12108 let (final_content, transitioned) = match post_result {
12109 PostLoopResult::NoTransition(content) => (content, false),
12110 PostLoopResult::Transitioned(content) => (content, true),
12111 PostLoopResult::NeedsRedispatch => {
12112 const MAX_REDISPATCH_DEPTH: u32 = 3;
12113 let current_depth = *self.redispatch_depth.read();
12114 let content = if current_depth >= MAX_REDISPATCH_DEPTH {
12115 warn!(
12116 depth = current_depth,
12117 "Post-transition re-dispatch depth limit reached (stream)"
12118 );
12119 let c = String::new();
12120 let _ = self.memory.add_message(ChatMessage::assistant(&c)).await;
12121 c
12122 } else {
12123 *self.redispatch_depth.write() += 1;
12124 if let Some(context) = self.active_turn_context.write().as_mut() {
12125 context.enter_redispatch();
12126 }
12127 info!(
12128 depth = current_depth + 1,
12129 "Re-dispatching for new state after transition (stream)"
12130 );
12131 let result = self.run_loop_internal(processed_input).await;
12132 *self.redispatch_depth.write() -= 1;
12133 if let Some(context) = self.active_turn_context.write().as_mut() {
12134 context.exit_redispatch();
12135 }
12136 match result {
12137 Ok(resp) => resp.content,
12138 Err(e) => {
12139 yield StreamChunk::error(e.to_string());
12140 return;
12141 }
12142 }
12143 };
12144 (content, true)
12145 }
12146 };
12147
12148 if transitioned {
12149 if include_state_events
12150 && let Some(state) = self.current_state()
12151 {
12152 yield StreamChunk::state_transition(None, state);
12153 }
12154 yield StreamChunk::content(&final_content);
12156 }
12157
12158 let final_response = self.build_agent_response(AgentResponseParts {
12160 content: final_content,
12161 all_tool_calls,
12162 reasoning_mode,
12163 auto_detected,
12164 iterations,
12165 thinking: thinking_content,
12166 reflection_metadata,
12167 });
12168 if let Err(e) = self.finish_turn_if_root(&final_response).await {
12169 yield StreamChunk::error(e.to_string());
12170 return;
12171 }
12172
12173 record_runtime_stream_final(&terminal, final_response);
12174 yield StreamChunk::Done {};
12175 return;
12176 }
12177 })
12178 }
12179
12180 fn run_loop_stream<'a>(
12183 &'a self,
12184 input: &'a str,
12185 terminal: RuntimeStreamTerminalSlot,
12186 ) -> Pin<Box<dyn Stream<Item = StreamChunk> + Send + 'a>> {
12187 Box::pin(async_stream::stream! {
12188 self.begin_root_turn();
12189 let _root_cleanup = RootTurnCleanup::new(self);
12190 self.hooks.on_message_received(input).await;
12191
12192 if !self.context_initialized.swap(true, Ordering::SeqCst) {
12194 if let Err(e) = self.context_manager.initialize().await {
12195 yield StreamChunk::error(e.to_string());
12196 return;
12197 }
12198 debug!("Context manager initialized (defaults, env, builtins)");
12199 }
12200
12201 if let Err(e) = self.check_turn_timeout().await {
12202 yield StreamChunk::error(e.to_string());
12203 return;
12204 }
12205 if let Err(e) = self.context_manager.refresh_per_turn().await {
12206 yield StreamChunk::error(e.to_string());
12207 return;
12208 }
12209
12210 self.clear_disambiguation_context();
12212
12213 if let Some(ref disambiguator) = self.disambiguation_manager {
12215 let disambiguation_context = match self.build_disambiguation_context().await {
12216 Ok(ctx) => ctx,
12217 Err(e) => {
12218 yield StreamChunk::error(e.to_string());
12219 return;
12220 }
12221 };
12222
12223 let state_override = self
12224 .state_machine
12225 .as_ref()
12226 .and_then(|sm| sm.current_definition())
12227 .and_then(|def| def.disambiguation.clone());
12228
12229 let state_generation = self
12230 .state_machine
12231 .as_ref()
12232 .map(|state_machine| state_machine.generation());
12233 let disambiguation_epoch = self.disambiguation_epoch.load(Ordering::SeqCst);
12234 let mut result = match self
12235 .observe_purpose(
12236 ObservationPurpose::DisambiguationDetection,
12237 disambiguator.process_input_with_override(
12238 input,
12239 &disambiguation_context,
12240 state_override.as_ref(),
12241 None,
12242 ),
12243 )
12244 .await
12245 {
12246 Ok(r) => r,
12247 Err(e) => {
12248 yield StreamChunk::error(e.to_string());
12249 return;
12250 }
12251 };
12252 let current_state_generation = self
12253 .state_machine
12254 .as_ref()
12255 .map(|state_machine| state_machine.generation());
12256 if current_state_generation != state_generation
12257 || self.disambiguation_epoch.load(Ordering::SeqCst) != disambiguation_epoch
12258 {
12259 disambiguator.clear_pending().await;
12260 *self.pending_skill_id.write() = None;
12261 result = DisambiguationResult::Abandoned { new_input: None };
12262 info!(
12263 confirmation_event = "invalidated",
12264 invalidation_reason = "state_generation_changed",
12265 "Streaming disambiguation result invalidated before redispatch"
12266 );
12267 }
12268 match result {
12269 DisambiguationResult::Clear => {
12270 debug!("Input is clear, proceeding normally (stream)");
12271 }
12272 DisambiguationResult::NeedsClarification {
12273 question,
12274 detection,
12275 } => {
12276 let admission = match self
12277 .admit_disambiguation_redispatch(
12278 disambiguation_epoch,
12279 state_generation,
12280 )
12281 .await
12282 {
12283 Ok(admission) => admission,
12284 Err(error) => {
12285 *self.pending_skill_id.write() = None;
12286 yield StreamChunk::error(error.to_string());
12287 return;
12288 }
12289 };
12290 let awaiting_confirmation = disambiguator.has_pending_confirmation().await;
12291 info!(
12292 ambiguity_type = ?detection.ambiguity_type,
12293 confidence = detection.confidence,
12294 "Input requires clarification (stream)"
12295 );
12296 if let Err(e) = self.commit_root_user_message(input).await {
12299 yield StreamChunk::error(e.to_string());
12300 return;
12301 }
12302 let _ = self
12303 .memory
12304 .add_message(ChatMessage::assistant(&question.question))
12305 .await;
12306 let status = if awaiting_confirmation {
12307 "awaiting_confirmation"
12308 } else {
12309 "awaiting_clarification"
12310 };
12311 let response = AgentResponse::new(&question.question).with_metadata(
12312 "disambiguation",
12313 serde_json::json!({ "status": status }),
12314 );
12315 drop(admission);
12316 if let Err(e) = self.finish_turn_if_root(&response).await {
12317 yield StreamChunk::error(e.to_string());
12318 return;
12319 }
12320 yield StreamChunk::content(&question.question);
12321 record_runtime_stream_final(&terminal, response);
12322 yield StreamChunk::Done {};
12323 return;
12324 }
12325 DisambiguationResult::Clarified {
12326 enriched_input,
12327 resolved,
12328 ..
12329 } => {
12330 let admission = match self
12331 .admit_disambiguation_redispatch(
12332 disambiguation_epoch,
12333 state_generation,
12334 )
12335 .await
12336 {
12337 Ok(admission) => admission,
12338 Err(error) => {
12339 *self.pending_skill_id.write() = None;
12340 yield StreamChunk::error(error.to_string());
12341 return;
12342 }
12343 };
12344 info!(
12345 resolved_count = resolved.len(),
12346 enriched = %enriched_input,
12347 "Input clarified (stream)"
12348 );
12349 for (key, value) in &resolved {
12350 let context_key = format!("disambiguation.{}", key);
12351 let _ = self.context_manager.set(&context_key, value.clone());
12352 }
12353 if let Some(intent) = resolved.get("intent") {
12354 let _ = self.context_manager.set("resolved_intent", intent.clone());
12355 }
12356 let _ = self
12357 .context_manager
12358 .set("disambiguation.resolved", serde_json::Value::Bool(true));
12359
12360 let skill_id = self.pending_skill_id.read().clone();
12364 if let Some(skill_id) = skill_id {
12365 info!(skill_id = %skill_id, "Re-checking skill disambiguation on clarified input (stream)");
12366 drop(admission);
12367 match self
12368 .recheck_skill_disambiguation(
12369 &skill_id,
12370 &enriched_input,
12371 disambiguation_epoch,
12372 state_generation,
12373 )
12374 .await
12375 {
12376 Ok(resp) => {
12377 yield StreamChunk::content(&resp.content);
12378 record_runtime_stream_final(&terminal, resp);
12379 yield StreamChunk::Done {};
12380 return;
12381 }
12382 Err(e) => {
12383 yield StreamChunk::error(e.to_string());
12384 return;
12385 }
12386 }
12387 }
12388
12389 drop(admission);
12391 let mut inner = self.run_loop_internal_stream(
12392 &enriched_input,
12393 Arc::clone(&terminal),
12394 );
12395 while let Some(chunk) = inner.next().await {
12396 yield chunk;
12397 }
12398 return;
12399 }
12400 DisambiguationResult::ProceedWithBestGuess { enriched_input } => {
12401 info!("Proceeding with best guess (stream)");
12402
12403 let skill_id = self.pending_skill_id.read().clone();
12405 if let Some(skill_id) = skill_id {
12406 info!(skill_id = %skill_id, "Re-checking skill disambiguation on best-guess input (stream)");
12407 match self
12408 .recheck_skill_disambiguation(
12409 &skill_id,
12410 &enriched_input,
12411 disambiguation_epoch,
12412 state_generation,
12413 )
12414 .await
12415 {
12416 Ok(resp) => {
12417 yield StreamChunk::content(&resp.content);
12418 record_runtime_stream_final(&terminal, resp);
12419 yield StreamChunk::Done {};
12420 return;
12421 }
12422 Err(e) => {
12423 yield StreamChunk::error(e.to_string());
12424 return;
12425 }
12426 }
12427 }
12428
12429 let mut inner = self.run_loop_internal_stream(
12430 &enriched_input,
12431 Arc::clone(&terminal),
12432 );
12433 while let Some(chunk) = inner.next().await {
12434 yield chunk;
12435 }
12436 return;
12437 }
12438 DisambiguationResult::GiveUp { reason } => {
12439 *self.pending_skill_id.write() = None;
12440 warn!(reason = %reason, "Disambiguation gave up (stream)");
12441 let apology = self
12442 .generate_localized_apology(
12443 "Generate a brief, polite apology saying you couldn't understand the request. Be concise.",
12444 &reason,
12445 )
12446 .await
12447 .unwrap_or_else(|_| {
12448 format!("I'm sorry, I couldn't understand your request: {}", reason)
12449 });
12450 let response = AgentResponse::new(&apology);
12451 if let Err(e) = self.finish_turn_if_root(&response).await {
12452 yield StreamChunk::error(e.to_string());
12453 return;
12454 }
12455 yield StreamChunk::content(&apology);
12456 record_runtime_stream_final(&terminal, response);
12457 yield StreamChunk::Done {};
12458 return;
12459 }
12460 DisambiguationResult::Escalate { reason } => {
12461 *self.pending_skill_id.write() = None;
12462 info!(reason = %reason, "Escalating to human (stream)");
12463 if let Some(ref hitl) = self.hitl_engine {
12464 let trigger =
12465 ApprovalTrigger::condition("disambiguation_escalation", reason.clone());
12466 let mut context_map = HashMap::new();
12467 context_map.insert("original_input".to_string(), serde_json::json!(input));
12468 context_map.insert("reason".to_string(), serde_json::json!(&reason));
12469 let check_result = HITLCheckResult::required(
12470 trigger,
12471 context_map,
12472 format!("User request needs human assistance: {}", reason),
12473 Some(hitl.config().default_timeout_seconds),
12474 );
12475 match self.request_hitl_approval(check_result).await {
12476 Ok(ApprovalResult::Approved | ApprovalResult::Modified { .. }) => {
12477 let mut inner = self.run_loop_internal_stream(
12478 input,
12479 Arc::clone(&terminal),
12480 );
12481 while let Some(chunk) = inner.next().await {
12482 yield chunk;
12483 }
12484 return;
12485 }
12486 Ok(_) => {}
12487 Err(e) => {
12488 yield StreamChunk::error(e.to_string());
12489 return;
12490 }
12491 }
12492 }
12493 let apology = self
12494 .generate_localized_apology(
12495 "Explain briefly that you're transferring the user to a human agent for help.",
12496 &reason,
12497 )
12498 .await
12499 .unwrap_or_else(|_| {
12500 format!("I need human assistance to help with your request: {}", reason)
12501 });
12502 let response = AgentResponse::new(&apology);
12503 if let Err(e) = self.finish_turn_if_root(&response).await {
12504 yield StreamChunk::error(e.to_string());
12505 return;
12506 }
12507 yield StreamChunk::content(&apology);
12508 record_runtime_stream_final(&terminal, response);
12509 yield StreamChunk::Done {};
12510 return;
12511 }
12512 DisambiguationResult::Abandoned { new_input } => {
12513 *self.pending_skill_id.write() = None;
12514
12515 info!(
12516 has_new_input = new_input.is_some(),
12517 "Clarification abandoned by user (stream)"
12518 );
12519
12520 if let Err(e) = self.commit_root_user_message(input).await {
12521 yield StreamChunk::error(e.to_string());
12522 return;
12523 }
12524
12525 match new_input {
12526 Some(fresh_input) => {
12527 let mut inner = self.run_loop_internal_stream(
12529 &fresh_input,
12530 Arc::clone(&terminal),
12531 );
12532 while let Some(chunk) = inner.next().await {
12533 yield chunk;
12534 }
12535 return;
12536 }
12537 None => {
12538 let ack = self
12540 .generate_localized_apology(
12541 "The user changed their mind about their previous request. \
12542 Generate a brief, friendly acknowledgment (e.g. 'OK, no problem. What else can I help with?'). \
12543 Do NOT apologize excessively. Be concise.",
12544 "User abandoned clarification",
12545 )
12546 .await
12547 .unwrap_or_else(|_| {
12548 "OK, no problem. What else can I help with?".to_string()
12549 });
12550
12551 let _ = self
12552 .memory
12553 .add_message(ChatMessage::assistant(&ack))
12554 .await;
12555
12556 let response = AgentResponse::new(&ack);
12557 if let Err(e) = self.finish_turn_if_root(&response).await {
12558 yield StreamChunk::error(e.to_string());
12559 return;
12560 }
12561 yield StreamChunk::content(&ack);
12562 record_runtime_stream_final(&terminal, response);
12563 yield StreamChunk::Done {};
12564 return;
12565 }
12566 }
12567 }
12568 }
12569 }
12570
12571 let mut inner = self.run_loop_internal_stream(input, Arc::clone(&terminal));
12573 while let Some(chunk) = inner.next().await {
12574 yield chunk;
12575 }
12576 })
12577 }
12578
12579 pub fn info(&self) -> AgentInfo {
12580 self.info.clone()
12581 }
12582
12583 pub fn skills(&self) -> &[SkillDefinition] {
12584 &self.skills
12585 }
12586
12587 async fn reset_runtime_state(&self) -> Result<()> {
12589 let _admission = self.disambiguation_admission.write().await;
12590 if self.state_transition_reserved.load(Ordering::SeqCst) {
12591 return Err(AgentError::Other(
12592 "Cannot reset while a state transition is in progress".to_string(),
12593 ));
12594 }
12595 self.disambiguation_epoch.fetch_add(1, Ordering::SeqCst);
12596 *self.pending_skill_id.write() = None;
12597 if let Some(disambiguator) = self.disambiguation_manager.as_ref() {
12598 disambiguator.clear_pending().await;
12599 }
12600 self.memory.clear().await?;
12601 self.active_native_exchanges.write().clear();
12602 *self.iteration_count.write() = 0;
12603 self.tool_call_history.write().clear();
12604 if let Some(ref sm) = self.state_machine {
12605 sm.reset();
12606 }
12607 Ok(())
12608 }
12609
12610 pub async fn reset(&self) -> Result<()> {
12612 self.reset_runtime_state().await
12613 }
12614
12615 pub fn max_context_tokens(&self) -> u32 {
12616 self.max_context_tokens
12617 }
12618
12619 pub fn llm_registry(&self) -> &Arc<LLMRegistry> {
12620 &self.llm_registry
12621 }
12622
12623 pub fn state_machine(&self) -> Option<&Arc<StateMachine>> {
12624 self.state_machine.as_ref()
12625 }
12626
12627 pub fn context_manager(&self) -> &Arc<ContextManager> {
12628 &self.context_manager
12629 }
12630
12631 pub fn tool_call_history(&self) -> Vec<ToolCallRecord> {
12632 self.tool_call_history.read().clone()
12633 }
12634
12635 pub fn memory_token_budget(&self) -> Option<&MemoryTokenBudget> {
12636 self.memory_token_budget.as_ref()
12637 }
12638
12639 pub fn parallel_tools_config(&self) -> &ParallelToolsConfig {
12640 &self.parallel_tools
12641 }
12642
12643 pub fn streaming_config(&self) -> &StreamingConfig {
12644 &self.streaming
12645 }
12646
12647 pub fn hooks(&self) -> &Arc<dyn AgentHooks> {
12648 &self.hooks
12649 }
12650
12651 pub fn hitl_engine(&self) -> Option<&HITLEngine> {
12652 self.hitl_engine.as_ref()
12653 }
12654
12655 pub fn approval_handler(&self) -> &Arc<dyn ApprovalHandler> {
12656 &self.approval_handler
12657 }
12658
12659 fn build_hitl_language_context(&self) -> HashMap<String, Value> {
12661 let mut ctx = HashMap::new();
12662 for key in &["user.language", "input.detected.language", "language"] {
12663 if let Some(val) = self.context_manager.get(key) {
12664 ctx.insert(key.to_string(), val);
12665 }
12666 }
12667 ctx
12668 }
12669
12670 async fn request_hitl_approval(&self, check_result: HITLCheckResult) -> Result<ApprovalResult> {
12672 let Some(request) = check_result.into_request() else {
12673 return Ok(ApprovalResult::Approved);
12674 };
12675
12676 self.hooks.on_approval_requested(&request).await;
12677
12678 let timeout = request.timeout;
12679
12680 let raw_result = if let Some(duration) = timeout {
12681 match tokio::time::timeout(
12682 duration,
12683 self.approval_handler.request_approval(request.clone()),
12684 )
12685 .await
12686 {
12687 Ok(result) => result,
12688 Err(_) => ApprovalResult::timeout(),
12689 }
12690 } else {
12691 self.approval_handler
12692 .request_approval(request.clone())
12693 .await
12694 };
12695
12696 self.hooks
12697 .on_approval_result(&request.id, &raw_result)
12698 .await;
12699
12700 let (outcome, effective_result): (ApprovalResolvedOutcome, Result<ApprovalResult>) =
12701 match &raw_result {
12702 ApprovalResult::Approved => (
12703 ApprovalResolvedOutcome::Approved,
12704 Ok(ApprovalResult::Approved),
12705 ),
12706 ApprovalResult::Rejected { reason } => (
12707 ApprovalResolvedOutcome::Rejected {
12708 reason: reason.clone(),
12709 },
12710 Ok(ApprovalResult::Rejected {
12711 reason: reason.clone(),
12712 }),
12713 ),
12714 ApprovalResult::Modified { changes } => (
12715 ApprovalResolvedOutcome::Modified {
12716 changes: changes.clone(),
12717 },
12718 Ok(ApprovalResult::Modified {
12719 changes: changes.clone(),
12720 }),
12721 ),
12722 ApprovalResult::Timeout => {
12723 if let Some(ref engine) = self.hitl_engine {
12724 match engine.config().on_timeout {
12725 TimeoutAction::Approve => (
12726 ApprovalResolvedOutcome::Approved,
12727 Ok(ApprovalResult::Approved),
12728 ),
12729 TimeoutAction::Reject => {
12730 let reason = Some("Timeout".to_string());
12731 (
12732 ApprovalResolvedOutcome::Rejected {
12733 reason: reason.clone(),
12734 },
12735 Ok(ApprovalResult::Rejected { reason }),
12736 )
12737 }
12738 TimeoutAction::Error => {
12739 let message = "HITL approval timeout".to_string();
12740 (
12741 ApprovalResolvedOutcome::Error {
12742 message: message.clone(),
12743 },
12744 Err(AgentError::Other(message)),
12745 )
12746 }
12747 }
12748 } else {
12749 let reason = Some("Timeout (no engine)".to_string());
12750 (
12751 ApprovalResolvedOutcome::Rejected {
12752 reason: reason.clone(),
12753 },
12754 Ok(ApprovalResult::Rejected { reason }),
12755 )
12756 }
12757 }
12758 };
12759
12760 self.hooks
12761 .on_approval_resolved(&request, &raw_result, &outcome)
12762 .await;
12763
12764 effective_result
12765 }
12766
12767 pub async fn check_state_hitl(&self, from: Option<&str>, to: &str) -> Result<bool> {
12768 if let Some(ref hitl_engine) = self.hitl_engine {
12769 let hitl_lang_ctx = self.build_hitl_language_context();
12770 let check_result = self
12771 .observe_purpose(
12772 ObservationPurpose::HitlLocalization,
12773 hitl_engine.check_state_transition_with_localization(
12774 from,
12775 to,
12776 &hitl_lang_ctx,
12777 self.approval_handler.as_ref(),
12778 Some(&self.llm_registry),
12779 ),
12780 )
12781 .await?;
12782 if check_result.is_required() {
12783 let result = self.request_hitl_approval(check_result).await?;
12784 return Ok(matches!(
12785 result,
12786 ApprovalResult::Approved | ApprovalResult::Modified { .. }
12787 ));
12788 }
12789 }
12790 Ok(true)
12791 }
12792
12793 async fn execute_tools_parallel(
12795 &self,
12796 tool_calls: &[ToolCall],
12797 ) -> Vec<(String, Result<String>)> {
12798 let can_run_parallel = tool_calls.iter().all(|tc| {
12799 self.tools
12800 .resolve(&tc.name)
12801 .map(|resolved| resolved.tool.classify_call(&tc.arguments).concurrency_safe)
12802 .unwrap_or(false)
12803 });
12804
12805 if !self.parallel_tools.enabled || tool_calls.len() <= 1 || !can_run_parallel {
12806 let mut results = Vec::new();
12807 for tc in tool_calls {
12808 let result = self
12809 .observe_purpose(
12810 current_observation_context()
12811 .map(|context| context.purpose)
12812 .unwrap_or_default(),
12813 self.execute_tool_smart(tc),
12814 )
12815 .await;
12816 results.push((tc.id.clone(), result));
12817 }
12818 return results;
12819 }
12820
12821 let chunks: Vec<_> = tool_calls
12822 .chunks(self.parallel_tools.max_parallel)
12823 .collect();
12824
12825 let mut all_results = Vec::new();
12826
12827 for chunk in chunks {
12828 let futures: Vec<_> = chunk
12829 .iter()
12830 .map(|tc| {
12831 let tc = tc.clone();
12832 async move {
12833 let result = self.execute_tool_smart(&tc).await;
12834 (tc.id.clone(), result)
12835 }
12836 })
12837 .collect();
12838
12839 let results = futures::future::join_all(futures).await;
12840 all_results.extend(results);
12841 }
12842
12843 all_results
12844 }
12845
12846 pub async fn chat_stream<'a>(
12850 &'a self,
12851 input: &'a str,
12852 ) -> Result<Pin<Box<dyn Stream<Item = StreamChunk> + Send + 'a>>> {
12853 let RootTurnAdmission {
12854 guard: root_turn_guard,
12855 identity_stack,
12856 } = self.acquire_root_turn().await?;
12857 scope_runtime_gate_identity_stack(&identity_stack, self.init_storage()).await?;
12861 info!(input_len = input.len(), "Starting streaming chat");
12862 let terminal = new_runtime_stream_terminal_slot();
12863 let inner = self.run_loop_stream(input, terminal);
12864 let observation_context = self.build_observation_context(None);
12865 let stream: Pin<Box<dyn Stream<Item = StreamChunk> + Send + 'a>> =
12866 Box::pin(async_stream::stream! {
12867 let mut root_turn_guard = Some(root_turn_guard);
12868 let mut inner = inner;
12869 loop {
12870 let next = scope_runtime_gate_identity_stack(&identity_stack, async {
12871 if let Some(context) = observation_context.as_ref() {
12872 with_observation_context(context.clone(), inner.next()).await
12873 } else {
12874 inner.next().await
12875 }
12876 })
12877 .await;
12878 match next {
12879 Some(StreamChunk::Done {}) => {
12880 while scope_runtime_gate_identity_stack(&identity_stack, async {
12881 if let Some(context) = observation_context.as_ref() {
12882 with_observation_context(context.clone(), inner.next())
12883 .await
12884 .is_some()
12885 } else {
12886 inner.next().await.is_some()
12887 }
12888 })
12889 .await
12890 {}
12891 if observation_context.is_some() {
12892 scope_runtime_gate_identity_stack(
12893 &identity_stack,
12894 self.export_observability_if_configured(),
12895 )
12896 .await;
12897 }
12898 drop(root_turn_guard.take());
12899 yield StreamChunk::Done {};
12900 return;
12901 }
12902 Some(chunk) => yield chunk,
12903 None => {
12904 if observation_context.is_some() {
12905 scope_runtime_gate_identity_stack(
12906 &identity_stack,
12907 self.export_observability_if_configured(),
12908 )
12909 .await;
12910 }
12911 drop(root_turn_guard.take());
12912 return;
12913 }
12914 }
12915 }
12916 });
12917 Ok(stream)
12918 }
12919
12920 pub async fn chat_stream_events<'a>(
12924 &'a self,
12925 input: &'a str,
12926 ) -> Result<Pin<Box<dyn Stream<Item = AgentStreamEvent> + Send + 'a>>> {
12927 let RootTurnAdmission {
12928 guard: root_turn_guard,
12929 identity_stack,
12930 } = self.acquire_root_turn().await?;
12931 scope_runtime_gate_identity_stack(&identity_stack, self.init_storage()).await?;
12935 info!(input_len = input.len(), "Starting streaming chat events");
12936 let terminal = new_runtime_stream_terminal_slot();
12937 let mut inner = self.run_loop_stream(input, Arc::clone(&terminal));
12938 let observation_context = self.build_observation_context(None);
12939 let stream: Pin<Box<dyn Stream<Item = AgentStreamEvent> + Send + 'a>> =
12940 Box::pin(async_stream::stream! {
12941 let mut root_turn_guard = Some(root_turn_guard);
12942 loop {
12943 let next = scope_runtime_gate_identity_stack(&identity_stack, async {
12944 if let Some(context) = observation_context.as_ref() {
12945 with_observation_context(context.clone(), inner.next()).await
12946 } else {
12947 inner.next().await
12948 }
12949 })
12950 .await;
12951 match next {
12952 Some(StreamChunk::Done {}) => {
12953 let terminal_event = { terminal.write().take() };
12954 if let Some(response) = terminal_event {
12955 while scope_runtime_gate_identity_stack(&identity_stack, async {
12956 if let Some(context) = observation_context.as_ref() {
12957 with_observation_context(context.clone(), inner.next())
12958 .await
12959 .is_some()
12960 } else {
12961 inner.next().await.is_some()
12962 }
12963 })
12964 .await
12965 {}
12966 if observation_context.is_some() {
12967 scope_runtime_gate_identity_stack(
12968 &identity_stack,
12969 self.export_observability_if_configured(),
12970 )
12971 .await;
12972 }
12973 drop(root_turn_guard.take());
12974 yield AgentStreamEvent::Final(response);
12975 return;
12976 }
12977 }
12978 Some(StreamChunk::Error { message }) => {
12979 let finalized = { terminal.read().is_some() };
12980 if finalized {
12981 continue;
12982 }
12983 while scope_runtime_gate_identity_stack(&identity_stack, async {
12984 if let Some(context) = observation_context.as_ref() {
12985 with_observation_context(context.clone(), inner.next())
12986 .await
12987 .is_some()
12988 } else {
12989 inner.next().await.is_some()
12990 }
12991 })
12992 .await
12993 {}
12994 if observation_context.is_some() {
12995 scope_runtime_gate_identity_stack(
12996 &identity_stack,
12997 self.export_observability_if_configured(),
12998 )
12999 .await;
13000 }
13001 drop(root_turn_guard.take());
13002 yield AgentStreamEvent::Chunk(StreamChunk::Error { message });
13003 return;
13004 }
13005 Some(chunk) => yield AgentStreamEvent::Chunk(chunk),
13006 None => {
13007 if observation_context.is_some() {
13008 scope_runtime_gate_identity_stack(
13009 &identity_stack,
13010 self.export_observability_if_configured(),
13011 )
13012 .await;
13013 }
13014 drop(root_turn_guard.take());
13015 return;
13016 }
13017 }
13018 }
13019 });
13020 Ok(stream)
13021 }
13022}
13023
13024#[async_trait]
13025impl ToolInvoker for RuntimeAgent {
13026 async fn invoke_tool(&self, request: ToolExecutionRequest) -> Result<ToolExecutionRecord> {
13027 self.execute_tool_record(request).await
13028 }
13029}
13030
13031#[async_trait]
13032impl Agent for RuntimeAgent {
13033 async fn chat(&self, input: &str) -> Result<AgentResponse> {
13035 let RootTurnAdmission {
13036 guard,
13037 identity_stack,
13038 } = self.acquire_root_turn().await?;
13039 let result = scope_runtime_gate_identity_stack(&identity_stack, async {
13040 let result = if let Some(context) = self.build_observation_context(None) {
13041 with_observation_context(context, self.run_loop(input)).await
13042 } else {
13043 self.run_loop(input).await
13044 };
13045 self.export_observability_if_configured().await;
13046 result
13047 })
13048 .await;
13049 drop(guard);
13050 result
13051 }
13052
13053 fn info(&self) -> AgentInfo {
13054 self.info.clone()
13055 }
13056
13057 async fn reset(&self) -> Result<()> {
13059 self.reset_runtime_state().await
13060 }
13061}
13062
13063fn background_maintenance_tags(
13073 label: &str,
13074 stage: &str,
13075 reason: Option<&str>,
13076 policy: Option<&crate::optimization::config::MaintenanceTaskPolicy>,
13077) -> HashMap<String, String> {
13078 let mut tags = HashMap::new();
13079 tags.insert("runtime.background".to_string(), "true".to_string());
13080 tags.insert("runtime.maintenance".to_string(), label.to_string());
13081 tags.insert("runtime.maintenance_stage".to_string(), stage.to_string());
13082 if let Some(policy) = policy {
13083 tags.insert(
13084 "runtime.await_before_next_turn".to_string(),
13085 await_before_next_turn_label(policy.await_before_next_turn).to_string(),
13086 );
13087 tags.insert(
13088 "runtime.maintenance_mode".to_string(),
13089 maintenance_mode_label(policy.mode).to_string(),
13090 );
13091 }
13092 if let Some(reason) = reason {
13093 tags.insert("runtime.reason".to_string(), reason.to_string());
13094 }
13095 tags
13096}
13097
13098fn await_before_next_turn_label(policy: AwaitBeforeNextTurn) -> &'static str {
13099 match policy {
13100 AwaitBeforeNextTurn::Never => "never",
13101 AwaitBeforeNextTurn::SameActor => "same_actor",
13102 AwaitBeforeNextTurn::Always => "always",
13103 }
13104}
13105
13106fn maintenance_mode_label(mode: MaintenanceMode) -> &'static str {
13107 match mode {
13108 MaintenanceMode::InlineSerial => "inline_serial",
13109 MaintenanceMode::InlineParallel => "inline_parallel",
13110 MaintenanceMode::Background => "background",
13111 }
13112}
13113
13114fn record_background_maintenance_event(
13116 manager: Option<&Arc<ObservabilityManager>>,
13117 label: &str,
13118 status: EventStatus,
13119 duration_ms: u64,
13120 stage: &str,
13121 reason: Option<String>,
13122 policy: Option<&crate::optimization::config::MaintenanceTaskPolicy>,
13123) {
13124 if let Some(manager) = manager {
13125 manager.record_lifecycle_event(
13126 EventType::MemoryOperation {
13127 operation: format!("{}_background_{}", label, stage),
13128 },
13129 ObservationPurpose::Other(format!("{}_maintenance", label)),
13130 status,
13131 duration_ms,
13132 background_maintenance_tags(label, stage, reason.as_deref(), policy),
13133 None,
13134 );
13135 }
13136}
13137
13138fn effective_maintenance_mode(mode: MaintenanceMode, force_parallel: bool) -> MaintenanceMode {
13139 if force_parallel && matches!(mode, MaintenanceMode::InlineSerial) {
13140 MaintenanceMode::InlineParallel
13141 } else {
13142 mode
13143 }
13144}
13145
13146fn observation_purpose_for_process(hint: ProcessPurposeHint) -> ObservationPurpose {
13147 match hint {
13148 ProcessPurposeHint::Detect => ObservationPurpose::ProcessDetect,
13149 ProcessPurposeHint::Extract => ObservationPurpose::ProcessExtract,
13150 ProcessPurposeHint::Validate => ObservationPurpose::ProcessValidate,
13151 ProcessPurposeHint::Transform | ProcessPurposeHint::Other => {
13152 ObservationPurpose::ProcessTransform
13153 }
13154 }
13155}
13156
13157fn new_tool_resource_locks() -> ToolResourceLocks {
13158 Arc::new(RwLock::new(HashMap::new()))
13159}
13160
13161fn tool_resource_lock_keys(
13166 _canonical_id: &str,
13167 args: &Value,
13168 bindings: &ai_agents_core::ToolPolicyBindings,
13169 classification: &ai_agents_core::ToolCallClassification,
13170) -> Vec<String> {
13171 if classification.concurrency_safe {
13172 return Vec::new();
13173 }
13174
13175 let mut keys = Vec::new();
13176 let mut has_path_resource = false;
13177 for binding in &bindings.path_fields {
13178 let value = value_at_argument_path(args, &binding.field)
13179 .cloned()
13180 .or_else(|| {
13181 binding
13182 .default_path
13183 .as_ref()
13184 .map(|path| Value::String(path.clone()))
13185 });
13186 if let Some(value) = value {
13187 collect_resource_strings(&value, |_| {
13188 has_path_resource = true;
13189 });
13190 }
13191 }
13192 for binding in &bindings.domain_fields {
13193 if let Some(value) = value_at_argument_path(args, &binding.field) {
13194 collect_resource_strings(value, |domain| {
13195 let normalized = if binding.is_url {
13196 normalized_url_resource_key(domain)
13197 } else {
13198 domain.trim().trim_end_matches('.').to_ascii_lowercase()
13199 };
13200 keys.push(format!("domain:{}", normalized));
13201 });
13202 }
13203 }
13204 for binding in &bindings.command_fields {
13205 if !matches!(binding.kind, ai_agents_core::CommandBindingKind::Cwd) {
13206 continue;
13207 }
13208 if let Some(value) = value_at_argument_path(args, &binding.field) {
13209 collect_resource_strings(value, |_| {
13210 has_path_resource = true;
13211 });
13212 }
13213 }
13214 if has_path_resource {
13215 keys.push("path-mutation:global".to_string());
13216 }
13217 if keys.is_empty() {
13218 keys.push("side-effect:unbound".to_string());
13219 }
13220 keys.sort();
13221 keys.dedup();
13222 keys
13223}
13224
13225fn value_at_argument_path<'a>(value: &'a Value, field: &str) -> Option<&'a Value> {
13226 let mut current = value;
13227 for segment in field.split('.') {
13228 if segment.is_empty() {
13229 return None;
13230 }
13231 current = current.get(segment)?;
13232 }
13233 Some(current)
13234}
13235
13236fn collect_resource_strings(value: &Value, mut collect: impl FnMut(&str)) {
13237 match value {
13238 Value::String(value) => collect(value),
13239 Value::Array(values) => {
13240 for value in values {
13241 if let Some(value) = value.as_str() {
13242 collect(value);
13243 }
13244 }
13245 }
13246 _ => {}
13247 }
13248}
13249
13250fn normalized_url_resource_key(value: &str) -> String {
13251 let value = value.trim();
13252 let Some((scheme, remainder)) = value.split_once("://") else {
13253 return value.to_ascii_lowercase();
13254 };
13255 let authority_end = remainder.find(['/', '?', '#']).unwrap_or(remainder.len());
13256 let (authority, suffix) = remainder.split_at(authority_end);
13257 format!(
13258 "{}://{}{}",
13259 scheme.to_ascii_lowercase(),
13260 authority.to_ascii_lowercase(),
13261 suffix
13262 )
13263}
13264
13265fn render_concurrent_template(
13266 template: &str,
13267 user_input: &str,
13268 context_values: &std::collections::HashMap<String, serde_json::Value>,
13269) -> Result<String> {
13270 let mut env = minijinja::Environment::new();
13271 env.add_template("concurrent", template)
13272 .map_err(|e| AgentError::Other(format!("Concurrent template parse error: {}", e)))?;
13273
13274 let mut ctx = std::collections::BTreeMap::new();
13275 ctx.insert("user_input".to_string(), minijinja::Value::from(user_input));
13276
13277 let context_obj = minijinja::Value::from_serialize(context_values);
13279 ctx.insert("context".to_string(), context_obj);
13280
13281 let tmpl = env
13282 .get_template("concurrent")
13283 .map_err(|e| AgentError::Other(format!("Concurrent template error: {}", e)))?;
13284
13285 tmpl.render(minijinja::Value::from_serialize(&ctx))
13286 .map_err(|e| AgentError::Other(format!("Concurrent template render error: {}", e)))
13287}
13288
13289#[cfg(test)]
13290mod tests {
13291 use super::*;
13292 use crate::AgentBuilder;
13293 use ai_agents_core::{LLMChunk, LLMConfig, LLMError, LLMFeature, Tool};
13294 use ai_agents_llm::mock::MockLLMProvider;
13295 use ai_agents_skills::{SkillDefinition, SkillStep};
13296 use ai_agents_tools::{
13297 CalculatorTool, CopyPathTool, DeletePathTool, FileWriteTool, MovePathTool, ToolAliases,
13298 ToolDescriptor, ToolProvider, ToolProviderError, ToolProviderType, WebFetchResolver,
13299 WebFetchTool, WebFetchTransport, WebFetchTransportRequest, WebFetchTransportResponse,
13300 };
13301
13302 fn mock_with_response(response: &str) -> MockLLMProvider {
13303 let mut mock = MockLLMProvider::new("test");
13304 mock.set_response(response);
13305 mock
13306 }
13307
13308 fn mock_with_responses(responses: Vec<&str>) -> MockLLMProvider {
13309 let mut mock = MockLLMProvider::new("test");
13310 mock.set_responses(responses.into_iter().map(String::from).collect(), true);
13311 mock
13312 }
13313
13314 fn signed_calculator_response(
13315 exchange_id: &str,
13316 call_id: &str,
13317 expression: &str,
13318 ) -> LLMResponse {
13319 let call = ToolCall {
13320 id: call_id.to_string(),
13321 name: "calculator".to_string(),
13322 arguments: serde_json::json!({"expression": expression}),
13323 };
13324 let state = ai_agents_core::NativeProviderState::new(
13325 exchange_id,
13326 "fixture",
13327 "native-tools",
13328 ai_agents_core::NativeProviderTarget::new("https://fixture.invalid/", "fixture-model")
13329 .unwrap(),
13330 serde_json::json!({
13331 "role": "model",
13332 "parts": [{
13333 "functionCall": {"name": "calculator", "args": {"expression": expression}},
13334 "thoughtSignature": format!("signature-{exchange_id}")
13335 }]
13336 }),
13337 vec![ai_agents_core::NativeCallBinding::new(call_id, 0).unwrap()],
13338 )
13339 .unwrap();
13340 LLMResponse::new("", FinishReason::ToolCall)
13341 .with_provider_state(state)
13342 .unwrap()
13343 .with_tool_calls(vec![call])
13344 .unwrap()
13345 }
13346
13347 struct TerminalHistoryProvider {
13348 calls: Arc<std::sync::atomic::AtomicU32>,
13349 }
13350
13351 struct DroppingSignedAssistantMemory {
13352 messages: RwLock<Vec<ChatMessage>>,
13353 }
13354
13355 struct DroppingEarlierSequentialMemory {
13356 messages: RwLock<Vec<ChatMessage>>,
13357 signed_seen: std::sync::atomic::AtomicUsize,
13358 }
13359
13360 #[async_trait]
13361 impl ai_agents_core::Memory for DroppingSignedAssistantMemory {
13362 async fn add_message(&self, message: ChatMessage) -> Result<()> {
13363 let signed = message.role == ai_agents_core::Role::Assistant
13364 && ai_agents_core::decode_native_tool_call_markers(&message.content)
13365 .map_err(|error| AgentError::LLM(error.to_string()))?
13366 .is_some_and(|batch| batch.provider_state().is_some());
13367 if !signed {
13368 self.messages.write().push(message);
13369 }
13370 Ok(())
13371 }
13372
13373 async fn get_messages(&self, limit: Option<usize>) -> Result<Vec<ChatMessage>> {
13374 let messages = self.messages.read();
13375 let start = limit
13376 .map(|limit| messages.len().saturating_sub(limit))
13377 .unwrap_or(0);
13378 Ok(messages[start..].to_vec())
13379 }
13380
13381 async fn clear(&self) -> Result<()> {
13382 self.messages.write().clear();
13383 Ok(())
13384 }
13385
13386 fn len(&self) -> usize {
13387 self.messages.read().len()
13388 }
13389
13390 async fn restore(&self, snapshot: ai_agents_core::MemorySnapshot) -> Result<()> {
13391 *self.messages.write() = snapshot.messages;
13392 Ok(())
13393 }
13394 }
13395
13396 #[async_trait]
13397 impl ai_agents_memory::Memory for DroppingSignedAssistantMemory {}
13398
13399 #[async_trait]
13400 impl ai_agents_core::Memory for DroppingEarlierSequentialMemory {
13401 async fn add_message(&self, message: ChatMessage) -> Result<()> {
13402 let signed = message.role == ai_agents_core::Role::Assistant
13403 && ai_agents_core::decode_native_tool_call_markers(&message.content)
13404 .map_err(|error| AgentError::LLM(error.to_string()))?
13405 .is_some_and(|batch| batch.provider_state().is_some());
13406 let mut messages = self.messages.write();
13407 if signed && self.signed_seen.fetch_add(1, Ordering::SeqCst) == 1 {
13408 messages.retain(|stored| !stored.content.contains("seq-call-1"));
13409 }
13410 messages.push(message);
13411 Ok(())
13412 }
13413
13414 async fn get_messages(&self, limit: Option<usize>) -> Result<Vec<ChatMessage>> {
13415 let messages = self.messages.read();
13416 let start = limit
13417 .map(|limit| messages.len().saturating_sub(limit))
13418 .unwrap_or(0);
13419 Ok(messages[start..].to_vec())
13420 }
13421
13422 async fn clear(&self) -> Result<()> {
13423 self.messages.write().clear();
13424 self.signed_seen.store(0, Ordering::SeqCst);
13425 Ok(())
13426 }
13427
13428 fn len(&self) -> usize {
13429 self.messages.read().len()
13430 }
13431
13432 async fn restore(&self, snapshot: ai_agents_core::MemorySnapshot) -> Result<()> {
13433 *self.messages.write() = snapshot.messages;
13434 self.signed_seen.store(0, Ordering::SeqCst);
13435 Ok(())
13436 }
13437 }
13438
13439 #[async_trait]
13440 impl ai_agents_memory::Memory for DroppingEarlierSequentialMemory {}
13441
13442 #[async_trait]
13443 impl LLMProvider for TerminalHistoryProvider {
13444 async fn complete(
13445 &self,
13446 _messages: &[ChatMessage],
13447 _config: Option<&LLMConfig>,
13448 ) -> std::result::Result<LLMResponse, LLMError> {
13449 self.calls.fetch_add(1, Ordering::SeqCst);
13450 Err(LLMError::Serialization(
13451 "native history integrity failure".to_string(),
13452 ))
13453 }
13454
13455 async fn complete_stream(
13456 &self,
13457 _messages: &[ChatMessage],
13458 _config: Option<&LLMConfig>,
13459 ) -> std::result::Result<
13460 Box<dyn Stream<Item = std::result::Result<LLMChunk, LLMError>> + Unpin + Send>,
13461 LLMError,
13462 > {
13463 Err(LLMError::Serialization(
13464 "native history integrity failure".to_string(),
13465 ))
13466 }
13467
13468 fn provider_name(&self) -> &str {
13469 "terminal-history"
13470 }
13471
13472 fn supports(&self, _feature: LLMFeature) -> bool {
13473 false
13474 }
13475
13476 fn is_terminal_error(&self, error: &LLMError) -> bool {
13477 matches!(error, LLMError::Serialization(_))
13478 }
13479 }
13480
13481 fn disambiguation_state_machine(
13483 state_enabled: Option<bool>,
13484 require_confirmation: bool,
13485 ) -> Arc<StateMachine> {
13486 let definition = ai_agents_state::StateDefinition {
13487 prompt: Some("Handle the resolved request.".to_string()),
13488 disambiguation: Some(ai_agents_disambiguation::StateDisambiguationOverride {
13489 enabled: state_enabled,
13490 require_confirmation,
13491 ..Default::default()
13492 }),
13493 ..Default::default()
13494 };
13495 let review = ai_agents_state::StateDefinition {
13496 prompt: Some("Review a fresh request.".to_string()),
13497 ..Default::default()
13498 };
13499 Arc::new(
13500 StateMachine::new(ai_agents_state::StateConfig {
13501 initial: "active".to_string(),
13502 states: std::collections::HashMap::from([
13503 ("active".to_string(), definition),
13504 ("review".to_string(), review),
13505 ]),
13506 global_transitions: Vec::new(),
13507 fallback: None,
13508 max_no_transition: None,
13509 regenerate_on_transition: true,
13510 })
13511 .unwrap(),
13512 )
13513 }
13514
13515 fn state_disambiguation_agent(
13517 responses: Vec<&str>,
13518 manager_enabled: bool,
13519 state_enabled: Option<bool>,
13520 require_confirmation: bool,
13521 ) -> (RuntimeAgent, MockLLMProvider) {
13522 state_disambiguation_agent_with_skills(
13523 responses,
13524 manager_enabled,
13525 state_enabled,
13526 require_confirmation,
13527 Vec::new(),
13528 )
13529 }
13530
13531 fn state_disambiguation_agent_with_skills(
13533 responses: Vec<&str>,
13534 manager_enabled: bool,
13535 state_enabled: Option<bool>,
13536 require_confirmation: bool,
13537 skills: Vec<SkillDefinition>,
13538 ) -> (RuntimeAgent, MockLLMProvider) {
13539 let mut mock = MockLLMProvider::new("state-confirmation");
13540 mock.set_responses(responses.into_iter().map(String::from).collect(), false);
13541 let observed = mock.clone();
13542 let agent = AgentBuilder::new()
13543 .system_prompt("Handle requests.")
13544 .llm(Arc::new(mock.clone()))
13545 .llm_alias("router", Arc::new(mock))
13546 .state_machine(disambiguation_state_machine(
13547 state_enabled,
13548 require_confirmation,
13549 ))
13550 .skills(skills)
13551 .build()
13552 .unwrap()
13553 .with_disambiguation(DisambiguationConfig {
13554 enabled: manager_enabled,
13555 ..Default::default()
13556 });
13557 (agent, observed)
13558 }
13559
13560 fn confirmation_skill() -> SkillDefinition {
13562 SkillDefinition {
13563 id: "send_report".to_string(),
13564 description: "Send a report after clarification".to_string(),
13565 trigger: "When the user asks to send a report".to_string(),
13566 steps: vec![SkillStep::Prompt {
13567 prompt: "Execute confirmed report skill for: {{ input }}".to_string(),
13568 llm: None,
13569 }],
13570 reasoning: None,
13571 reflection: None,
13572 disambiguation: Some(ai_agents_disambiguation::SkillDisambiguationOverride {
13573 enabled: Some(true),
13574 ..Default::default()
13575 }),
13576 }
13577 }
13578
13579 fn confirmation_skill_call_count(observed: &MockLLMProvider) -> usize {
13581 observed
13582 .call_history()
13583 .iter()
13584 .filter(|call| {
13585 call.messages
13586 .iter()
13587 .any(|message| message.content.contains("Execute confirmed report skill"))
13588 })
13589 .count()
13590 }
13591
13592 struct BlockingRuntimeConfirmationObserver {
13593 entered: tokio::sync::Barrier,
13594 release: tokio::sync::Notify,
13595 }
13596
13597 impl BlockingRuntimeConfirmationObserver {
13598 fn new() -> Self {
13599 Self {
13600 entered: tokio::sync::Barrier::new(2),
13601 release: tokio::sync::Notify::new(),
13602 }
13603 }
13604 }
13605
13606 struct ResetOnTransitionHooks {
13607 agent: parking_lot::Mutex<Option<Weak<RuntimeAgent>>>,
13608 invoked: AtomicBool,
13609 }
13610
13611 #[async_trait]
13612 impl AgentHooks for ResetOnTransitionHooks {
13613 async fn on_state_transition(&self, _from: Option<&str>, _to: &str, _reason: &str) {
13614 if self.invoked.swap(true, Ordering::SeqCst) {
13615 return;
13616 }
13617 let agent = self.agent.lock().as_ref().and_then(Weak::upgrade);
13618 if let Some(agent) = agent {
13619 agent.reset().await.unwrap();
13620 }
13621 }
13622 }
13623
13624 impl ClarificationObserver for BlockingRuntimeConfirmationObserver {
13625 fn observe_question<'a>(
13626 &'a self,
13627 future: ClarificationQuestionFuture<'a>,
13628 ) -> ClarificationQuestionFuture<'a> {
13629 future
13630 }
13631
13632 fn observe_parse<'a>(
13633 &'a self,
13634 future: ClarificationParseFuture<'a>,
13635 ) -> ClarificationParseFuture<'a> {
13636 future
13637 }
13638
13639 fn observe_confirmation_parse<'a>(
13640 &'a self,
13641 future: ConfirmationParseFuture<'a>,
13642 ) -> ConfirmationParseFuture<'a> {
13643 Box::pin(async move {
13644 self.entered.wait().await;
13645 self.release.notified().await;
13646 future.await
13647 })
13648 }
13649 }
13650
13651 #[tokio::test]
13652 async fn state_confirmation_blocks_redispatch_until_explicit_agreement() {
13653 let (agent, observed) = state_disambiguation_agent(
13654 vec![
13655 r#"{"is_ambiguous":true,"confidence":0.2,"ambiguity_type":"missing_target","reasoning":"target missing","what_is_unclear":["target"],"detected_language":"en"}"#,
13656 r#"{"question":"What should I send?","options":null}"#,
13657 r#"{"status":"answered","selected_option":null,"enriched_input":"Send the report to Ada","resolved":{"intent":"send_report"}}"#,
13658 r#"{"question":"Should I send the report to Ada?"}"#,
13659 r#"{"status":"confirmed"}"#,
13660 "Request executed.",
13661 ],
13662 true,
13663 None,
13664 true,
13665 );
13666
13667 let clarification = agent.chat("Send it").await.unwrap();
13668 assert_eq!(clarification.content, "What should I send?");
13669 assert_eq!(observed.call_count(), 2);
13670
13671 let confirmation = agent.chat("The report to Ada").await.unwrap();
13672 assert_eq!(confirmation.content, "Should I send the report to Ada?");
13673 assert_eq!(
13674 confirmation
13675 .metadata
13676 .as_ref()
13677 .and_then(|metadata| metadata.get("disambiguation"))
13678 .and_then(|metadata| metadata.get("status"))
13679 .and_then(Value::as_str),
13680 Some("awaiting_confirmation")
13681 );
13682 assert_eq!(observed.call_count(), 4);
13683
13684 let completed = agent.chat("Yes").await.unwrap();
13685 assert_eq!(completed.content, "Request executed.");
13686 assert_eq!(observed.call_count(), 6);
13687 }
13688
13689 #[tokio::test]
13690 async fn streaming_state_confirmation_ends_the_turn_before_redispatch() {
13691 let (agent, observed) = state_disambiguation_agent(
13692 vec![
13693 r#"{"is_ambiguous":true,"confidence":0.2,"ambiguity_type":"missing_target","reasoning":"target missing","what_is_unclear":["target"],"detected_language":"en"}"#,
13694 r#"{"question":"What should I send?","options":null}"#,
13695 r#"{"status":"answered","selected_option":null,"enriched_input":"Send the report to Ada","resolved":{"intent":"send_report"}}"#,
13696 r#"{"question":"Should I send the report to Ada?"}"#,
13697 r#"{"status":"confirmed"}"#,
13698 "Request executed.",
13699 ],
13700 true,
13701 None,
13702 true,
13703 );
13704
13705 let mut clarification_stream = agent.chat_stream("Send it").await.unwrap();
13706 let mut clarification = String::new();
13707 while let Some(chunk) = clarification_stream.next().await {
13708 match chunk {
13709 StreamChunk::Content { text } => clarification.push_str(&text),
13710 StreamChunk::Done {} => break,
13711 StreamChunk::Error { message } => panic!("unexpected stream error: {message}"),
13712 _ => {}
13713 }
13714 }
13715 assert_eq!(clarification, "What should I send?");
13716 assert_eq!(observed.call_count(), 2);
13717
13718 let mut confirmation_stream = agent.chat_stream_events("The report to Ada").await.unwrap();
13719 let mut confirmation = None;
13720 while let Some(event) = confirmation_stream.next().await {
13721 match event {
13722 AgentStreamEvent::Final(response) => confirmation = Some(response),
13723 AgentStreamEvent::Chunk(StreamChunk::Error { message }) => {
13724 panic!("unexpected stream error: {message}")
13725 }
13726 AgentStreamEvent::Chunk(_) => {}
13727 }
13728 }
13729 let confirmation = confirmation.expect("confirmation must finalize");
13730 assert_eq!(confirmation.content, "Should I send the report to Ada?");
13731 assert_eq!(
13732 confirmation
13733 .metadata
13734 .as_ref()
13735 .and_then(|metadata| metadata.get("disambiguation"))
13736 .and_then(|metadata| metadata.get("status"))
13737 .and_then(Value::as_str),
13738 Some("awaiting_confirmation")
13739 );
13740 assert_eq!(observed.call_count(), 4);
13741
13742 let mut completed_stream = agent.chat_stream("Yes").await.unwrap();
13743 let mut completed = String::new();
13744 while let Some(chunk) = completed_stream.next().await {
13745 match chunk {
13746 StreamChunk::Content { text } => completed.push_str(&text),
13747 StreamChunk::Done {} => break,
13748 StreamChunk::Error { message } => panic!("unexpected stream error: {message}"),
13749 _ => {}
13750 }
13751 }
13752 assert_eq!(completed, "Request executed.");
13753 assert_eq!(observed.call_count(), 6);
13754 }
13755
13756 #[tokio::test]
13758 async fn root_turn_gate_serializes_blocking_and_streaming_entry_points() {
13759 let (complete_entered, mut complete_events) = tokio::sync::mpsc::unbounded_channel();
13760 let agent = Arc::new(
13761 AgentBuilder::new()
13762 .system_prompt("Serialize root turns.")
13763 .llm(Arc::new(RootTurnProbeProvider { complete_entered }))
13764 .build()
13765 .unwrap(),
13766 );
13767 let blocking_agent = Arc::clone(&agent);
13768
13769 let legacy_stream = agent.chat_stream("stream owner").await.unwrap();
13770 assert!(agent.root_turn_gate.try_lock().is_err());
13771 let blocking = tokio::spawn(async move { blocking_agent.chat("blocked").await.unwrap() });
13772 assert!(
13773 tokio::time::timeout(std::time::Duration::from_millis(50), complete_events.recv())
13774 .await
13775 .is_err(),
13776 "blocking turn reached the provider while the legacy stream owned the root gate"
13777 );
13778
13779 drop(legacy_stream);
13780 assert_eq!(
13781 tokio::time::timeout(std::time::Duration::from_secs(2), complete_events.recv())
13782 .await
13783 .expect("blocking turn did not enter after stream drop"),
13784 Some(())
13785 );
13786 let response = tokio::time::timeout(std::time::Duration::from_secs(2), blocking)
13787 .await
13788 .expect("blocking turn did not finish after stream drop")
13789 .unwrap();
13790 assert_eq!(response.content, "blocking complete");
13791
13792 let mut event_stream = agent.chat_stream_events("event terminal").await.unwrap();
13793 assert!(agent.root_turn_gate.try_lock().is_err());
13794 let mut saw_final = false;
13795 while let Some(event) = event_stream.next().await {
13796 if matches!(event, AgentStreamEvent::Final(_)) {
13797 saw_final = true;
13798 break;
13799 }
13800 }
13801 assert!(saw_final);
13802 assert!(
13803 agent.root_turn_gate.try_lock().is_ok(),
13804 "authoritative terminal event retained the root gate"
13805 );
13806 }
13807
13808 #[tokio::test]
13810 async fn response_hook_rejects_same_runtime_chat_reentry() {
13811 let hooks = Arc::new(ResponseChatHooks {
13812 target: parking_lot::Mutex::new(None),
13813 invoked: AtomicBool::new(false),
13814 nested_result: parking_lot::Mutex::new(None),
13815 });
13816 let agent = Arc::new(
13817 AgentBuilder::new()
13818 .system_prompt("Reject response hook reentry.")
13819 .llm(Arc::new(mock_with_response("outer response")))
13820 .hooks(hooks.clone())
13821 .build()
13822 .unwrap(),
13823 );
13824 *hooks.target.lock() = Some(Arc::downgrade(&agent));
13825
13826 let response = tokio::time::timeout(
13827 std::time::Duration::from_secs(2),
13828 agent.chat("outer request"),
13829 )
13830 .await
13831 .expect("same-runtime response hook reentry must fail without deadlocking")
13832 .unwrap();
13833
13834 assert_eq!(response.content, "outer response");
13835 let nested_result = hooks
13836 .nested_result
13837 .lock()
13838 .clone()
13839 .expect("response hook must record its nested call");
13840 let error = nested_result.expect_err("same-runtime nested chat must be rejected");
13841 assert!(error.contains("reentrant root turn ownership"));
13842 }
13843
13844 #[tokio::test]
13846 async fn root_turn_gate_allows_nested_runtime_and_rejects_cycles() {
13847 let agent_a = AgentBuilder::new()
13848 .system_prompt("Runtime A.")
13849 .llm(Arc::new(mock_with_response("response A")))
13850 .build()
13851 .unwrap();
13852 let agent_b = AgentBuilder::new()
13853 .system_prompt("Runtime B.")
13854 .llm(Arc::new(mock_with_response("response B")))
13855 .build()
13856 .unwrap();
13857 let RootTurnAdmission {
13858 guard: guard_a,
13859 identity_stack: stack_a,
13860 } = agent_a.acquire_root_turn().await.unwrap();
13861
13862 let cycle_error = scope_runtime_gate_identity_stack(&stack_a, async {
13863 let RootTurnAdmission {
13864 guard: guard_b,
13865 identity_stack: stack_b,
13866 } = agent_b
13867 .acquire_root_turn()
13868 .await
13869 .expect("runtime B must acquire a different gate");
13870 let result =
13871 scope_runtime_gate_identity_stack(&stack_b, agent_a.acquire_root_turn()).await;
13872 drop(guard_b);
13873 match result {
13874 Err(error) => error,
13875 Ok(_) => panic!("runtime A accepted a repeated gate identity"),
13876 }
13877 })
13878 .await;
13879 drop(guard_a);
13880
13881 assert!(
13882 cycle_error
13883 .to_string()
13884 .contains("reentrant root turn ownership")
13885 );
13886 }
13887
13888 #[tokio::test]
13890 async fn concurrent_orchestration_propagates_root_gate_ancestry() {
13891 let registry = Arc::new(crate::spawner::AgentRegistry::new());
13892 let hooks_a = Arc::new(ConcurrentResponseHooks {
13893 registry: Arc::downgrade(®istry),
13894 child_id: "runtime-b".to_string(),
13895 invoked: AtomicBool::new(false),
13896 nested_result: parking_lot::Mutex::new(None),
13897 });
13898 let hooks_b = Arc::new(ResponseChatHooks {
13899 target: parking_lot::Mutex::new(None),
13900 invoked: AtomicBool::new(false),
13901 nested_result: parking_lot::Mutex::new(None),
13902 });
13903 let agent_a = AgentBuilder::new()
13904 .system_prompt("Runtime A dispatches runtime B concurrently.")
13905 .llm(Arc::new(mock_with_response("response A")))
13906 .hooks(hooks_a.clone())
13907 .build()
13908 .unwrap();
13909 let agent_b = AgentBuilder::new()
13910 .system_prompt("Runtime B attempts to re-enter runtime A.")
13911 .llm(Arc::new(mock_with_response("response B")))
13912 .hooks(hooks_b.clone())
13913 .build()
13914 .unwrap();
13915 let spec_a = crate::spec::AgentSpec {
13916 name: "runtime-a".to_string(),
13917 system_prompt: "Runtime A dispatches runtime B concurrently.".to_string(),
13918 ..crate::spec::AgentSpec::default()
13919 };
13920 let spec_b = crate::spec::AgentSpec {
13921 name: "runtime-b".to_string(),
13922 system_prompt: "Runtime B attempts to re-enter runtime A.".to_string(),
13923 ..crate::spec::AgentSpec::default()
13924 };
13925 registry
13926 .register(crate::spawner::SpawnedAgent::from_runtime(
13927 "runtime-a".to_string(),
13928 agent_a,
13929 spec_a,
13930 ))
13931 .await
13932 .unwrap();
13933 registry
13934 .register(crate::spawner::SpawnedAgent::from_runtime(
13935 "runtime-b".to_string(),
13936 agent_b,
13937 spec_b,
13938 ))
13939 .await
13940 .unwrap();
13941 let runtime_a = registry.get("runtime-a").unwrap();
13942 *hooks_b.target.lock() = Some(Arc::downgrade(&runtime_a));
13943
13944 let response = tokio::time::timeout(
13945 std::time::Duration::from_secs(2),
13946 runtime_a.chat("outer concurrent request"),
13947 )
13948 .await
13949 .expect("concurrent orchestration cycle must fail without deadlocking")
13950 .unwrap();
13951
13952 assert_eq!(response.content, "response A");
13953 let child_result = hooks_a
13954 .nested_result
13955 .lock()
13956 .clone()
13957 .expect("runtime A hook must record runtime B completion");
13958 assert_eq!(child_result.unwrap(), "response B");
13959 let cycle_result = hooks_b
13960 .nested_result
13961 .lock()
13962 .clone()
13963 .expect("runtime B hook must record runtime A reentry");
13964 assert!(
13965 cycle_result
13966 .expect_err("runtime A accepted a repeated gate identity")
13967 .contains("reentrant root turn ownership")
13968 );
13969 }
13970
13971 #[tokio::test]
13973 async fn confirmed_skill_route_executes_exactly_once() {
13974 let (agent, observed) = state_disambiguation_agent_with_skills(
13975 vec![
13976 r#"{"is_ambiguous":false,"confidence":0.99,"ambiguity_type":null,"reasoning":"top-level clear","what_is_unclear":[],"detected_language":"en"}"#,
13977 "send_report",
13978 r#"{"is_ambiguous":true,"confidence":0.2,"ambiguity_type":"missing_target","reasoning":"target missing","what_is_unclear":["target"],"detected_language":"en"}"#,
13979 r#"{"question":"What should I send?","options":null}"#,
13980 r#"{"status":"answered","selected_option":null,"enriched_input":"Send the report to Ada","resolved":{"intent":"send_report"}}"#,
13981 r#"{"question":"Should I send the report to Ada?"}"#,
13982 r#"{"status":"confirmed"}"#,
13983 r#"{"is_ambiguous":false,"confidence":0.99,"ambiguity_type":null,"reasoning":"resolved","what_is_unclear":[],"detected_language":"en"}"#,
13984 "Report skill executed.",
13985 ],
13986 true,
13987 None,
13988 true,
13989 vec![confirmation_skill()],
13990 );
13991
13992 let clarification = agent.chat("Send it").await.unwrap();
13993 assert_eq!(clarification.content, "What should I send?");
13994 assert_eq!(confirmation_skill_call_count(&observed), 0);
13995
13996 let confirmation = agent.chat("The report to Ada").await.unwrap();
13997 assert_eq!(confirmation.content, "Should I send the report to Ada?");
13998 assert_eq!(
13999 confirmation
14000 .metadata
14001 .as_ref()
14002 .and_then(|metadata| metadata.get("disambiguation"))
14003 .and_then(|metadata| metadata.get("status"))
14004 .and_then(Value::as_str),
14005 Some("awaiting_confirmation")
14006 );
14007 assert_eq!(confirmation_skill_call_count(&observed), 0);
14008
14009 let completed = agent.chat("Yes").await.unwrap();
14010 assert_eq!(completed.content, "Report skill executed.");
14011 assert_eq!(confirmation_skill_call_count(&observed), 1);
14012 assert!(agent.pending_skill_id.read().is_none());
14013 let messages = agent.memory.get_messages(None).await.unwrap();
14014 assert!(!messages.iter().any(|message| message.content == "Yes"));
14015 }
14016
14017 #[tokio::test]
14019 async fn confirmed_skill_recheck_preserves_new_clarification_metadata() {
14020 let (agent, observed) = state_disambiguation_agent_with_skills(
14021 vec![
14022 r#"{"is_ambiguous":false,"confidence":0.99,"ambiguity_type":null,"reasoning":"top-level clear","what_is_unclear":[],"detected_language":"en"}"#,
14023 "send_report",
14024 r#"{"is_ambiguous":true,"confidence":0.2,"ambiguity_type":"missing_target","reasoning":"target missing","what_is_unclear":["target"],"detected_language":"en"}"#,
14025 r#"{"question":"What should I send?","options":null}"#,
14026 r#"{"status":"answered","selected_option":null,"enriched_input":"Send the report to Ada","resolved":{"intent":"send_report"}}"#,
14027 r#"{"question":"Should I send the report to Ada?"}"#,
14028 r#"{"status":"confirmed"}"#,
14029 r#"{"is_ambiguous":true,"confidence":0.3,"ambiguity_type":"missing_parameters","reasoning":"timing missing","what_is_unclear":["timing"],"detected_language":"en"}"#,
14030 r#"{"question":"When should I send it?","options":null}"#,
14031 ],
14032 true,
14033 None,
14034 true,
14035 vec![confirmation_skill()],
14036 );
14037
14038 agent.chat("Send it").await.unwrap();
14039 agent.chat("The report to Ada").await.unwrap();
14040 let follow_up = agent.chat("Yes").await.unwrap();
14041
14042 assert_eq!(follow_up.content, "When should I send it?");
14043 let metadata = follow_up
14044 .metadata
14045 .as_ref()
14046 .and_then(|metadata| metadata.get("disambiguation"))
14047 .unwrap();
14048 assert_eq!(
14049 metadata.get("status").and_then(Value::as_str),
14050 Some("awaiting_clarification")
14051 );
14052 assert_eq!(
14053 metadata.get("skill_id").and_then(Value::as_str),
14054 Some("send_report")
14055 );
14056 assert!(metadata.get("detection").is_some());
14057 assert_eq!(confirmation_skill_call_count(&observed), 0);
14058 }
14059
14060 #[tokio::test]
14062 async fn rejected_skill_confirmation_never_executes() {
14063 let (agent, observed) = state_disambiguation_agent_with_skills(
14064 vec![
14065 r#"{"is_ambiguous":false,"confidence":0.99,"ambiguity_type":null,"reasoning":"top-level clear","what_is_unclear":[],"detected_language":"en"}"#,
14066 "send_report",
14067 r#"{"is_ambiguous":true,"confidence":0.2,"ambiguity_type":"missing_target","reasoning":"target missing","what_is_unclear":["target"],"detected_language":"en"}"#,
14068 r#"{"question":"What should I send?","options":null}"#,
14069 r#"{"status":"answered","selected_option":null,"enriched_input":"Send the report to Ada","resolved":{"intent":"send_report"}}"#,
14070 r#"{"question":"Should I send the report to Ada?"}"#,
14071 r#"{"status":"rejected"}"#,
14072 "Confirmation rejected.",
14073 ],
14074 true,
14075 None,
14076 true,
14077 vec![confirmation_skill()],
14078 );
14079
14080 agent.chat("Send it").await.unwrap();
14081 agent.chat("The report to Ada").await.unwrap();
14082 let rejected = agent.chat("No").await.unwrap();
14083
14084 assert_eq!(rejected.content, "Confirmation rejected.");
14085 assert_eq!(confirmation_skill_call_count(&observed), 0);
14086 assert!(agent.pending_skill_id.read().is_none());
14087 }
14088
14089 #[tokio::test]
14091 async fn reset_invalidates_pending_skill_confirmation_before_streaming_input() {
14092 let (agent, observed) = state_disambiguation_agent_with_skills(
14093 vec![
14094 r#"{"is_ambiguous":false,"confidence":0.99,"ambiguity_type":null,"reasoning":"top-level clear","what_is_unclear":[],"detected_language":"en"}"#,
14095 "send_report",
14096 r#"{"is_ambiguous":true,"confidence":0.2,"ambiguity_type":"missing_target","reasoning":"target missing","what_is_unclear":["target"],"detected_language":"en"}"#,
14097 r#"{"question":"What should I send?","options":null}"#,
14098 r#"{"status":"answered","selected_option":null,"enriched_input":"Send the report to Ada","resolved":{"intent":"send_report"}}"#,
14099 r#"{"question":"Should I send the report to Ada?"}"#,
14100 r#"{"is_ambiguous":false,"confidence":0.99,"ambiguity_type":null,"reasoning":"fresh input","what_is_unclear":[],"detected_language":"en"}"#,
14101 "none",
14102 "Fresh response.",
14103 ],
14104 true,
14105 None,
14106 true,
14107 vec![confirmation_skill()],
14108 );
14109
14110 agent.chat("Send it").await.unwrap();
14111 agent.chat("The report to Ada").await.unwrap();
14112 agent.reset().await.unwrap();
14113 assert!(agent.pending_skill_id.read().is_none());
14114 assert!(
14115 !agent
14116 .disambiguation_manager()
14117 .unwrap()
14118 .has_pending_clarification()
14119 .await
14120 );
14121
14122 let mut stream = agent.chat_stream("Yes").await.unwrap();
14123 let mut content = String::new();
14124 while let Some(chunk) = stream.next().await {
14125 match chunk {
14126 StreamChunk::Content { text } => content.push_str(&text),
14127 StreamChunk::Done {} => break,
14128 StreamChunk::Error { message } => panic!("unexpected stream error: {message}"),
14129 _ => {}
14130 }
14131 }
14132
14133 assert_eq!(content, "Fresh response.");
14134 assert_eq!(confirmation_skill_call_count(&observed), 0);
14135 }
14136
14137 #[tokio::test]
14139 async fn trait_reset_clears_pending_skill_confirmation() {
14140 let (agent, _) = state_disambiguation_agent_with_skills(
14141 vec![
14142 r#"{"is_ambiguous":false,"confidence":0.99,"ambiguity_type":null,"reasoning":"top-level clear","what_is_unclear":[],"detected_language":"en"}"#,
14143 "send_report",
14144 r#"{"is_ambiguous":true,"confidence":0.2,"ambiguity_type":"missing_target","reasoning":"target missing","what_is_unclear":["target"],"detected_language":"en"}"#,
14145 r#"{"question":"What should I send?","options":null}"#,
14146 r#"{"status":"answered","selected_option":null,"enriched_input":"Send the report to Ada","resolved":{"intent":"send_report"}}"#,
14147 r#"{"question":"Should I send the report to Ada?"}"#,
14148 ],
14149 true,
14150 None,
14151 true,
14152 vec![confirmation_skill()],
14153 );
14154
14155 agent.chat("Send it").await.unwrap();
14156 agent.chat("The report to Ada").await.unwrap();
14157 <RuntimeAgent as Agent>::reset(&agent).await.unwrap();
14158
14159 assert!(agent.pending_skill_id.read().is_none());
14160 assert!(
14161 !agent
14162 .disambiguation_manager()
14163 .unwrap()
14164 .has_pending_clarification()
14165 .await
14166 );
14167 }
14168
14169 #[tokio::test]
14171 async fn state_change_invalidates_pending_skill_confirmation() {
14172 let (agent, observed) = state_disambiguation_agent_with_skills(
14173 vec![
14174 r#"{"is_ambiguous":false,"confidence":0.99,"ambiguity_type":null,"reasoning":"top-level clear","what_is_unclear":[],"detected_language":"en"}"#,
14175 "send_report",
14176 r#"{"is_ambiguous":true,"confidence":0.2,"ambiguity_type":"missing_target","reasoning":"target missing","what_is_unclear":["target"],"detected_language":"en"}"#,
14177 r#"{"question":"What should I send?","options":null}"#,
14178 r#"{"status":"answered","selected_option":null,"enriched_input":"Send the report to Ada","resolved":{"intent":"send_report"}}"#,
14179 r#"{"question":"Should I send the report to Ada?"}"#,
14180 r#"{"is_ambiguous":false,"confidence":0.99,"ambiguity_type":null,"reasoning":"fresh input","what_is_unclear":[],"detected_language":"en"}"#,
14181 "none",
14182 "Fresh response.",
14183 ],
14184 true,
14185 None,
14186 true,
14187 vec![confirmation_skill()],
14188 );
14189
14190 agent.chat("Send it").await.unwrap();
14191 agent.chat("The report to Ada").await.unwrap();
14192 agent.transition_to("review").await.unwrap();
14193 let cancelled = agent.chat("Yes").await.unwrap();
14194
14195 assert_eq!(cancelled.content, "Fresh response.");
14196 assert_eq!(confirmation_skill_call_count(&observed), 0);
14197 assert!(agent.pending_skill_id.read().is_none());
14198 }
14199
14200 #[tokio::test]
14202 async fn in_flight_confirmation_cannot_redispatch_after_reset() {
14203 let (mut agent, observed) = state_disambiguation_agent_with_skills(
14204 vec![
14205 r#"{"is_ambiguous":false,"confidence":0.99,"ambiguity_type":null,"reasoning":"top-level clear","what_is_unclear":[],"detected_language":"en"}"#,
14206 "send_report",
14207 r#"{"is_ambiguous":true,"confidence":0.2,"ambiguity_type":"missing_target","reasoning":"target missing","what_is_unclear":["target"],"detected_language":"en"}"#,
14208 r#"{"question":"What should I send?","options":null}"#,
14209 r#"{"status":"answered","selected_option":null,"enriched_input":"Send the report to Ada","resolved":{"intent":"send_report"}}"#,
14210 r#"{"question":"Should I send the report to Ada?"}"#,
14211 r#"{"status":"confirmed"}"#,
14212 "Confirmation cancelled.",
14213 ],
14214 true,
14215 None,
14216 true,
14217 vec![confirmation_skill()],
14218 );
14219 let observer = Arc::new(BlockingRuntimeConfirmationObserver::new());
14220 let manager = agent
14221 .disambiguation_manager
14222 .take()
14223 .unwrap()
14224 .with_clarification_observer(observer.clone());
14225 agent.disambiguation_manager = Some(manager);
14226 let agent = Arc::new(agent);
14227
14228 agent.chat("Send it").await.unwrap();
14229 agent.chat("The report to Ada").await.unwrap();
14230
14231 let confirming_agent = Arc::clone(&agent);
14232 let confirmation = tokio::spawn(async move { confirming_agent.chat("Yes").await });
14233 observer.entered.wait().await;
14234 agent.reset().await.unwrap();
14235 observer.release.notify_one();
14236
14237 let response = confirmation.await.unwrap().unwrap();
14238 assert_eq!(response.content, "Confirmation cancelled.");
14239 assert_eq!(confirmation_skill_call_count(&observed), 0);
14240 assert!(agent.pending_skill_id.read().is_none());
14241 }
14242
14243 #[tokio::test]
14245 async fn queued_reset_prevents_stale_confirmation_question_publication() {
14246 let (agent, observed) = state_disambiguation_agent(
14247 vec![
14248 r#"{"is_ambiguous":true,"confidence":0.2,"ambiguity_type":"missing_target","reasoning":"target missing","what_is_unclear":["target"],"detected_language":"en"}"#,
14249 r#"{"question":"What should I send?","options":null}"#,
14250 r#"{"status":"answered","selected_option":null,"enriched_input":"Send the report to Ada","resolved":{"intent":"send_report"}}"#,
14251 r#"{"question":"Should I send the report to Ada?"}"#,
14252 ],
14253 true,
14254 None,
14255 true,
14256 );
14257 let agent = Arc::new(agent);
14258 agent.chat("Send it").await.unwrap();
14259
14260 let admission = agent.disambiguation_admission.write().await;
14261 let (started_tx, started_rx) = tokio::sync::oneshot::channel();
14262 let resetting_agent = Arc::clone(&agent);
14263 let reset = tokio::spawn(async move {
14264 let _ = started_tx.send(());
14265 resetting_agent.reset().await
14266 });
14267 started_rx.await.unwrap();
14268 tokio::task::yield_now().await;
14269
14270 let responding_agent = Arc::clone(&agent);
14271 let response =
14272 tokio::spawn(async move { responding_agent.chat("The report to Ada").await });
14273 tokio::time::timeout(std::time::Duration::from_secs(2), async {
14274 while observed.call_count() < 4 {
14275 tokio::task::yield_now().await;
14276 }
14277 })
14278 .await
14279 .expect("clarification processing must reach terminal publication");
14280 drop(admission);
14281
14282 reset.await.unwrap().unwrap();
14283 let error = response.await.unwrap().unwrap_err();
14284 assert!(error.to_string().contains("ownership changed"));
14285 assert!(
14286 !agent
14287 .disambiguation_manager()
14288 .unwrap()
14289 .has_pending_clarification()
14290 .await
14291 );
14292 assert!(agent.memory.get_messages(None).await.unwrap().is_empty());
14293 }
14294
14295 #[tokio::test]
14297 async fn queued_reset_prevents_stale_skill_clarification_publication() {
14298 let (agent, observed) = state_disambiguation_agent_with_skills(
14299 vec![
14300 r#"{"is_ambiguous":false,"confidence":0.99,"ambiguity_type":null,"reasoning":"top-level clear","what_is_unclear":[],"detected_language":"en"}"#,
14301 "send_report",
14302 r#"{"is_ambiguous":true,"confidence":0.2,"ambiguity_type":"missing_target","reasoning":"target missing","what_is_unclear":["target"],"detected_language":"en"}"#,
14303 r#"{"question":"What should I send?","options":null}"#,
14304 ],
14305 true,
14306 None,
14307 true,
14308 vec![confirmation_skill()],
14309 );
14310 let agent = Arc::new(agent);
14311 let admission = agent.disambiguation_admission.write().await;
14312 let (started_tx, started_rx) = tokio::sync::oneshot::channel();
14313 let resetting_agent = Arc::clone(&agent);
14314 let reset = tokio::spawn(async move {
14315 let _ = started_tx.send(());
14316 resetting_agent.reset().await
14317 });
14318 started_rx.await.unwrap();
14319 tokio::task::yield_now().await;
14320
14321 let responding_agent = Arc::clone(&agent);
14322 let response = tokio::spawn(async move { responding_agent.chat("Send it").await });
14323 tokio::time::timeout(std::time::Duration::from_secs(2), async {
14324 while observed.call_count() < 4 {
14325 tokio::task::yield_now().await;
14326 }
14327 })
14328 .await
14329 .expect("skill clarification must reach terminal publication");
14330 drop(admission);
14331
14332 reset.await.unwrap().unwrap();
14333 let error = response.await.unwrap().unwrap_err();
14334 assert!(error.to_string().contains("ownership changed"));
14335 assert_eq!(confirmation_skill_call_count(&observed), 0);
14336 assert!(agent.pending_skill_id.read().is_none());
14337 assert!(agent.memory.get_messages(None).await.unwrap().is_empty());
14338 }
14339
14340 #[tokio::test]
14342 async fn transition_hook_can_reset_without_admission_deadlock() {
14343 let hooks = Arc::new(ResetOnTransitionHooks {
14344 agent: parking_lot::Mutex::new(None),
14345 invoked: AtomicBool::new(false),
14346 });
14347 let agent = Arc::new(
14348 AgentBuilder::new()
14349 .system_prompt("Test transition hook reentrancy.")
14350 .llm(Arc::new(mock_with_response("done")))
14351 .state_machine(disambiguation_state_machine(None, false))
14352 .build()
14353 .unwrap()
14354 .with_hooks(hooks.clone()),
14355 );
14356 *hooks.agent.lock() = Some(Arc::downgrade(&agent));
14357
14358 let transitioned = tokio::time::timeout(
14359 std::time::Duration::from_secs(2),
14360 agent.apply_transition_target("active", "review", "test transition", None),
14361 )
14362 .await
14363 .expect("transition hook reset must not deadlock")
14364 .unwrap();
14365
14366 assert!(transitioned);
14367 assert!(hooks.invoked.load(Ordering::SeqCst));
14368 assert_eq!(agent.current_state().as_deref(), Some("active"));
14369 }
14370
14371 #[tokio::test]
14373 async fn concurrent_transition_cannot_duplicate_exit_actions() {
14374 let gate = PathMutationGate::new();
14375 let active = ai_agents_state::StateDefinition {
14376 on_exit: vec![StateAction::Tool {
14377 tool: "transition_exit".to_string(),
14378 args: Some(serde_json::json!({"path": "./transition-exit.txt"})),
14379 }],
14380 ..Default::default()
14381 };
14382 let state_machine = Arc::new(
14383 StateMachine::new(ai_agents_state::StateConfig {
14384 initial: "active".to_string(),
14385 states: HashMap::from([
14386 ("active".to_string(), active),
14387 (
14388 "review".to_string(),
14389 ai_agents_state::StateDefinition::default(),
14390 ),
14391 ]),
14392 global_transitions: Vec::new(),
14393 fallback: None,
14394 max_no_transition: None,
14395 regenerate_on_transition: true,
14396 })
14397 .unwrap(),
14398 );
14399 let agent = Arc::new(
14400 AgentBuilder::new()
14401 .system_prompt("Test transition reservation.")
14402 .llm(Arc::new(mock_with_response("done")))
14403 .tool(Arc::new(BlockingPathMutationTool {
14404 id: "transition_exit",
14405 path_fields: vec![ai_agents_core::PathPolicyBinding::write("path")],
14406 gate: gate.clone(),
14407 }))
14408 .state_machine(state_machine)
14409 .build()
14410 .unwrap(),
14411 );
14412
14413 let first_agent = Arc::clone(&agent);
14414 let first = tokio::spawn(async move { first_agent.transition_to("review").await });
14415 tokio::time::timeout(std::time::Duration::from_secs(2), gate.wait_until_entered())
14416 .await
14417 .expect("reserved transition must enter its exit action");
14418
14419 let second = tokio::time::timeout(
14420 std::time::Duration::from_secs(2),
14421 agent.transition_to("review"),
14422 )
14423 .await
14424 .expect("competing transition must fail without waiting for the exit action")
14425 .unwrap_err();
14426 assert!(second.to_string().contains("already in progress"));
14427
14428 gate.release();
14429 first.await.unwrap().unwrap();
14430 assert_eq!(agent.current_state().as_deref(), Some("review"));
14431 }
14432
14433 #[tokio::test]
14435 async fn concurrent_transition_cannot_overtake_enter_actions() {
14436 let gate = PathMutationGate::new();
14437 let review = ai_agents_state::StateDefinition {
14438 on_enter: vec![StateAction::Tool {
14439 tool: "transition_enter".to_string(),
14440 args: Some(serde_json::json!({"path": "./transition-enter.txt"})),
14441 }],
14442 ..Default::default()
14443 };
14444 let state_machine = Arc::new(
14445 StateMachine::new(ai_agents_state::StateConfig {
14446 initial: "active".to_string(),
14447 states: HashMap::from([
14448 (
14449 "active".to_string(),
14450 ai_agents_state::StateDefinition::default(),
14451 ),
14452 ("review".to_string(), review),
14453 ]),
14454 global_transitions: Vec::new(),
14455 fallback: None,
14456 max_no_transition: None,
14457 regenerate_on_transition: true,
14458 })
14459 .unwrap(),
14460 );
14461 let agent = Arc::new(
14462 AgentBuilder::new()
14463 .system_prompt("Test transition lifecycle reservation.")
14464 .llm(Arc::new(mock_with_response("done")))
14465 .tool(Arc::new(BlockingPathMutationTool {
14466 id: "transition_enter",
14467 path_fields: vec![ai_agents_core::PathPolicyBinding::write("path")],
14468 gate: gate.clone(),
14469 }))
14470 .state_machine(state_machine)
14471 .build()
14472 .unwrap(),
14473 );
14474
14475 let first_agent = Arc::clone(&agent);
14476 let first = tokio::spawn(async move { first_agent.transition_to("review").await });
14477 tokio::time::timeout(std::time::Duration::from_secs(2), gate.wait_until_entered())
14478 .await
14479 .expect("committed transition must enter its destination action");
14480
14481 let second = agent.transition_to("active").await.unwrap_err();
14482 assert!(second.to_string().contains("already in progress"));
14483 assert!(agent.reset().await.is_err());
14484
14485 gate.release();
14486 first.await.unwrap().unwrap();
14487 assert_eq!(agent.current_state().as_deref(), Some("review"));
14488 }
14489
14490 #[tokio::test]
14492 async fn same_state_restore_invalidates_pending_skill_confirmation() {
14493 let (agent, observed) = state_disambiguation_agent_with_skills(
14494 vec![
14495 r#"{"is_ambiguous":false,"confidence":0.99,"ambiguity_type":null,"reasoning":"top-level clear","what_is_unclear":[],"detected_language":"en"}"#,
14496 "send_report",
14497 r#"{"is_ambiguous":true,"confidence":0.2,"ambiguity_type":"missing_target","reasoning":"target missing","what_is_unclear":["target"],"detected_language":"en"}"#,
14498 r#"{"question":"What should I send?","options":null}"#,
14499 r#"{"status":"answered","selected_option":null,"enriched_input":"Send the report to Ada","resolved":{"intent":"send_report"}}"#,
14500 r#"{"question":"Should I send the report to Ada?"}"#,
14501 ],
14502 true,
14503 None,
14504 true,
14505 vec![confirmation_skill()],
14506 );
14507
14508 agent.chat("Send it").await.unwrap();
14509 agent.chat("The report to Ada").await.unwrap();
14510 let snapshot = agent.save_state().await.unwrap();
14511 assert_eq!(agent.current_state().as_deref(), Some("active"));
14512
14513 agent.restore_state(snapshot).await.unwrap();
14514
14515 assert_eq!(agent.current_state().as_deref(), Some("active"));
14516 assert!(agent.pending_skill_id.read().is_none());
14517 assert!(
14518 !agent
14519 .disambiguation_manager()
14520 .unwrap()
14521 .has_pending_clarification()
14522 .await
14523 );
14524 assert_eq!(confirmation_skill_call_count(&observed), 0);
14525 }
14526
14527 #[tokio::test]
14529 async fn direct_state_generation_change_invalidates_confirmation() {
14530 let (agent, observed) = state_disambiguation_agent_with_skills(
14531 vec![
14532 r#"{"is_ambiguous":false,"confidence":0.99,"ambiguity_type":null,"reasoning":"top-level clear","what_is_unclear":[],"detected_language":"en"}"#,
14533 "send_report",
14534 r#"{"is_ambiguous":true,"confidence":0.2,"ambiguity_type":"missing_target","reasoning":"target missing","what_is_unclear":["target"],"detected_language":"en"}"#,
14535 r#"{"question":"What should I send?","options":null}"#,
14536 r#"{"status":"answered","selected_option":null,"enriched_input":"Send the report to Ada","resolved":{"intent":"send_report"}}"#,
14537 r#"{"question":"Should I send the report to Ada?"}"#,
14538 "Confirmation cancelled.",
14539 ],
14540 true,
14541 None,
14542 true,
14543 vec![confirmation_skill()],
14544 );
14545
14546 agent.chat("Send it").await.unwrap();
14547 agent.chat("The report to Ada").await.unwrap();
14548 let state_machine = agent.state_machine().unwrap();
14549 state_machine
14550 .transition_to("review", "external test")
14551 .unwrap();
14552 state_machine
14553 .transition_to("active", "external test")
14554 .unwrap();
14555
14556 let response = agent.chat("Yes").await.unwrap();
14557
14558 assert_eq!(response.content, "Confirmation cancelled.");
14559 assert_eq!(confirmation_skill_call_count(&observed), 0);
14560 assert!(agent.pending_skill_id.read().is_none());
14561 }
14562
14563 #[tokio::test]
14564 async fn state_confirmation_does_not_add_a_question_for_clear_input() {
14565 let (agent, observed) = state_disambiguation_agent(
14566 vec![
14567 r#"{"is_ambiguous":false,"confidence":0.99,"ambiguity_type":null,"reasoning":"clear","what_is_unclear":[],"detected_language":"en"}"#,
14568 "Request executed.",
14569 ],
14570 true,
14571 None,
14572 true,
14573 );
14574
14575 let response = agent.chat("Send the report to Ada").await.unwrap();
14576
14577 assert_eq!(response.content, "Request executed.");
14578 assert_eq!(observed.call_count(), 2);
14579 }
14580
14581 #[tokio::test]
14582 async fn state_override_cannot_activate_a_disabled_top_level_manager() {
14583 let (agent, observed) =
14584 state_disambiguation_agent(vec!["Request executed."], false, Some(true), true);
14585
14586 assert!(!agent.has_disambiguation());
14587 let response = agent.chat("Send it").await.unwrap();
14588
14589 assert_eq!(response.content, "Request executed.");
14590 assert_eq!(observed.call_count(), 1);
14591 }
14592
14593 #[tokio::test]
14594 async fn native_required_choice_executes_through_the_shared_tool_path() {
14595 let mut mock = MockLLMProvider::new("native-required");
14596 mock.set_tool_choice(Some(ToolChoice::Required));
14597 let native_call = ToolCall {
14598 id: "provider-call-1".to_string(),
14599 name: "calculator".to_string(),
14600 arguments: serde_json::json!({"expression": "2 + 2"}),
14601 };
14602 let provider_state = ai_agents_core::NativeProviderState::new(
14603 "fixture-exchange-1",
14604 "fixture",
14605 "native-tools",
14606 ai_agents_core::NativeProviderTarget::new("https://fixture.invalid/", "fixture-model")
14607 .unwrap(),
14608 serde_json::json!({
14609 "role": "model",
14610 "parts": [{
14611 "functionCall": {"name": "calculator", "args": {"expression": "2 + 2"}},
14612 "thoughtSignature": "fixture-signature"
14613 }]
14614 }),
14615 vec![ai_agents_core::NativeCallBinding::new("provider-call-1", 0).unwrap()],
14616 )
14617 .unwrap();
14618 mock.add_response(
14619 LLMResponse::new("", FinishReason::ToolCall)
14620 .with_provider_state(provider_state)
14621 .unwrap()
14622 .with_tool_calls(vec![native_call])
14623 .unwrap(),
14624 );
14625 mock.add_response(LLMResponse::new("The answer is 4.", FinishReason::Stop));
14626 let observed = mock.clone();
14627 let agent = AgentBuilder::new()
14628 .system_prompt("Use the calculator when needed.")
14629 .llm(Arc::new(mock))
14630 .tool(Arc::new(CalculatorTool::new()))
14631 .build()
14632 .unwrap();
14633
14634 let response = agent.chat("What is 2 + 2?").await.unwrap();
14635
14636 assert_eq!(response.content, "The answer is 4.");
14637 assert_eq!(
14638 response.tool_calls.as_ref().unwrap()[0].id,
14639 "provider-call-1"
14640 );
14641 let calls = observed.call_history();
14642 assert_eq!(calls.len(), 2);
14643 assert!(matches!(
14644 calls[0].request.as_ref().map(|request| &request.choice),
14645 Some(ToolChoice::Required)
14646 ));
14647 assert!(matches!(
14648 calls[1].request.as_ref().map(|request| &request.choice),
14649 Some(ToolChoice::Auto)
14650 ));
14651 let replay_batch = calls[1]
14652 .messages
14653 .iter()
14654 .find_map(|message| {
14655 ai_agents_core::decode_native_tool_call_markers(&message.content).unwrap()
14656 })
14657 .expect("signed native call marker must be replayed");
14658 assert_eq!(
14659 replay_batch.provider_state().unwrap().exchange_id(),
14660 "fixture-exchange-1"
14661 );
14662 assert!(calls[1].messages.iter().any(|message| {
14663 ai_agents_core::decode_native_tool_result_markers(&message.content)
14664 .is_ok_and(|results| results.is_some())
14665 }));
14666 }
14667
14668 #[tokio::test]
14669 async fn custom_memory_loss_stops_before_signed_tool_execution() {
14670 let mut mock = MockLLMProvider::new("native-custom-memory");
14671 mock.set_tool_choice(Some(ToolChoice::Required));
14672 let call = ToolCall {
14673 id: "provider-call-drop".to_string(),
14674 name: "calculator".to_string(),
14675 arguments: serde_json::json!({"expression": "3 + 4"}),
14676 };
14677 let state = ai_agents_core::NativeProviderState::new(
14678 "fixture-exchange-drop",
14679 "fixture",
14680 "native-tools",
14681 ai_agents_core::NativeProviderTarget::new("https://fixture.invalid/", "fixture-model")
14682 .unwrap(),
14683 serde_json::json!({
14684 "role": "model",
14685 "parts": [{
14686 "functionCall": {"name": "calculator", "args": {"expression": "3 + 4"}},
14687 "thoughtSignature": "fixture-signature-drop"
14688 }]
14689 }),
14690 vec![ai_agents_core::NativeCallBinding::new("provider-call-drop", 0).unwrap()],
14691 )
14692 .unwrap();
14693 mock.add_response(
14694 LLMResponse::new("", FinishReason::ToolCall)
14695 .with_provider_state(state)
14696 .unwrap()
14697 .with_tool_calls(vec![call])
14698 .unwrap(),
14699 );
14700 let agent = AgentBuilder::new()
14701 .system_prompt("Use the calculator.")
14702 .llm(Arc::new(mock))
14703 .memory(Arc::new(DroppingSignedAssistantMemory {
14704 messages: RwLock::new(Vec::new()),
14705 }))
14706 .tool(Arc::new(CalculatorTool::new()))
14707 .build()
14708 .unwrap();
14709
14710 let error = agent.chat("What is 3 + 4?").await.unwrap_err();
14711
14712 assert!(
14713 error
14714 .to_string()
14715 .contains("removed before provider continuation")
14716 );
14717 assert!(agent.tool_call_history.read().is_empty());
14718 }
14719
14720 #[tokio::test]
14721 async fn sequential_signed_history_validates_every_prior_exchange() {
14722 let mut mock = MockLLMProvider::new("native-sequential-memory");
14723 mock.set_tool_choice(Some(ToolChoice::Required));
14724 mock.add_response(signed_calculator_response(
14725 "seq-exchange-1",
14726 "seq-call-1",
14727 "1 + 1",
14728 ));
14729 mock.add_response(signed_calculator_response(
14730 "seq-exchange-2",
14731 "seq-call-2",
14732 "2 + 2",
14733 ));
14734 let agent = AgentBuilder::new()
14735 .system_prompt("Use the calculator sequentially.")
14736 .llm(Arc::new(mock))
14737 .memory(Arc::new(DroppingEarlierSequentialMemory {
14738 messages: RwLock::new(Vec::new()),
14739 signed_seen: std::sync::atomic::AtomicUsize::new(0),
14740 }))
14741 .tool(Arc::new(CalculatorTool::new()))
14742 .build()
14743 .unwrap();
14744
14745 let error = agent.chat("Calculate twice.").await.unwrap_err();
14746
14747 assert!(error.to_string().contains("seq-exchange-1"));
14748 assert_eq!(agent.tool_call_history.read().len(), 1);
14749 }
14750
14751 #[tokio::test]
14752 async fn post_transition_signed_hitl_rejection_stops_before_continuation() {
14753 let mut native = MockLLMProvider::new("post-transition-native");
14754 native.set_tool_choice(Some(ToolChoice::Auto));
14755 let call = ToolCall {
14756 id: "post-transition-call".to_string(),
14757 name: "echo".to_string(),
14758 arguments: serde_json::json!({"message": "hello"}),
14759 };
14760 let state = ai_agents_core::NativeProviderState::new(
14761 "post-transition-exchange",
14762 "fixture",
14763 "native-tools",
14764 ai_agents_core::NativeProviderTarget::new("https://fixture.invalid/", "fixture-model")
14765 .unwrap(),
14766 serde_json::json!({
14767 "role": "model",
14768 "parts": [{
14769 "functionCall": {"name": "echo", "args": {"message": "hello"}},
14770 "thoughtSignature": "post-transition-signature"
14771 }]
14772 }),
14773 vec![ai_agents_core::NativeCallBinding::new("post-transition-call", 0).unwrap()],
14774 )
14775 .unwrap();
14776 native.add_response(
14777 LLMResponse::new("", FinishReason::ToolCall)
14778 .with_provider_state(state)
14779 .unwrap()
14780 .with_tool_calls(vec![call])
14781 .unwrap(),
14782 );
14783 let observed_native = native.clone();
14784 let yaml = r#"
14785name: PostTransitionNativeReject
14786system_prompt: test
14787tools: [echo]
14788hitl:
14789 tools:
14790 echo:
14791 require_approval: true
14792states:
14793 initial: intake
14794 states:
14795 intake:
14796 prompt: intake
14797 transitions:
14798 - to: active
14799 guard:
14800 context:
14801 route:
14802 eq: active
14803 active:
14804 prompt: active
14805 llm: native
14806"#;
14807 let agent = AgentBuilder::from_yaml(yaml)
14808 .unwrap()
14809 .llm(Arc::new(mock_with_response("stale intake response")))
14810 .llm_alias("native", Arc::new(native))
14811 .auto_configure_features()
14812 .unwrap()
14813 .build()
14814 .unwrap();
14815 agent
14816 .set_context("route", serde_json::json!("active"))
14817 .unwrap();
14818
14819 let error = agent.chat("move to active").await.unwrap_err();
14820
14821 assert!(matches!(error, AgentError::HITLRejected(_)));
14822 assert_eq!(observed_native.call_count(), 1);
14823 }
14824
14825 #[test]
14826 fn runtime_overflow_removes_a_past_signed_user_turn_as_one_prefix() {
14827 let call = ToolCall {
14828 id: "overflow-call".to_string(),
14829 name: "calculator".to_string(),
14830 arguments: serde_json::json!({"expression": "1 + 1"}),
14831 };
14832 let state = ai_agents_core::NativeProviderState::new(
14833 "overflow-exchange",
14834 "google",
14835 "generateContent",
14836 ai_agents_core::NativeProviderTarget::new("https://example.invalid/", "gemini-3")
14837 .unwrap(),
14838 serde_json::json!({
14839 "role": "model",
14840 "parts": [{
14841 "functionCall": {"name": "calculator", "args": {"expression": "1 + 1"}},
14842 "thoughtSignature": "overflow-signature"
14843 }]
14844 }),
14845 vec![ai_agents_core::NativeCallBinding::new("overflow-call", 0).unwrap()],
14846 )
14847 .unwrap();
14848 let call_marker = ai_agents_core::encode_native_tool_call_markers(
14849 std::slice::from_ref(&call),
14850 Some(&state),
14851 )
14852 .unwrap();
14853 let result_marker = ai_agents_core::encode_native_tool_result_marker(
14854 &call,
14855 serde_json::json!({"result": 2}),
14856 )
14857 .unwrap();
14858 let history = vec![
14859 ChatMessage::user("old question"),
14860 ChatMessage::assistant(call_marker),
14861 ChatMessage::function("calculator", result_marker),
14862 ChatMessage::assistant("old answer"),
14863 ChatMessage::user("new question"),
14864 ];
14865
14866 let removable = RuntimeAgent::native_safe_prefix_at_least(&history, 1).unwrap();
14867
14868 assert_eq!(removable, 4);
14869 }
14870
14871 #[test]
14872 fn auxiliary_projection_does_not_interpret_user_marker_text() {
14873 let user_text = serde_json::json!({
14874 "_ai_agents_native_tool_call": true,
14875 "id": "",
14876 "tool": "user-data",
14877 "arguments": {}
14878 })
14879 .to_string();
14880
14881 let projected =
14882 RuntimeAgent::readable_native_messages(vec![ChatMessage::user(&user_text)]).unwrap();
14883
14884 assert_eq!(projected[0].content, user_text);
14885 }
14886
14887 #[tokio::test]
14888 async fn terminal_provider_history_error_skips_retry_and_static_fallback() {
14889 let calls = Arc::new(std::sync::atomic::AtomicU32::new(0));
14890 let recovery = RecoveryManager::new(ai_agents_recovery::ErrorRecoveryConfig {
14891 default: ai_agents_recovery::RetryConfig {
14892 max_retries: 3,
14893 ..Default::default()
14894 },
14895 llm: ai_agents_recovery::LLMRecoveryConfig {
14896 on_failure: LLMFailureAction::FallbackResponse {
14897 message: "must not be returned".to_string(),
14898 },
14899 ..Default::default()
14900 },
14901 ..Default::default()
14902 });
14903 let agent = AgentBuilder::new()
14904 .system_prompt("Reject corrupted native history.")
14905 .llm(Arc::new(TerminalHistoryProvider {
14906 calls: Arc::clone(&calls),
14907 }))
14908 .recovery_manager(recovery)
14909 .build()
14910 .unwrap();
14911
14912 let error = agent.chat("continue").await.unwrap_err();
14913
14914 assert!(
14915 error
14916 .to_string()
14917 .contains("native history integrity failure")
14918 );
14919 assert_eq!(calls.load(Ordering::SeqCst), 1);
14920 }
14921
14922 #[tokio::test]
14923 async fn prompt_fallback_uses_one_corrective_retry() {
14924 let mut mock = MockLLMProvider::new("prompt-required");
14925 mock.set_tool_choice(Some(ToolChoice::Required));
14926 mock.set_native_tool_support(false);
14927 mock.set_responses(
14928 vec![
14929 "I can calculate that.".to_string(),
14930 r#"{"tool":"calculator","arguments":{"expression":"2 + 2"}}"#.to_string(),
14931 "The answer is 4.".to_string(),
14932 ],
14933 false,
14934 );
14935 let observed = mock.clone();
14936 let agent = AgentBuilder::new()
14937 .system_prompt("Use tools.")
14938 .llm(Arc::new(mock))
14939 .tool(Arc::new(CalculatorTool::new()))
14940 .build()
14941 .unwrap();
14942
14943 let response = agent.chat("What is 2 + 2?").await.unwrap();
14944
14945 assert_eq!(response.content, "The answer is 4.");
14946 assert_eq!(observed.call_count(), 3);
14947 let corrective = &observed.call_history()[1].messages;
14948 assert!(
14949 corrective
14950 .last()
14951 .unwrap()
14952 .content
14953 .contains("previous response")
14954 );
14955 }
14956
14957 #[tokio::test]
14958 async fn prompt_fallback_fails_after_one_noncompliant_retry() {
14959 let mut mock = MockLLMProvider::new("prompt-required-failure");
14960 mock.set_tool_choice(Some(ToolChoice::Required));
14961 mock.set_native_tool_support(false);
14962 mock.set_responses(
14963 vec!["No tool.".to_string(), "Still no tool.".to_string()],
14964 false,
14965 );
14966 let observed = mock.clone();
14967 let agent = AgentBuilder::new()
14968 .system_prompt("Use tools.")
14969 .llm(Arc::new(mock))
14970 .tool(Arc::new(CalculatorTool::new()))
14971 .build()
14972 .unwrap();
14973
14974 let error = agent.chat("What is 2 + 2?").await.unwrap_err();
14975
14976 assert!(error.to_string().contains("one corrective retry"));
14977 assert_eq!(observed.call_count(), 2);
14978 }
14979
14980 #[tokio::test]
14981 async fn specific_choice_cannot_widen_the_effective_grant() {
14982 let mut mock = MockLLMProvider::new("specific-outside-grant");
14983 mock.set_tool_choice(Some(ToolChoice::Specific("random".to_string())));
14984 let observed = mock.clone();
14985 let agent = AgentBuilder::new()
14986 .system_prompt("Use tools.")
14987 .llm(Arc::new(mock))
14988 .tool(Arc::new(CalculatorTool::new()))
14989 .build()
14990 .unwrap();
14991
14992 let error = agent.chat("Generate a value.").await.unwrap_err();
14993
14994 assert!(error.to_string().contains("is not registered"));
14995 assert_eq!(observed.call_count(), 0);
14996 }
14997
14998 #[tokio::test]
14999 async fn none_choice_exposes_no_tool_protocol() {
15000 let mut mock = MockLLMProvider::new("no-tools");
15001 mock.set_tool_choice(Some(ToolChoice::None));
15002 mock.set_response(r#"{"tool":"calculator","arguments":{"expression":"2 + 2"}}"#);
15003 let observed = mock.clone();
15004 let agent = AgentBuilder::new()
15005 .system_prompt("Answer directly.")
15006 .llm(Arc::new(mock))
15007 .tool(Arc::new(CalculatorTool::new()))
15008 .build()
15009 .unwrap();
15010
15011 let response = agent.chat("Hello").await.unwrap();
15012
15013 assert!(response.tool_calls.is_none());
15014 assert_eq!(observed.call_count(), 1);
15015 let call = observed.last_call().unwrap();
15016 assert!(call.request.is_none());
15017 assert!(
15018 call.messages
15019 .iter()
15020 .all(|message| !message.content.contains("Available tools:"))
15021 );
15022 }
15023
15024 struct RuntimeStorage {
15025 capabilities: Box<[StorageCapability]>,
15026 snapshots: RwLock<HashMap<String, AgentSnapshot>>,
15027 metadata: RwLock<HashMap<String, ai_agents_core::SessionMetadata>>,
15028 metadata_save_calls: AtomicU64,
15029 metadata_load_calls: AtomicU64,
15030 fail_metadata_save: AtomicBool,
15031 fail_metadata_load: AtomicBool,
15032 }
15033
15034 impl RuntimeStorage {
15035 fn new(capabilities: impl IntoIterator<Item = StorageCapability>) -> Self {
15036 Self {
15037 capabilities: capabilities.into_iter().collect(),
15038 snapshots: RwLock::new(HashMap::new()),
15039 metadata: RwLock::new(HashMap::new()),
15040 metadata_save_calls: AtomicU64::new(0),
15041 metadata_load_calls: AtomicU64::new(0),
15042 fail_metadata_save: AtomicBool::new(false),
15043 fail_metadata_load: AtomicBool::new(false),
15044 }
15045 }
15046 }
15047
15048 #[async_trait]
15049 impl AgentStorage for RuntimeStorage {
15050 fn supports(&self, capability: StorageCapability) -> bool {
15051 self.capabilities.contains(&capability)
15052 }
15053
15054 async fn save(&self, session_id: &str, snapshot: &AgentSnapshot) -> Result<()> {
15055 self.snapshots
15056 .write()
15057 .insert(session_id.to_string(), snapshot.clone());
15058 Ok(())
15059 }
15060
15061 async fn load(&self, session_id: &str) -> Result<Option<AgentSnapshot>> {
15062 Ok(self.snapshots.read().get(session_id).cloned())
15063 }
15064
15065 async fn delete(&self, session_id: &str) -> Result<()> {
15066 self.snapshots.write().remove(session_id);
15067 Ok(())
15068 }
15069
15070 async fn list_sessions(&self) -> Result<Vec<String>> {
15071 Ok(self.snapshots.read().keys().cloned().collect())
15072 }
15073
15074 async fn save_snapshot_with_metadata(
15075 &self,
15076 session_id: &str,
15077 snapshot: &AgentSnapshot,
15078 metadata: &ai_agents_core::SessionMetadata,
15079 ) -> Result<()> {
15080 self.metadata_save_calls.fetch_add(1, Ordering::SeqCst);
15081 if self.fail_metadata_save.load(Ordering::SeqCst) {
15082 return Err(AgentError::Persistence("metadata save failed".into()));
15083 }
15084 self.snapshots
15085 .write()
15086 .insert(session_id.to_string(), snapshot.clone());
15087 self.metadata
15088 .write()
15089 .insert(session_id.to_string(), metadata.clone());
15090 Ok(())
15091 }
15092
15093 async fn save_metadata(
15094 &self,
15095 session_id: &str,
15096 metadata: &ai_agents_core::SessionMetadata,
15097 ) -> Result<()> {
15098 self.metadata_save_calls.fetch_add(1, Ordering::SeqCst);
15099 if self.fail_metadata_save.load(Ordering::SeqCst) {
15100 return Err(AgentError::Persistence("metadata save failed".into()));
15101 }
15102 self.metadata
15103 .write()
15104 .insert(session_id.to_string(), metadata.clone());
15105 Ok(())
15106 }
15107
15108 async fn load_metadata(
15109 &self,
15110 session_id: &str,
15111 ) -> Result<Option<ai_agents_core::SessionMetadata>> {
15112 self.metadata_load_calls.fetch_add(1, Ordering::SeqCst);
15113 if self.fail_metadata_load.load(Ordering::SeqCst) {
15114 return Err(AgentError::Persistence("metadata load failed".into()));
15115 }
15116 Ok(self.metadata.read().get(session_id).cloned())
15117 }
15118 }
15119
15120 fn runtime_storage_agent() -> RuntimeAgent {
15121 AgentBuilder::new()
15122 .system_prompt("Test runtime storage integration.")
15123 .llm(Arc::new(mock_with_response("done")))
15124 .build()
15125 .unwrap()
15126 }
15127
15128 fn restore_spec(id: &str) -> crate::spec::AgentSpec {
15129 crate::spec::AgentSpec {
15130 name: id.to_string(),
15131 system_prompt: format!("Restore child {id}."),
15132 ..crate::spec::AgentSpec::default()
15133 }
15134 }
15135
15136 fn restore_entry(id: &str) -> ai_agents_core::SpawnedAgentEntry {
15137 ai_agents_core::SpawnedAgentEntry {
15138 id: id.to_string(),
15139 name: id.to_string(),
15140 spec_yaml: serde_yaml::to_string(&restore_spec(id)).unwrap(),
15141 }
15142 }
15143
15144 fn restore_spawner(
15145 storage: Arc<RuntimeStorage>,
15146 max_agents: usize,
15147 ) -> (
15148 Arc<crate::spawner::AgentSpawner>,
15149 Arc<crate::spawner::AgentRegistry>,
15150 ) {
15151 let mut llms = LLMRegistry::new();
15152 llms.register("default", Arc::new(mock_with_response("done")));
15153 (
15154 Arc::new(
15155 crate::spawner::AgentSpawner::new()
15156 .with_shared_llms(llms)
15157 .with_shared_storage(storage)
15158 .with_max_agents(max_agents),
15159 ),
15160 Arc::new(crate::spawner::AgentRegistry::new()),
15161 )
15162 }
15163
15164 async fn save_restore_target(
15165 parent: &RuntimeAgent,
15166 storage: &RuntimeStorage,
15167 session_id: &str,
15168 entries: Vec<ai_agents_core::SpawnedAgentEntry>,
15169 ) {
15170 let mut snapshot = parent.save_state().await.unwrap();
15171 snapshot.spawned_agents = Some(entries);
15172 storage.save(session_id, &snapshot).await.unwrap();
15173 storage
15174 .save_metadata(session_id, &ai_agents_core::SessionMetadata::default())
15175 .await
15176 .unwrap();
15177 }
15178
15179 #[tokio::test]
15180 async fn storage_init_requires_storage_for_actor_facts() {
15181 let facts = ai_agents_facts::FactsConfig {
15182 enabled: true,
15183 ..Default::default()
15184 };
15185 let agent = runtime_storage_agent().with_facts_config(None, Some(facts));
15186
15187 let error = agent.init_storage().await.unwrap_err();
15188 assert!(matches!(
15189 error,
15190 AgentError::Config(message)
15191 if message.contains("actor facts or actor memory")
15192 && message.contains("none is configured or injected")
15193 ));
15194 }
15195
15196 #[tokio::test]
15197 async fn storage_init_validates_actor_facts_capability() {
15198 let storage = Arc::new(RuntimeStorage::new([StorageCapability::Snapshot]));
15199 let actor_memory = ai_agents_facts::ActorMemoryConfig {
15200 enabled: true,
15201 ..Default::default()
15202 };
15203 let agent = runtime_storage_agent()
15204 .with_storage(storage)
15205 .with_facts_config(Some(actor_memory), None);
15206
15207 assert!(matches!(
15208 agent.init_storage().await,
15209 Err(AgentError::UnsupportedStorageCapability(
15210 StorageCapability::ActorFacts
15211 ))
15212 ));
15213 }
15214
15215 #[tokio::test]
15216 async fn blocking_chat_rejects_unsupported_required_storage() {
15217 let storage = Arc::new(RuntimeStorage::new([StorageCapability::Snapshot]));
15218 let facts = ai_agents_facts::FactsConfig {
15219 enabled: true,
15220 ..Default::default()
15221 };
15222 let agent = runtime_storage_agent()
15223 .with_storage(storage)
15224 .with_facts_config(None, Some(facts));
15225
15226 assert!(matches!(
15227 agent.chat("hello").await,
15228 Err(AgentError::UnsupportedStorageCapability(
15229 StorageCapability::ActorFacts
15230 ))
15231 ));
15232 }
15233
15234 #[tokio::test]
15235 async fn streaming_chat_rejects_unsupported_required_storage_before_stream_creation() {
15236 let storage = Arc::new(RuntimeStorage::new([StorageCapability::Snapshot]));
15237 let config = ai_agents_relationships::RelationshipConfig {
15238 enabled: true,
15239 ..Default::default()
15240 };
15241 let manager = Arc::new(RelationshipManager::from_config(config).unwrap());
15242 let agent = runtime_storage_agent()
15243 .with_storage(storage)
15244 .with_relationships(manager);
15245
15246 assert!(matches!(
15247 agent.chat_stream("hello").await,
15248 Err(AgentError::UnsupportedStorageCapability(
15249 StorageCapability::ActorRelationships
15250 ))
15251 ));
15252 }
15253
15254 #[tokio::test]
15255 async fn storage_init_completes_facts_for_injected_storage() {
15256 let storage = Arc::new(RuntimeStorage::new([
15257 StorageCapability::Snapshot,
15258 StorageCapability::ActorFacts,
15259 ]));
15260 let facts = ai_agents_facts::FactsConfig {
15261 enabled: true,
15262 ..Default::default()
15263 };
15264 let agent = runtime_storage_agent()
15265 .with_storage(storage)
15266 .with_facts_config(None, Some(facts));
15267
15268 agent.init_storage().await.unwrap();
15269 assert!(agent.fact_store().is_some());
15270 }
15271
15272 #[tokio::test]
15273 async fn storage_init_requires_storage_for_persistent_relationships() {
15274 let config = ai_agents_relationships::RelationshipConfig {
15275 enabled: true,
15276 ..Default::default()
15277 };
15278 let manager = Arc::new(RelationshipManager::from_config(config).unwrap());
15279 let agent = runtime_storage_agent().with_relationships(manager);
15280
15281 let error = agent.init_storage().await.unwrap_err();
15282 assert!(matches!(
15283 error,
15284 AgentError::Config(message)
15285 if message.contains("persistent relationships")
15286 && message.contains("none is configured or injected")
15287 ));
15288 }
15289
15290 #[tokio::test]
15291 async fn storage_init_validates_persistent_relationships_capability() {
15292 let storage = Arc::new(RuntimeStorage::new([StorageCapability::Snapshot]));
15293 let config = ai_agents_relationships::RelationshipConfig {
15294 enabled: true,
15295 ..Default::default()
15296 };
15297 let manager = Arc::new(RelationshipManager::from_config(config).unwrap());
15298 let agent = runtime_storage_agent()
15299 .with_storage(storage)
15300 .with_relationships(manager);
15301
15302 assert!(matches!(
15303 agent.init_storage().await,
15304 Err(AgentError::UnsupportedStorageCapability(
15305 StorageCapability::ActorRelationships
15306 ))
15307 ));
15308 }
15309
15310 #[tokio::test]
15311 async fn session_restore_updates_identity_and_clears_stale_actor_binding() {
15312 let storage = Arc::new(RuntimeStorage::new([
15313 StorageCapability::Snapshot,
15314 StorageCapability::SessionMetadata,
15315 ]));
15316 let agent = runtime_storage_agent().with_storage(storage.clone());
15317 agent.set_actor_id("old-actor").unwrap();
15318 agent.save_session("old").await.unwrap();
15319 storage
15320 .save("target", &agent.save_state().await.unwrap())
15321 .await
15322 .unwrap();
15323 storage
15324 .save_metadata("target", &ai_agents_core::SessionMetadata::default())
15325 .await
15326 .unwrap();
15327
15328 assert!(agent.load_session("target").await.unwrap());
15329
15330 assert_eq!(agent.current_session_id.read().as_deref(), Some("target"));
15331 assert_eq!(agent.actor_id(), None);
15332 }
15333
15334 #[tokio::test]
15335 async fn complete_restore_reconciles_growth_shrink_and_empty_topologies() {
15336 let storage = Arc::new(RuntimeStorage::new([
15337 StorageCapability::Snapshot,
15338 StorageCapability::SessionMetadata,
15339 ]));
15340 let (spawner, registry) = restore_spawner(storage.clone(), 3);
15341 let parent = runtime_storage_agent()
15342 .with_storage(storage.clone())
15343 .with_spawner_handles(Arc::clone(&spawner), Arc::clone(®istry));
15344
15345 for id in ["a", "b"] {
15346 let spawned = spawner
15347 .spawn_with_id(id.to_string(), restore_spec(id))
15348 .await
15349 .unwrap();
15350 spawned.agent.save_session("grow").await.unwrap();
15351 registry.register(spawned).await.unwrap();
15352 }
15353 let staged_c = crate::spawner::storage::NamespacedStorage::new(storage.clone(), "c");
15354 staged_c
15355 .save("grow", &AgentSnapshot::new("c".into()))
15356 .await
15357 .unwrap();
15358 staged_c
15359 .save_metadata("grow", &ai_agents_core::SessionMetadata::default())
15360 .await
15361 .unwrap();
15362 save_restore_target(
15363 &parent,
15364 storage.as_ref(),
15365 "grow",
15366 vec![restore_entry("a"), restore_entry("b"), restore_entry("c")],
15367 )
15368 .await;
15369
15370 assert_eq!(parent.restore_session_full("grow").await.unwrap(), 3);
15371 assert_eq!(registry.count(), 3);
15372 assert!(registry.contains("c"));
15373 assert_eq!(spawner.spawned_count(), 3);
15374
15375 for id in ["a", "b"] {
15376 registry
15377 .get(id)
15378 .unwrap()
15379 .save_session("shrink")
15380 .await
15381 .unwrap();
15382 }
15383 save_restore_target(
15384 &parent,
15385 storage.as_ref(),
15386 "shrink",
15387 vec![restore_entry("a"), restore_entry("b")],
15388 )
15389 .await;
15390
15391 assert_eq!(parent.restore_session_full("shrink").await.unwrap(), 2);
15392 assert_eq!(registry.count(), 2);
15393 assert!(!registry.contains("c"));
15394 assert_eq!(spawner.spawned_count(), 2);
15395
15396 save_restore_target(&parent, storage.as_ref(), "empty", Vec::new()).await;
15397
15398 assert_eq!(parent.restore_session_full("empty").await.unwrap(), 0);
15399 assert_eq!(registry.count(), 0);
15400 assert_eq!(spawner.spawned_count(), 0);
15401 assert_eq!(parent.current_session_id.read().as_deref(), Some("empty"));
15402 }
15403
15404 #[tokio::test]
15405 async fn storage_session_metadata_is_called_only_when_advertised() {
15406 let storage = Arc::new(RuntimeStorage::new([StorageCapability::Snapshot]));
15407 storage.fail_metadata_save.store(true, Ordering::SeqCst);
15408 storage.fail_metadata_load.store(true, Ordering::SeqCst);
15409 let agent = runtime_storage_agent().with_storage(storage.clone());
15410
15411 agent.save_session("session").await.unwrap();
15412 assert!(agent.load_session("session").await.unwrap());
15413 assert_eq!(storage.metadata_save_calls.load(Ordering::SeqCst), 0);
15414 assert_eq!(storage.metadata_load_calls.load(Ordering::SeqCst), 0);
15415 }
15416
15417 #[cfg(feature = "sqlite")]
15418 #[tokio::test]
15419 async fn sqlite_runtime_save_filter_reopen_and_reload_stay_consistent() {
15420 let directory =
15421 std::env::temp_dir().join(format!("ai-agents-runtime-sqlite-{}", uuid::Uuid::new_v4()));
15422 let path = directory.join("sessions.sqlite");
15423 let path_string = path.to_string_lossy().into_owned();
15424 let storage = Arc::new(
15425 ai_agents_storage::SqliteStorage::new(&path_string)
15426 .await
15427 .unwrap(),
15428 );
15429 let agent = runtime_storage_agent().with_storage(storage.clone());
15430 agent.set_session_metadata(ai_agents_core::SessionMetadata {
15431 tags: vec!["initial".into()],
15432 ..Default::default()
15433 });
15434 agent.chat("persist this turn").await.unwrap();
15435 agent.save_session("session").await.unwrap();
15436
15437 agent.set_session_metadata(ai_agents_core::SessionMetadata {
15438 tags: vec!["updated".into()],
15439 ..Default::default()
15440 });
15441 agent.save_session("session").await.unwrap();
15442 assert!(
15443 agent
15444 .list_sessions_filtered(&ai_agents_core::SessionFilter {
15445 tags: Some(vec!["initial".into()]),
15446 ..Default::default()
15447 })
15448 .await
15449 .unwrap()
15450 .is_empty()
15451 );
15452 assert_eq!(
15453 agent
15454 .list_sessions_filtered(&ai_agents_core::SessionFilter {
15455 tags: Some(vec!["updated".into()]),
15456 ..Default::default()
15457 })
15458 .await
15459 .unwrap()
15460 .len(),
15461 1
15462 );
15463 drop(agent);
15464 storage.close().await;
15465 drop(storage);
15466
15467 let reopened_storage = Arc::new(
15468 ai_agents_storage::SqliteStorage::new(&path_string)
15469 .await
15470 .unwrap(),
15471 );
15472 let restored = runtime_storage_agent().with_storage(reopened_storage.clone());
15473 assert!(restored.load_session("session").await.unwrap());
15474 assert_eq!(restored.session_metadata().tags, vec!["updated"]);
15475 assert_eq!(
15476 restored.current_session_id.read().as_deref(),
15477 Some("session")
15478 );
15479 assert!(restored.save_state().await.unwrap().memory.messages.len() >= 2);
15480 assert_eq!(
15481 restored
15482 .list_sessions_filtered(&ai_agents_core::SessionFilter {
15483 tags: Some(vec!["updated".into()]),
15484 ..Default::default()
15485 })
15486 .await
15487 .unwrap()
15488 .len(),
15489 1
15490 );
15491
15492 drop(restored);
15493 reopened_storage.close().await;
15494 drop(reopened_storage);
15495 crate::remove_sqlite_test_directory(&directory)
15496 .await
15497 .unwrap();
15498 }
15499
15500 #[tokio::test]
15501 async fn storage_session_metadata_backend_failures_propagate() {
15502 let storage = Arc::new(RuntimeStorage::new([
15503 StorageCapability::Snapshot,
15504 StorageCapability::SessionMetadata,
15505 ]));
15506 let agent = runtime_storage_agent().with_storage(storage.clone());
15507
15508 agent.save_session("session").await.unwrap();
15509 storage
15510 .save("target", &agent.save_state().await.unwrap())
15511 .await
15512 .unwrap();
15513 storage.fail_metadata_load.store(true, Ordering::SeqCst);
15514 assert!(matches!(
15515 agent.load_session("target").await,
15516 Err(AgentError::Persistence(message)) if message == "metadata load failed"
15517 ));
15518 assert_eq!(agent.current_session_id.read().as_deref(), Some("session"));
15519
15520 storage.fail_metadata_save.store(true, Ordering::SeqCst);
15521 assert!(matches!(
15522 agent.save_session("session").await,
15523 Err(AgentError::Persistence(message)) if message == "metadata save failed"
15524 ));
15525 }
15526
15527 struct ProviderFutureDropSignal {
15528 dropped: Arc<AtomicBool>,
15529 }
15530
15531 impl Drop for ProviderFutureDropSignal {
15532 fn drop(&mut self) {
15533 self.dropped.store(true, Ordering::SeqCst);
15534 }
15535 }
15536
15537 struct BufferedLockingProvider {
15538 lock: Arc<tokio::sync::Mutex<()>>,
15539 stream_started: Arc<tokio::sync::Notify>,
15540 stream_dropped: Arc<AtomicBool>,
15541 committed_after_drop: Arc<AtomicBool>,
15542 }
15543
15544 #[async_trait]
15545 impl LLMProvider for BufferedLockingProvider {
15546 async fn complete(
15547 &self,
15548 _messages: &[ChatMessage],
15549 _config: Option<&LLMConfig>,
15550 ) -> std::result::Result<LLMResponse, LLMError> {
15551 let _guard = self.lock.lock().await;
15552 self.committed_after_drop
15553 .store(self.stream_dropped.load(Ordering::SeqCst), Ordering::SeqCst);
15554 Ok(LLMResponse::new(
15555 "Committed technical response.",
15556 FinishReason::Stop,
15557 ))
15558 }
15559
15560 async fn complete_stream(
15561 &self,
15562 _messages: &[ChatMessage],
15563 _config: Option<&LLMConfig>,
15564 ) -> std::result::Result<
15565 Box<dyn Stream<Item = std::result::Result<LLMChunk, LLMError>> + Unpin + Send>,
15566 LLMError,
15567 > {
15568 let _guard = self.lock.lock().await;
15569 let _drop_signal = ProviderFutureDropSignal {
15570 dropped: Arc::clone(&self.stream_dropped),
15571 };
15572 self.stream_started.notify_one();
15573 std::future::pending().await
15574 }
15575
15576 fn provider_name(&self) -> &str {
15577 "buffered-locking"
15578 }
15579
15580 fn supports(&self, _feature: LLMFeature) -> bool {
15581 false
15582 }
15583 }
15584
15585 struct PendingDropStream {
15586 dropped: Arc<AtomicBool>,
15587 dropped_notify: Arc<tokio::sync::Notify>,
15588 }
15589
15590 impl Stream for PendingDropStream {
15591 type Item = std::result::Result<LLMChunk, LLMError>;
15592
15593 fn poll_next(
15594 self: Pin<&mut Self>,
15595 _cx: &mut std::task::Context<'_>,
15596 ) -> std::task::Poll<Option<Self::Item>> {
15597 std::task::Poll::Pending
15598 }
15599 }
15600
15601 impl Drop for PendingDropStream {
15602 fn drop(&mut self) {
15603 self.dropped.store(true, Ordering::SeqCst);
15604 self.dropped_notify.notify_one();
15605 }
15606 }
15607
15608 struct EstablishedStreamProvider {
15609 stream_started: Arc<tokio::sync::Notify>,
15610 stream_dropped: Arc<AtomicBool>,
15611 stream_dropped_notify: Arc<tokio::sync::Notify>,
15612 committed_after_drop: Arc<AtomicBool>,
15613 }
15614
15615 #[async_trait]
15616 impl LLMProvider for EstablishedStreamProvider {
15617 async fn complete(
15618 &self,
15619 _messages: &[ChatMessage],
15620 _config: Option<&LLMConfig>,
15621 ) -> std::result::Result<LLMResponse, LLMError> {
15622 if !self.stream_dropped.load(Ordering::SeqCst) {
15623 self.stream_dropped_notify.notified().await;
15624 }
15625 self.committed_after_drop
15626 .store(self.stream_dropped.load(Ordering::SeqCst), Ordering::SeqCst);
15627 Ok(LLMResponse::new(
15628 "Committed technical response.",
15629 FinishReason::Stop,
15630 ))
15631 }
15632
15633 async fn complete_stream(
15634 &self,
15635 _messages: &[ChatMessage],
15636 _config: Option<&LLMConfig>,
15637 ) -> std::result::Result<
15638 Box<dyn Stream<Item = std::result::Result<LLMChunk, LLMError>> + Unpin + Send>,
15639 LLMError,
15640 > {
15641 self.stream_started.notify_one();
15642 Ok(Box::new(PendingDropStream {
15643 dropped: Arc::clone(&self.stream_dropped),
15644 dropped_notify: Arc::clone(&self.stream_dropped_notify),
15645 }))
15646 }
15647
15648 fn provider_name(&self) -> &str {
15649 "established-stream"
15650 }
15651
15652 fn supports(&self, _feature: LLMFeature) -> bool {
15653 false
15654 }
15655 }
15656
15657 struct FirstCallLockingProvider {
15658 lock: Arc<tokio::sync::Mutex<()>>,
15659 first_started: Arc<tokio::sync::Notify>,
15660 first_dropped: Arc<AtomicBool>,
15661 committed_after_drop: Arc<AtomicBool>,
15662 calls: AtomicU64,
15663 }
15664
15665 #[async_trait]
15666 impl LLMProvider for FirstCallLockingProvider {
15667 async fn complete(
15668 &self,
15669 _messages: &[ChatMessage],
15670 _config: Option<&LLMConfig>,
15671 ) -> std::result::Result<LLMResponse, LLMError> {
15672 let _guard = self.lock.lock().await;
15673 let call = self.calls.fetch_add(1, Ordering::SeqCst);
15674 if call == 0 {
15675 let _drop_signal = ProviderFutureDropSignal {
15676 dropped: Arc::clone(&self.first_dropped),
15677 };
15678 self.first_started.notify_one();
15679 return std::future::pending().await;
15680 }
15681 self.committed_after_drop
15682 .store(self.first_dropped.load(Ordering::SeqCst), Ordering::SeqCst);
15683 Ok(LLMResponse::new(
15684 "Committed technical response.",
15685 FinishReason::Stop,
15686 ))
15687 }
15688
15689 async fn complete_stream(
15690 &self,
15691 _messages: &[ChatMessage],
15692 _config: Option<&LLMConfig>,
15693 ) -> std::result::Result<
15694 Box<dyn Stream<Item = std::result::Result<LLMChunk, LLMError>> + Unpin + Send>,
15695 LLMError,
15696 > {
15697 Err(LLMError::Other(
15698 "streaming is not used in this test".to_string(),
15699 ))
15700 }
15701
15702 fn provider_name(&self) -> &str {
15703 "first-call-locking"
15704 }
15705
15706 fn supports(&self, _feature: LLMFeature) -> bool {
15707 false
15708 }
15709 }
15710
15711 struct RoutingAfterProviderStart {
15712 provider_started: Arc<tokio::sync::Notify>,
15713 }
15714
15715 #[async_trait]
15716 impl LLMProvider for RoutingAfterProviderStart {
15717 async fn complete(
15718 &self,
15719 _messages: &[ChatMessage],
15720 _config: Option<&LLMConfig>,
15721 ) -> std::result::Result<LLMResponse, LLMError> {
15722 self.provider_started.notified().await;
15723 Ok(LLMResponse::new("1", FinishReason::Stop))
15724 }
15725
15726 async fn complete_stream(
15727 &self,
15728 _messages: &[ChatMessage],
15729 _config: Option<&LLMConfig>,
15730 ) -> std::result::Result<
15731 Box<dyn Stream<Item = std::result::Result<LLMChunk, LLMError>> + Unpin + Send>,
15732 LLMError,
15733 > {
15734 Err(LLMError::Other(
15735 "streaming is not used in this test".to_string(),
15736 ))
15737 }
15738
15739 fn provider_name(&self) -> &str {
15740 "routing-after-start"
15741 }
15742
15743 fn supports(&self, _feature: LLMFeature) -> bool {
15744 false
15745 }
15746 }
15747
15748 struct ResponseCountingHooks {
15750 responses: Arc<std::sync::atomic::AtomicUsize>,
15751 }
15752
15753 struct RootTurnProbeProvider {
15755 complete_entered: tokio::sync::mpsc::UnboundedSender<()>,
15756 }
15757
15758 struct ResponseChatHooks {
15760 target: parking_lot::Mutex<Option<Weak<RuntimeAgent>>>,
15761 invoked: AtomicBool,
15762 nested_result: parking_lot::Mutex<Option<std::result::Result<String, String>>>,
15763 }
15764
15765 struct ConcurrentResponseHooks {
15767 registry: Weak<crate::spawner::AgentRegistry>,
15768 child_id: String,
15769 invoked: AtomicBool,
15770 nested_result: parking_lot::Mutex<Option<std::result::Result<String, String>>>,
15771 }
15772
15773 struct RetryDeadlineTool {
15775 calls: Arc<std::sync::atomic::AtomicUsize>,
15776 deadlines: Arc<parking_lot::Mutex<Vec<chrono::DateTime<chrono::Utc>>>>,
15777 remaining_ms: Arc<parking_lot::Mutex<Vec<i64>>>,
15778 }
15779
15780 struct ToolLifecycleRecordingHooks {
15782 events: parking_lot::Mutex<Vec<String>>,
15783 records: parking_lot::Mutex<Vec<ToolExecutionRecord>>,
15784 }
15785
15786 impl ToolLifecycleRecordingHooks {
15787 fn new() -> Self {
15789 Self {
15790 events: parking_lot::Mutex::new(Vec::new()),
15791 records: parking_lot::Mutex::new(Vec::new()),
15792 }
15793 }
15794
15795 fn events(&self) -> Vec<String> {
15797 self.events.lock().clone()
15798 }
15799
15800 fn records(&self) -> Vec<ToolExecutionRecord> {
15802 self.records.lock().clone()
15803 }
15804 }
15805
15806 struct ContextEchoTool;
15808
15809 #[async_trait]
15810 impl LLMProvider for RootTurnProbeProvider {
15811 async fn complete(
15812 &self,
15813 _messages: &[ChatMessage],
15814 _config: Option<&LLMConfig>,
15815 ) -> std::result::Result<LLMResponse, LLMError> {
15816 let _ = self.complete_entered.send(());
15817 Ok(LLMResponse::new("blocking complete", FinishReason::Stop))
15818 }
15819
15820 async fn complete_stream(
15821 &self,
15822 _messages: &[ChatMessage],
15823 _config: Option<&LLMConfig>,
15824 ) -> std::result::Result<
15825 Box<dyn Stream<Item = std::result::Result<LLMChunk, LLMError>> + Unpin + Send>,
15826 LLMError,
15827 > {
15828 Ok(Box::new(futures::stream::iter(vec![Ok(
15829 LLMChunk::final_chunk("stream complete", FinishReason::Stop, None),
15830 )])))
15831 }
15832
15833 fn provider_name(&self) -> &str {
15834 "root-turn-probe"
15835 }
15836
15837 fn supports(&self, feature: LLMFeature) -> bool {
15838 matches!(feature, LLMFeature::Streaming)
15839 }
15840 }
15841
15842 #[async_trait]
15843 impl ai_agents_core::Tool for ContextEchoTool {
15844 fn id(&self) -> &str {
15845 "context_echo"
15846 }
15847
15848 fn name(&self) -> &str {
15849 "Context Echo"
15850 }
15851
15852 fn description(&self) -> &str {
15853 "Returns selected execution context fields."
15854 }
15855
15856 fn input_schema(&self) -> Value {
15857 serde_json::json!({"type": "object"})
15858 }
15859
15860 fn policy_bindings(&self) -> ai_agents_core::ToolPolicyBindings {
15861 ai_agents_core::ToolPolicyBindings {
15862 path_fields: vec![ai_agents_core::PathPolicyBinding::read("path")],
15863 result_limit_fields: vec![ai_agents_core::ResultLimitBinding::new(
15864 "max_results",
15865 ai_agents_core::ResultLimitKind::MaxResults,
15866 )],
15867 ..Default::default()
15868 }
15869 }
15870
15871 async fn execute(
15872 &self,
15873 _args: Value,
15874 ctx: ai_agents_core::ToolExecutionContext,
15875 ) -> ToolResult {
15876 ToolResult::ok(
15877 serde_json::json!({
15878 "requested_name": ctx.requested_name,
15879 "canonical_id": ctx.canonical_id,
15880 "display_name": ctx.display_name,
15881 "max_results": ctx.limits.max_results,
15882 "custom_config": ctx.custom_config,
15883 })
15884 .to_string(),
15885 )
15886 }
15887 }
15888
15889 #[async_trait]
15890 impl ai_agents_core::Tool for RetryDeadlineTool {
15891 fn id(&self) -> &str {
15892 "retry_deadline"
15893 }
15894
15895 fn name(&self) -> &str {
15896 "Retry Deadline"
15897 }
15898
15899 fn description(&self) -> &str {
15900 "Records one deadline per retry invocation."
15901 }
15902
15903 fn input_schema(&self) -> Value {
15904 serde_json::json!({"type": "object"})
15905 }
15906
15907 fn safety_metadata(&self) -> ai_agents_core::ToolSafetyMetadata {
15908 ai_agents_core::ToolSafetyMetadata::compute()
15909 }
15910
15911 fn classify_call(&self, _args: &Value) -> ai_agents_core::ToolCallClassification {
15912 let mut classification =
15913 ai_agents_core::ToolCallClassification::from_metadata(&self.safety_metadata());
15914 classification.timeout_ms = Some(1_000);
15915 classification.safely_retryable = true;
15916 classification
15917 }
15918
15919 async fn execute(
15921 &self,
15922 _args: Value,
15923 ctx: ai_agents_core::ToolExecutionContext,
15924 ) -> ToolResult {
15925 let deadline = ctx
15926 .deadline
15927 .expect("each invocation must receive a deadline");
15928 self.remaining_ms.lock().push(
15929 deadline
15930 .signed_duration_since(chrono::Utc::now())
15931 .num_milliseconds(),
15932 );
15933 self.deadlines.lock().push(deadline);
15934 let call = self.calls.fetch_add(1, Ordering::SeqCst);
15935 if call == 0 {
15936 tokio::time::sleep(std::time::Duration::from_millis(20)).await;
15937 ToolResult::error("retry")
15938 } else {
15939 ToolResult::ok("done")
15940 }
15941 }
15942 }
15943
15944 struct ClassifiedTimeoutTool {
15946 id: &'static str,
15947 calls: Arc<std::sync::atomic::AtomicUsize>,
15948 timeout_ms: u64,
15949 sleep_ms: u64,
15950 requires_approval: bool,
15951 remaining_ms: Arc<parking_lot::Mutex<Vec<i64>>>,
15952 }
15953
15954 struct ApprovalModifiedTimeoutTool {
15956 calls: Arc<std::sync::atomic::AtomicUsize>,
15957 }
15958
15959 struct SlowTool;
15961
15962 struct FlakyWriteTool {
15964 calls: Arc<std::sync::atomic::AtomicUsize>,
15965 }
15966
15967 struct LockedWriteTool {
15969 active: Arc<std::sync::atomic::AtomicUsize>,
15970 max_active: Arc<std::sync::atomic::AtomicUsize>,
15971 }
15972
15973 struct MultiResourceWriteTool {
15974 active: Arc<std::sync::atomic::AtomicUsize>,
15975 max_active: Arc<std::sync::atomic::AtomicUsize>,
15976 }
15977
15978 #[derive(Clone)]
15979 struct PathMutationGate {
15980 entered: Arc<AtomicBool>,
15981 entered_notify: Arc<tokio::sync::Notify>,
15982 release: Arc<tokio::sync::Notify>,
15983 }
15984
15985 impl PathMutationGate {
15986 fn new() -> Self {
15987 Self {
15988 entered: Arc::new(AtomicBool::new(false)),
15989 entered_notify: Arc::new(tokio::sync::Notify::new()),
15990 release: Arc::new(tokio::sync::Notify::new()),
15991 }
15992 }
15993
15994 async fn wait_until_entered(&self) {
15995 if !self.entered.load(Ordering::SeqCst) {
15996 self.entered_notify.notified().await;
15997 }
15998 }
15999
16000 fn release(&self) {
16001 self.release.notify_one();
16002 }
16003 }
16004
16005 struct BlockingPathMutationTool {
16006 id: &'static str,
16007 path_fields: Vec<ai_agents_core::PathPolicyBinding>,
16008 gate: PathMutationGate,
16009 }
16010
16011 struct NoBindingWriteTool {
16012 active: Arc<std::sync::atomic::AtomicUsize>,
16013 max_active: Arc<std::sync::atomic::AtomicUsize>,
16014 }
16015
16016 struct RecoveryTestTool {
16017 id: String,
16018 succeeds: bool,
16019 calls: Arc<std::sync::atomic::AtomicUsize>,
16020 max_output_chars: Option<usize>,
16021 }
16022
16023 struct BlockingApprovalHandler {
16024 entered: Arc<tokio::sync::Barrier>,
16025 release: Arc<tokio::sync::Notify>,
16026 result: ApprovalResult,
16027 }
16028
16029 struct CountingApprovalHandler {
16030 calls: Arc<std::sync::atomic::AtomicUsize>,
16031 }
16032
16033 struct DriftingFallbackProvider {
16035 refreshed: AtomicBool,
16036 primary_calls: Arc<std::sync::atomic::AtomicUsize>,
16037 secondary_calls: Arc<std::sync::atomic::AtomicUsize>,
16038 }
16039
16040 struct RefreshFallbackProviderHooks {
16042 agent: parking_lot::Mutex<Option<Weak<RuntimeAgent>>>,
16043 lifecycle: Arc<ToolLifecycleRecordingHooks>,
16044 }
16045
16046 struct RuntimeWebFetchTransport {
16047 calls: Arc<std::sync::atomic::AtomicUsize>,
16048 }
16049
16050 struct RuntimeWebFetchResolver;
16051
16052 struct ReentrantToolHooks {
16053 agent: parking_lot::Mutex<Option<Weak<RuntimeAgent>>>,
16054 invoked: AtomicBool,
16055 nested_success: AtomicBool,
16056 }
16057
16058 #[async_trait]
16059 impl ai_agents_core::Tool for ClassifiedTimeoutTool {
16060 fn id(&self) -> &str {
16062 self.id
16063 }
16064
16065 fn name(&self) -> &str {
16067 "Classified Timeout"
16068 }
16069
16070 fn description(&self) -> &str {
16072 "Records and waits under one call-level timeout."
16073 }
16074
16075 fn input_schema(&self) -> Value {
16077 serde_json::json!({"type": "object"})
16078 }
16079
16080 fn classify_call(&self, _args: &Value) -> ai_agents_core::ToolCallClassification {
16082 let mut classification =
16083 ai_agents_core::ToolCallClassification::from_metadata(&self.safety_metadata());
16084 classification.timeout_ms = Some(self.timeout_ms);
16085 classification.requires_approval = self.requires_approval;
16086 classification
16087 }
16088
16089 async fn execute(
16091 &self,
16092 _args: Value,
16093 ctx: ai_agents_core::ToolExecutionContext,
16094 ) -> ToolResult {
16095 self.calls.fetch_add(1, Ordering::SeqCst);
16096 let deadline = ctx
16097 .deadline
16098 .expect("each invocation must receive a deadline");
16099 self.remaining_ms.lock().push(
16100 deadline
16101 .signed_duration_since(chrono::Utc::now())
16102 .num_milliseconds(),
16103 );
16104 tokio::time::sleep(Duration::from_millis(self.sleep_ms)).await;
16105 ToolResult::ok("done")
16106 }
16107 }
16108
16109 #[async_trait]
16110 impl ai_agents_core::Tool for ApprovalModifiedTimeoutTool {
16111 fn id(&self) -> &str {
16113 "approval_modified_timeout"
16114 }
16115
16116 fn name(&self) -> &str {
16118 "Approval Modified Timeout"
16119 }
16120
16121 fn description(&self) -> &str {
16123 "Becomes invalid only after approval modifies its arguments."
16124 }
16125
16126 fn input_schema(&self) -> Value {
16128 serde_json::json!({"type": "object"})
16129 }
16130
16131 fn policy_bindings(&self) -> ai_agents_core::ToolPolicyBindings {
16133 ai_agents_core::ToolPolicyBindings {
16134 path_fields: vec![ai_agents_core::PathPolicyBinding::write("path")],
16135 ..Default::default()
16136 }
16137 }
16138
16139 fn safety_metadata(&self) -> ai_agents_core::ToolSafetyMetadata {
16141 ai_agents_core::ToolSafetyMetadata {
16142 read_only: false,
16143 concurrency_safe: false,
16144 operation: ai_agents_core::ToolOperationKind::Write,
16145 side_effect_level: ai_agents_core::ToolSideEffectLevel::LocalWrite,
16146 requires_network: false,
16147 destructive: false,
16148 open_world: false,
16149 host_dependent: false,
16150 requires_user_interaction: false,
16151 supports_cancellation: true,
16152 default_requires_approval: true,
16153 should_defer_schema: false,
16154 max_output_chars: Some(1024),
16155 max_result_size_chars: Some(1024),
16156 }
16157 }
16158
16159 fn classify_call(&self, args: &Value) -> ai_agents_core::ToolCallClassification {
16161 let mut classification =
16162 ai_agents_core::ToolCallClassification::from_metadata(&self.safety_metadata());
16163 classification.timeout_ms = Some(if args["invalid_timeout"].as_bool() == Some(true) {
16164 u64::MAX
16165 } else {
16166 1_000
16167 });
16168 classification
16169 }
16170
16171 async fn execute(
16173 &self,
16174 _args: Value,
16175 _ctx: ai_agents_core::ToolExecutionContext,
16176 ) -> ToolResult {
16177 self.calls.fetch_add(1, Ordering::SeqCst);
16178 ToolResult::ok("unexpected")
16179 }
16180 }
16181
16182 #[async_trait]
16183 impl ai_agents_core::Tool for SlowTool {
16184 fn id(&self) -> &str {
16185 "slow"
16186 }
16187
16188 fn name(&self) -> &str {
16189 "Slow"
16190 }
16191
16192 fn description(&self) -> &str {
16193 "Waits until cancelled or timed out."
16194 }
16195
16196 fn input_schema(&self) -> Value {
16197 serde_json::json!({"type": "object"})
16198 }
16199
16200 async fn execute(
16201 &self,
16202 _args: Value,
16203 _ctx: ai_agents_core::ToolExecutionContext,
16204 ) -> ToolResult {
16205 tokio::time::sleep(std::time::Duration::from_secs(5)).await;
16206 ToolResult::ok("done")
16207 }
16208 }
16209
16210 #[async_trait]
16211 impl ai_agents_core::Tool for FlakyWriteTool {
16212 fn id(&self) -> &str {
16213 "flaky_write"
16214 }
16215
16216 fn name(&self) -> &str {
16217 "Flaky Write"
16218 }
16219
16220 fn description(&self) -> &str {
16221 "Fails on the first write attempt."
16222 }
16223
16224 fn input_schema(&self) -> Value {
16225 serde_json::json!({"type": "object", "properties": {"path": {"type": "string"}}})
16226 }
16227
16228 fn policy_bindings(&self) -> ai_agents_core::ToolPolicyBindings {
16229 ai_agents_core::ToolPolicyBindings {
16230 path_fields: vec![ai_agents_core::PathPolicyBinding::write("path")],
16231 ..Default::default()
16232 }
16233 }
16234
16235 fn safety_metadata(&self) -> ai_agents_core::ToolSafetyMetadata {
16236 ai_agents_core::ToolSafetyMetadata {
16237 read_only: false,
16238 concurrency_safe: false,
16239 operation: ai_agents_core::ToolOperationKind::Write,
16240 side_effect_level: ai_agents_core::ToolSideEffectLevel::LocalWrite,
16241 requires_network: false,
16242 destructive: false,
16243 open_world: false,
16244 host_dependent: false,
16245 requires_user_interaction: false,
16246 supports_cancellation: true,
16247 default_requires_approval: false,
16248 should_defer_schema: false,
16249 max_output_chars: Some(1024),
16250 max_result_size_chars: Some(1024),
16251 }
16252 }
16253
16254 fn classify_call(&self, _args: &Value) -> ai_agents_core::ToolCallClassification {
16255 let mut classification =
16256 ai_agents_core::ToolCallClassification::from_metadata(&self.safety_metadata());
16257 classification.safely_retryable = false;
16258 classification
16259 }
16260
16261 async fn execute(
16262 &self,
16263 _args: Value,
16264 _ctx: ai_agents_core::ToolExecutionContext,
16265 ) -> ToolResult {
16266 let call = self.calls.fetch_add(1, Ordering::SeqCst);
16267 if call == 0 {
16268 ToolResult::error("first failure")
16269 } else {
16270 ToolResult::ok("second success")
16271 }
16272 }
16273 }
16274
16275 #[async_trait]
16276 impl ai_agents_core::Tool for LockedWriteTool {
16277 fn id(&self) -> &str {
16278 "locked_write"
16279 }
16280
16281 fn name(&self) -> &str {
16282 "Locked Write"
16283 }
16284
16285 fn description(&self) -> &str {
16286 "Tracks concurrent execution on one resource."
16287 }
16288
16289 fn input_schema(&self) -> Value {
16290 serde_json::json!({"type": "object", "properties": {"path": {"type": "string"}}})
16291 }
16292
16293 fn policy_bindings(&self) -> ai_agents_core::ToolPolicyBindings {
16294 ai_agents_core::ToolPolicyBindings {
16295 path_fields: vec![ai_agents_core::PathPolicyBinding::write("path")],
16296 ..Default::default()
16297 }
16298 }
16299
16300 fn safety_metadata(&self) -> ai_agents_core::ToolSafetyMetadata {
16301 ai_agents_core::ToolSafetyMetadata {
16302 read_only: false,
16303 concurrency_safe: false,
16304 operation: ai_agents_core::ToolOperationKind::Write,
16305 side_effect_level: ai_agents_core::ToolSideEffectLevel::LocalWrite,
16306 requires_network: false,
16307 destructive: false,
16308 open_world: false,
16309 host_dependent: false,
16310 requires_user_interaction: false,
16311 supports_cancellation: true,
16312 default_requires_approval: false,
16313 should_defer_schema: false,
16314 max_output_chars: Some(1024),
16315 max_result_size_chars: Some(1024),
16316 }
16317 }
16318
16319 async fn execute(
16320 &self,
16321 _args: Value,
16322 _ctx: ai_agents_core::ToolExecutionContext,
16323 ) -> ToolResult {
16324 let active = self.active.fetch_add(1, Ordering::SeqCst) + 1;
16325 loop {
16326 let current_max = self.max_active.load(Ordering::SeqCst);
16327 if active <= current_max {
16328 break;
16329 }
16330 if self
16331 .max_active
16332 .compare_exchange(current_max, active, Ordering::SeqCst, Ordering::SeqCst)
16333 .is_ok()
16334 {
16335 break;
16336 }
16337 }
16338 tokio::time::sleep(std::time::Duration::from_millis(100)).await;
16339 self.active.fetch_sub(1, Ordering::SeqCst);
16340 ToolResult::ok("done")
16341 }
16342 }
16343
16344 #[async_trait]
16345 impl ai_agents_core::Tool for MultiResourceWriteTool {
16346 fn id(&self) -> &str {
16347 "multi_resource_write"
16348 }
16349
16350 fn name(&self) -> &str {
16351 "Multi Resource Write"
16352 }
16353
16354 fn description(&self) -> &str {
16355 "Tracks concurrent execution across source and destination resources."
16356 }
16357
16358 fn input_schema(&self) -> Value {
16359 serde_json::json!({"type": "object"})
16360 }
16361
16362 fn policy_bindings(&self) -> ai_agents_core::ToolPolicyBindings {
16363 ai_agents_core::ToolPolicyBindings {
16364 path_fields: vec![
16365 ai_agents_core::PathPolicyBinding::read_write("source_path"),
16366 ai_agents_core::PathPolicyBinding::write("destination_path"),
16367 ],
16368 ..Default::default()
16369 }
16370 }
16371
16372 fn safety_metadata(&self) -> ai_agents_core::ToolSafetyMetadata {
16373 LockedWriteTool {
16374 active: Arc::clone(&self.active),
16375 max_active: Arc::clone(&self.max_active),
16376 }
16377 .safety_metadata()
16378 }
16379
16380 async fn execute(
16381 &self,
16382 _args: Value,
16383 _ctx: ai_agents_core::ToolExecutionContext,
16384 ) -> ToolResult {
16385 let active = self.active.fetch_add(1, Ordering::SeqCst) + 1;
16386 self.max_active.fetch_max(active, Ordering::SeqCst);
16387 tokio::time::sleep(std::time::Duration::from_millis(75)).await;
16388 self.active.fetch_sub(1, Ordering::SeqCst);
16389 ToolResult::ok("done")
16390 }
16391 }
16392
16393 #[async_trait]
16394 impl ai_agents_core::Tool for BlockingPathMutationTool {
16395 fn id(&self) -> &str {
16396 self.id
16397 }
16398
16399 fn name(&self) -> &str {
16400 self.id
16401 }
16402
16403 fn description(&self) -> &str {
16404 "Blocks a path mutation until the test releases it."
16405 }
16406
16407 fn input_schema(&self) -> Value {
16408 serde_json::json!({"type": "object"})
16409 }
16410
16411 fn policy_bindings(&self) -> ai_agents_core::ToolPolicyBindings {
16412 ai_agents_core::ToolPolicyBindings {
16413 path_fields: self.path_fields.clone(),
16414 ..Default::default()
16415 }
16416 }
16417
16418 fn safety_metadata(&self) -> ai_agents_core::ToolSafetyMetadata {
16419 ai_agents_core::ToolSafetyMetadata {
16420 read_only: false,
16421 concurrency_safe: false,
16422 operation: ai_agents_core::ToolOperationKind::Write,
16423 side_effect_level: ai_agents_core::ToolSideEffectLevel::LocalWrite,
16424 requires_network: false,
16425 destructive: false,
16426 open_world: false,
16427 host_dependent: false,
16428 requires_user_interaction: false,
16429 supports_cancellation: true,
16430 default_requires_approval: false,
16431 should_defer_schema: false,
16432 max_output_chars: Some(1024),
16433 max_result_size_chars: Some(1024),
16434 }
16435 }
16436
16437 async fn execute(
16438 &self,
16439 _args: Value,
16440 _ctx: ai_agents_core::ToolExecutionContext,
16441 ) -> ToolResult {
16442 self.gate.entered.store(true, Ordering::SeqCst);
16443 self.gate.entered_notify.notify_one();
16444 self.gate.release.notified().await;
16445 ToolResult::ok("done")
16446 }
16447 }
16448
16449 #[async_trait]
16450 impl ai_agents_core::Tool for NoBindingWriteTool {
16451 fn id(&self) -> &str {
16452 "no_binding_write"
16453 }
16454
16455 fn name(&self) -> &str {
16456 "No Binding Write"
16457 }
16458
16459 fn description(&self) -> &str {
16460 "Tracks concurrent execution without resource bindings."
16461 }
16462
16463 fn input_schema(&self) -> Value {
16464 serde_json::json!({"type": "object"})
16465 }
16466
16467 fn safety_metadata(&self) -> ai_agents_core::ToolSafetyMetadata {
16468 LockedWriteTool {
16469 active: Arc::clone(&self.active),
16470 max_active: Arc::clone(&self.max_active),
16471 }
16472 .safety_metadata()
16473 }
16474
16475 async fn execute(
16476 &self,
16477 _args: Value,
16478 _ctx: ai_agents_core::ToolExecutionContext,
16479 ) -> ToolResult {
16480 let active = self.active.fetch_add(1, Ordering::SeqCst) + 1;
16481 self.max_active.fetch_max(active, Ordering::SeqCst);
16482 tokio::time::sleep(std::time::Duration::from_millis(75)).await;
16483 self.active.fetch_sub(1, Ordering::SeqCst);
16484 ToolResult::ok("done")
16485 }
16486 }
16487
16488 #[async_trait]
16489 impl ai_agents_core::Tool for RecoveryTestTool {
16490 fn id(&self) -> &str {
16491 &self.id
16492 }
16493
16494 fn name(&self) -> &str {
16495 &self.id
16496 }
16497
16498 fn description(&self) -> &str {
16499 "Records recovery execution and returns a configured result."
16500 }
16501
16502 fn input_schema(&self) -> Value {
16503 serde_json::json!({"type": "object"})
16504 }
16505
16506 fn policy_bindings(&self) -> ai_agents_core::ToolPolicyBindings {
16507 ai_agents_core::ToolPolicyBindings {
16508 path_fields: vec![ai_agents_core::PathPolicyBinding::write("path")],
16509 ..Default::default()
16510 }
16511 }
16512
16513 fn safety_metadata(&self) -> ai_agents_core::ToolSafetyMetadata {
16515 ai_agents_core::ToolSafetyMetadata {
16516 read_only: false,
16517 concurrency_safe: false,
16518 operation: ai_agents_core::ToolOperationKind::Write,
16519 side_effect_level: ai_agents_core::ToolSideEffectLevel::LocalWrite,
16520 requires_network: false,
16521 destructive: false,
16522 open_world: false,
16523 host_dependent: false,
16524 requires_user_interaction: false,
16525 supports_cancellation: true,
16526 default_requires_approval: false,
16527 should_defer_schema: false,
16528 max_output_chars: Some(self.max_output_chars.unwrap_or(1024)),
16529 max_result_size_chars: Some(1024),
16530 }
16531 }
16532
16533 async fn execute(
16535 &self,
16536 _args: Value,
16537 _ctx: ai_agents_core::ToolExecutionContext,
16538 ) -> ToolResult {
16539 self.calls.fetch_add(1, Ordering::SeqCst);
16540 let mut result = if self.succeeds {
16541 ToolResult::ok(format!("{} succeeded", self.id))
16542 } else {
16543 ToolResult::error(format!("{} failed", self.id))
16544 };
16545 result.metadata = Some(HashMap::from([(
16546 "recovery_test_tool".to_string(),
16547 Value::String(self.id.clone()),
16548 )]));
16549 result
16550 }
16551 }
16552
16553 #[async_trait]
16554 impl WebFetchTransport for RuntimeWebFetchTransport {
16555 async fn send(
16557 &self,
16558 _request: WebFetchTransportRequest,
16559 ) -> std::result::Result<WebFetchTransportResponse, String> {
16560 Err("validated addresses are required".to_string())
16561 }
16562
16563 async fn send_validated(
16565 &self,
16566 _request: WebFetchTransportRequest,
16567 _addresses: &[std::net::SocketAddr],
16568 ) -> std::result::Result<WebFetchTransportResponse, String> {
16569 self.calls.fetch_add(1, Ordering::SeqCst);
16570 Ok(WebFetchTransportResponse {
16571 status: 200,
16572 content_type: Some("text/plain".to_string()),
16573 location: None,
16574 body: b"approved".to_vec(),
16575 })
16576 }
16577 }
16578
16579 #[async_trait]
16580 impl WebFetchResolver for RuntimeWebFetchResolver {
16581 async fn resolve(
16583 &self,
16584 _host: &str,
16585 _port: u16,
16586 ) -> std::result::Result<Vec<std::net::IpAddr>, String> {
16587 Ok(vec![std::net::IpAddr::V4(std::net::Ipv4Addr::new(
16588 93, 184, 216, 34,
16589 ))])
16590 }
16591 }
16592
16593 #[async_trait]
16594 impl ToolProvider for DriftingFallbackProvider {
16595 fn id(&self) -> &str {
16597 "drifting_fallback"
16598 }
16599
16600 fn name(&self) -> &str {
16602 "Drifting Fallback"
16603 }
16604
16605 fn provider_type(&self) -> ToolProviderType {
16607 ToolProviderType::Custom
16608 }
16609
16610 async fn list_tools(&self) -> Vec<ToolDescriptor> {
16612 let alias = ToolAliases::new().with_name("en", "fallback alias");
16613 let mut primary = ToolDescriptor::new(
16614 "primary",
16615 "Primary",
16616 "Fails before fallback.",
16617 serde_json::json!({"type": "object"}),
16618 );
16619 let mut secondary = ToolDescriptor::new(
16620 "secondary",
16621 "Secondary",
16622 "Must not execute after final canonical drift.",
16623 serde_json::json!({"type": "object"}),
16624 );
16625 if self.refreshed.load(Ordering::SeqCst) {
16626 primary = primary.with_aliases(alias);
16627 } else {
16628 secondary = secondary.with_aliases(alias);
16629 }
16630 vec![primary, secondary]
16631 }
16632
16633 async fn get_tool(&self, tool_id: &str) -> Option<Arc<dyn Tool>> {
16635 let calls = match tool_id {
16636 "primary" => Arc::clone(&self.primary_calls),
16637 "secondary" => Arc::clone(&self.secondary_calls),
16638 _ => return None,
16639 };
16640 Some(Arc::new(RecoveryTestTool {
16641 id: tool_id.to_string(),
16642 succeeds: false,
16643 calls,
16644 max_output_chars: None,
16645 }))
16646 }
16647
16648 fn supports_refresh(&self) -> bool {
16650 true
16651 }
16652
16653 async fn refresh(&self) -> std::result::Result<(), ToolProviderError> {
16655 self.refreshed.store(true, Ordering::SeqCst);
16656 Ok(())
16657 }
16658 }
16659
16660 #[async_trait]
16661 impl AgentHooks for RefreshFallbackProviderHooks {
16662 async fn on_tool_start(&self, tool: &str, args: &Value) {
16664 self.lifecycle.on_tool_start(tool, args).await;
16665 if tool != "secondary" {
16666 return;
16667 }
16668 let agent = self.agent.lock().as_ref().and_then(Weak::upgrade);
16669 if let Some(agent) = agent {
16670 agent
16671 .tools
16672 .refresh_provider("drifting_fallback")
16673 .await
16674 .unwrap();
16675 }
16676 }
16677
16678 async fn on_tool_complete(&self, tool: &str, result: &ToolResult, duration_ms: u64) {
16679 self.lifecycle
16680 .on_tool_complete(tool, result, duration_ms)
16681 .await;
16682 }
16683
16684 async fn on_tool_execution_record(&self, record: &ToolExecutionRecord) {
16685 self.lifecycle.on_tool_execution_record(record).await;
16686 }
16687
16688 async fn on_error(&self, error: &AgentError) {
16689 self.lifecycle.on_error(error).await;
16690 }
16691 }
16692
16693 #[async_trait]
16694 impl ApprovalHandler for BlockingApprovalHandler {
16695 async fn request_approval(
16696 &self,
16697 _request: ai_agents_hitl::ApprovalRequest,
16698 ) -> ApprovalResult {
16699 self.entered.wait().await;
16700 self.release.notified().await;
16701 self.result.clone()
16702 }
16703 }
16704
16705 #[async_trait]
16706 impl ApprovalHandler for CountingApprovalHandler {
16707 async fn request_approval(
16708 &self,
16709 _request: ai_agents_hitl::ApprovalRequest,
16710 ) -> ApprovalResult {
16711 self.calls.fetch_add(1, Ordering::SeqCst);
16712 ApprovalResult::Approved
16713 }
16714 }
16715
16716 #[async_trait]
16717 impl AgentHooks for ReentrantToolHooks {
16718 async fn on_tool_complete(&self, tool: &str, _result: &ToolResult, _duration_ms: u64) {
16719 if tool != "reentrant_write" || self.invoked.swap(true, Ordering::SeqCst) {
16720 return;
16721 }
16722 let agent = self.agent.lock().as_ref().and_then(Weak::upgrade);
16723 if let Some(agent) = agent {
16724 let result = agent
16725 .invoke_tool(ToolExecutionRequest::new(
16726 "nested-hook-call",
16727 "reentrant_write",
16728 serde_json::json!({"path": "./hook.txt"}),
16729 ToolCallSource::Manual,
16730 ))
16731 .await;
16732 self.nested_success
16733 .store(result.is_ok_and(|record| record.success), Ordering::SeqCst);
16734 }
16735 }
16736 }
16737
16738 #[async_trait]
16739 impl AgentHooks for ResponseCountingHooks {
16740 async fn on_response(&self, _response: &AgentResponse) {
16741 self.responses.fetch_add(1, Ordering::SeqCst);
16742 }
16743 }
16744
16745 #[async_trait]
16746 impl AgentHooks for ResponseChatHooks {
16747 async fn on_response(&self, _response: &AgentResponse) {
16749 if self.invoked.swap(true, Ordering::SeqCst) {
16750 return;
16751 }
16752 let target = self.target.lock().as_ref().and_then(Weak::upgrade);
16753 let result = if let Some(target) = target {
16754 target
16755 .chat("nested response hook call")
16756 .await
16757 .map(|response| response.content)
16758 .map_err(|error| error.to_string())
16759 } else {
16760 Err("response hook target is unavailable".to_string())
16761 };
16762 *self.nested_result.lock() = Some(result);
16763 }
16764 }
16765
16766 #[async_trait]
16767 impl AgentHooks for ConcurrentResponseHooks {
16768 async fn on_response(&self, _response: &AgentResponse) {
16770 if self.invoked.swap(true, Ordering::SeqCst) {
16771 return;
16772 }
16773 let Some(registry) = self.registry.upgrade() else {
16774 *self.nested_result.lock() =
16775 Some(Err("concurrent registry is unavailable".to_string()));
16776 return;
16777 };
16778 let agents = [ai_agents_state::ConcurrentAgentRef::Id(
16779 self.child_id.clone(),
16780 )];
16781 let aggregation = ai_agents_state::AggregationConfig {
16782 strategy: ai_agents_state::AggregationStrategy::FirstWins,
16783 synthesizer_llm: None,
16784 synthesizer_prompt: None,
16785 vote: None,
16786 };
16787 let result = crate::orchestration::concurrent(
16788 ®istry,
16789 "nested concurrent response hook call",
16790 &agents,
16791 &aggregation,
16792 None,
16793 Some(1),
16794 None,
16795 ai_agents_state::PartialFailureAction::Abort,
16796 None,
16797 )
16798 .await
16799 .map(|result| result.response.content)
16800 .map_err(|error| error.to_string());
16801 *self.nested_result.lock() = Some(result);
16802 }
16803 }
16804
16805 #[async_trait]
16806 impl AgentHooks for ToolLifecycleRecordingHooks {
16807 async fn on_tool_start(&self, tool: &str, _args: &Value) {
16808 self.events.lock().push(format!("start:{tool}"));
16809 }
16810
16811 async fn on_tool_complete(&self, tool: &str, result: &ToolResult, _duration_ms: u64) {
16812 self.events
16813 .lock()
16814 .push(format!("complete:{tool}:{}", result.success));
16815 }
16816
16817 async fn on_tool_execution_record(&self, record: &ToolExecutionRecord) {
16818 self.events.lock().push(format!(
16819 "record:{}:{}",
16820 record.canonical_id, record.executed
16821 ));
16822 self.records.lock().push(record.clone());
16823 }
16824
16825 async fn on_error(&self, _error: &AgentError) {
16827 self.events.lock().push("error".to_string());
16828 }
16829 }
16830
16831 struct ApprovalRecordingHooks {
16832 events: parking_lot::Mutex<Vec<String>>,
16833 }
16834
16835 impl ApprovalRecordingHooks {
16836 fn new() -> Self {
16837 Self {
16838 events: parking_lot::Mutex::new(Vec::new()),
16839 }
16840 }
16841
16842 fn events(&self) -> Vec<String> {
16843 self.events.lock().clone()
16844 }
16845 }
16846
16847 #[async_trait]
16848 impl AgentHooks for ApprovalRecordingHooks {
16849 async fn on_approval_result(&self, request_id: &str, result: &ApprovalResult) {
16850 self.events.lock().push(format!(
16851 "raw:{}:{}",
16852 request_id,
16853 approval_result_name(result)
16854 ));
16855 }
16856
16857 async fn on_approval_resolved(
16858 &self,
16859 request: &ai_agents_hitl::ApprovalRequest,
16860 raw_result: &ApprovalResult,
16861 outcome: &ApprovalResolvedOutcome,
16862 ) {
16863 self.events.lock().push(format!(
16864 "resolved:{}:{}:{}",
16865 request.id,
16866 approval_result_name(raw_result),
16867 approval_outcome_name(outcome)
16868 ));
16869 }
16870 }
16871
16872 fn approval_result_name(result: &ApprovalResult) -> &'static str {
16873 match result {
16874 ApprovalResult::Approved => "approved",
16875 ApprovalResult::Rejected { .. } => "rejected",
16876 ApprovalResult::Modified { .. } => "modified",
16877 ApprovalResult::Timeout => "timeout",
16878 }
16879 }
16880
16881 fn approval_outcome_name(outcome: &ApprovalResolvedOutcome) -> &'static str {
16882 match outcome {
16883 ApprovalResolvedOutcome::Approved => "approved",
16884 ApprovalResolvedOutcome::Rejected { .. } => "rejected",
16885 ApprovalResolvedOutcome::Modified { .. } => "modified",
16886 ApprovalResolvedOutcome::Error { .. } => "error",
16887 }
16888 }
16889
16890 fn assert_correlated_approval_events(
16891 events: &[String],
16892 raw_status: &str,
16893 outcome_status: &str,
16894 ) {
16895 assert_eq!(events.len(), 2);
16896 let raw: Vec<_> = events[0].split(':').collect();
16897 let resolved: Vec<_> = events[1].split(':').collect();
16898 assert_eq!(raw[0], "raw");
16899 assert_eq!(resolved[0], "resolved");
16900 assert_eq!(raw[1], resolved[1]);
16901 assert_eq!(raw[2], raw_status);
16902 assert_eq!(resolved[2], raw_status);
16903 assert_eq!(resolved[3], outcome_status);
16904 }
16905
16906 fn approval_security_config(policy_enabled: bool) -> ToolSecurityConfig {
16907 let mut security = ToolSecurityConfig {
16908 enabled: true,
16909 fail_closed: true,
16910 ..Default::default()
16911 };
16912 let policy = ai_agents_tools::ToolPolicyConfig {
16913 enabled: policy_enabled,
16914 write_paths: vec![".".to_string()],
16915 require_confirmation: true,
16916 ..Default::default()
16917 };
16918 security.tools.insert("locked_write".to_string(), policy);
16919 security
16920 }
16921
16922 struct MutationTestWorkspace {
16923 root: std::path::PathBuf,
16924 }
16925
16926 impl MutationTestWorkspace {
16927 fn new() -> Self {
16928 let root = std::env::temp_dir().join(format!(
16929 "ai-agents-runtime-mutation-{}",
16930 uuid::Uuid::new_v4()
16931 ));
16932 std::fs::create_dir_all(&root).unwrap();
16933 Self { root }
16934 }
16935 }
16936
16937 impl Drop for MutationTestWorkspace {
16938 fn drop(&mut self) {
16939 let _ = std::fs::remove_dir_all(&self.root);
16940 }
16941 }
16942
16943 async fn wait_for_resource_lock_strong_count(locks: &ToolResourceLocks, minimum: usize) {
16944 tokio::time::timeout(std::time::Duration::from_secs(2), async {
16945 loop {
16946 let strong_count = locks
16947 .read()
16948 .get("path-mutation:global")
16949 .map_or(0, |lock| lock.strong_count());
16950 if strong_count >= minimum {
16951 break;
16952 }
16953 tokio::task::yield_now().await;
16954 }
16955 })
16956 .await
16957 .expect("path mutation call did not reach the shared lock");
16958 }
16959
16960 async fn assert_path_mutation_pair_serialized(
16961 first_id: &'static str,
16962 first_fields: Vec<ai_agents_core::PathPolicyBinding>,
16963 first_args: Value,
16964 second_id: &'static str,
16965 second_fields: Vec<ai_agents_core::PathPolicyBinding>,
16966 second_args: Value,
16967 ) {
16968 let locks = new_tool_resource_locks();
16969 let first_gate = PathMutationGate::new();
16970 let second_gate = PathMutationGate::new();
16971 second_gate.release();
16972 let agent = Arc::new(
16973 AgentBuilder::new()
16974 .system_prompt("Test global path mutation locking.")
16975 .llm(Arc::new(mock_with_response("done")))
16976 .tool(Arc::new(BlockingPathMutationTool {
16977 id: first_id,
16978 path_fields: first_fields,
16979 gate: first_gate.clone(),
16980 }))
16981 .tool(Arc::new(BlockingPathMutationTool {
16982 id: second_id,
16983 path_fields: second_fields,
16984 gate: second_gate.clone(),
16985 }))
16986 .build()
16987 .unwrap()
16988 .with_shared_resource_locks(Arc::clone(&locks)),
16989 );
16990
16991 let first = {
16992 let agent = Arc::clone(&agent);
16993 tokio::spawn(async move {
16994 agent
16995 .invoke_tool(ToolExecutionRequest::new(
16996 format!("{}-first", first_id),
16997 first_id,
16998 first_args,
16999 ToolCallSource::Manual,
17000 ))
17001 .await
17002 .unwrap()
17003 })
17004 };
17005 first_gate.wait_until_entered().await;
17006
17007 let second = {
17008 let agent = Arc::clone(&agent);
17009 tokio::spawn(async move {
17010 agent
17011 .invoke_tool(ToolExecutionRequest::new(
17012 format!("{}-second", second_id),
17013 second_id,
17014 second_args,
17015 ToolCallSource::Manual,
17016 ))
17017 .await
17018 .unwrap()
17019 })
17020 };
17021 wait_for_resource_lock_strong_count(&locks, 2).await;
17022 assert!(!second_gate.entered.load(Ordering::SeqCst));
17023 assert!(!second.is_finished());
17024
17025 first_gate.release();
17026 let (first, second) = tokio::time::timeout(std::time::Duration::from_secs(2), async {
17027 tokio::join!(first, second)
17028 })
17029 .await
17030 .expect("serialized path mutation calls did not finish");
17031 assert!(first.unwrap().success);
17032 assert!(second.unwrap().success);
17033 assert!(second_gate.entered.load(Ordering::SeqCst));
17034 assert!(locks.read().is_empty());
17035 }
17036
17037 #[derive(Clone, Copy)]
17038 enum MutationDenial {
17039 Policy,
17040 Approval,
17041 }
17042
17043 fn mutation_denial_security_config(
17044 tool_id: &str,
17045 workspace: &std::path::Path,
17046 denial: MutationDenial,
17047 ) -> ToolSecurityConfig {
17048 let workspace = workspace.to_string_lossy().into_owned();
17049 let mut policy = ai_agents_tools::ToolPolicyConfig {
17050 read_paths: vec![workspace.clone()],
17051 write_paths: vec![workspace.clone()],
17052 ..Default::default()
17053 };
17054 match denial {
17055 MutationDenial::Policy => policy.blocked_paths = vec![workspace],
17056 MutationDenial::Approval => policy.require_confirmation = true,
17057 }
17058
17059 let mut security = ToolSecurityConfig {
17060 enabled: true,
17061 fail_closed: true,
17062 ..Default::default()
17063 };
17064 security.tools.insert(tool_id.to_string(), policy);
17065 security
17066 }
17067
17068 async fn assert_path_mutation_denied(tool: Arc<dyn Tool>, denial: MutationDenial) {
17069 let workspace = MutationTestWorkspace::new();
17070 let tool_id = tool.id().to_string();
17071 let preserved = workspace.root.join(format!("{}-preserved.txt", tool_id));
17072 let destination = workspace.root.join(format!("{}-destination.txt", tool_id));
17073 std::fs::write(&preserved, "preserved").unwrap();
17074 let arguments = match tool_id.as_str() {
17075 "copy_path" | "move_path" => serde_json::json!({
17076 "source_path": preserved.to_string_lossy(),
17077 "destination_path": destination.to_string_lossy(),
17078 "dry_run": false
17079 }),
17080 "delete_path" => serde_json::json!({
17081 "path": preserved.to_string_lossy(),
17082 "recursive": false,
17083 "dry_run": false
17084 }),
17085 _ => panic!("unsupported mutation tool: {}", tool_id),
17086 };
17087 let security = mutation_denial_security_config(&tool_id, &workspace.root, denial);
17088 let builder = AgentBuilder::new()
17089 .system_prompt("Test mutation denial.")
17090 .llm(Arc::new(mock_with_response("done")))
17091 .tool(tool)
17092 .tool_security(ToolSecurityEngine::new(security));
17093 let builder = match denial {
17094 MutationDenial::Policy => builder,
17095 MutationDenial::Approval => builder
17096 .hitl_engine(HITLEngine::new(ai_agents_hitl::HITLConfig::default()))
17097 .approval_handler(Arc::new(RejectAllHandler::new())),
17098 };
17099 let agent = builder.build().unwrap();
17100
17101 let record = agent
17102 .invoke_tool(ToolExecutionRequest::new(
17103 format!("{}-denied", tool_id),
17104 tool_id.clone(),
17105 arguments,
17106 ToolCallSource::Manual,
17107 ))
17108 .await
17109 .unwrap();
17110
17111 assert!(!record.executed, "{} must not be invoked", tool_id);
17112 assert!(!record.success);
17113 match denial {
17114 MutationDenial::Policy => {
17115 assert_eq!(record.policy.outcome, PermissionOutcome::Deny);
17116 assert!(record.approval.as_ref().is_some_and(|approval| matches!(
17117 &approval.status,
17118 ToolApprovalStatus::NotRequired
17119 )));
17120 }
17121 MutationDenial::Approval => {
17122 assert_eq!(record.policy.outcome, PermissionOutcome::RequiresApproval);
17123 assert!(record.approval.as_ref().is_some_and(|approval| matches!(
17124 &approval.status,
17125 ToolApprovalStatus::Rejected
17126 )));
17127 }
17128 }
17129 assert_eq!(std::fs::read_to_string(&preserved).unwrap(), "preserved");
17130 assert!(!destination.exists());
17131 }
17132
17133 fn recovery_manager_with_fallbacks(
17134 fallbacks: impl IntoIterator<Item = (String, String)>,
17135 ) -> RecoveryManager {
17136 use ai_agents_recovery::{ErrorRecoveryConfig, ToolRecoveryConfig, ToolRetryConfig};
17137
17138 let per_tool = fallbacks
17139 .into_iter()
17140 .map(|(tool, fallback_tool)| {
17141 (
17142 tool,
17143 ToolRetryConfig {
17144 max_retries: 0,
17145 timeout_ms: Some(1_000),
17146 on_failure: ToolFailureAction::Fallback { fallback_tool },
17147 },
17148 )
17149 })
17150 .collect();
17151 RecoveryManager::new(ErrorRecoveryConfig {
17152 tools: ToolRecoveryConfig {
17153 per_tool,
17154 ..Default::default()
17155 },
17156 ..Default::default()
17157 })
17158 }
17159
17160 fn approval_check() -> HITLCheckResult {
17161 HITLCheckResult::required(
17162 ApprovalTrigger::tool("test", serde_json::json!({})),
17163 HashMap::new(),
17164 "Approve?",
17165 None,
17166 )
17167 }
17168
17169 fn agent_with_approval_result(
17170 raw_result: ApprovalResult,
17171 timeout_action: TimeoutAction,
17172 hooks: Arc<ApprovalRecordingHooks>,
17173 ) -> RuntimeAgent {
17174 use ai_agents_hitl::{CallbackHandler, HITLConfig};
17175
17176 let config = HITLConfig {
17177 on_timeout: timeout_action,
17178 ..Default::default()
17179 };
17180 let handler = CallbackHandler::new(move |_| raw_result.clone());
17181 AgentBuilder::new()
17182 .system_prompt("Test HITL hooks.")
17183 .llm(Arc::new(mock_with_response("done")))
17184 .build()
17185 .unwrap()
17186 .with_hooks(hooks)
17187 .with_hitl(HITLEngine::new(config), Arc::new(handler))
17188 }
17189
17190 #[tokio::test]
17191 async fn approval_hooks_expose_direct_effective_decisions_after_raw_results() {
17192 let cases = vec![
17193 (ApprovalResult::Approved, "approved"),
17194 (
17195 ApprovalResult::Rejected {
17196 reason: Some("denied".to_string()),
17197 },
17198 "rejected",
17199 ),
17200 (
17201 ApprovalResult::Modified {
17202 changes: HashMap::from([("value".to_string(), serde_json::json!(2))]),
17203 },
17204 "modified",
17205 ),
17206 ];
17207
17208 for (raw_result, expected) in cases {
17209 let hooks = Arc::new(ApprovalRecordingHooks::new());
17210 let agent =
17211 agent_with_approval_result(raw_result, TimeoutAction::Reject, hooks.clone());
17212
17213 let result = agent.request_hitl_approval(approval_check()).await.unwrap();
17214
17215 assert_eq!(approval_result_name(&result), expected);
17216 assert_correlated_approval_events(&hooks.events(), expected, expected);
17217 }
17218 }
17219
17220 #[tokio::test]
17221 async fn approval_hooks_expose_timeout_policy_decisions() {
17222 for (timeout_action, expected) in [
17223 (TimeoutAction::Approve, "approved"),
17224 (TimeoutAction::Reject, "rejected"),
17225 ] {
17226 let hooks = Arc::new(ApprovalRecordingHooks::new());
17227 let agent =
17228 agent_with_approval_result(ApprovalResult::Timeout, timeout_action, hooks.clone());
17229
17230 let result = agent.request_hitl_approval(approval_check()).await.unwrap();
17231
17232 assert_eq!(approval_result_name(&result), expected);
17233 assert_correlated_approval_events(&hooks.events(), "timeout", expected);
17234 }
17235 }
17236
17237 #[tokio::test]
17238 async fn timeout_error_fires_correlated_resolved_error_before_returning() {
17239 let hooks = Arc::new(ApprovalRecordingHooks::new());
17240 let agent = agent_with_approval_result(
17241 ApprovalResult::Timeout,
17242 TimeoutAction::Error,
17243 hooks.clone(),
17244 );
17245
17246 let error = agent
17247 .request_hitl_approval(approval_check())
17248 .await
17249 .unwrap_err();
17250
17251 assert!(error.to_string().contains("HITL approval timeout"));
17252 assert_correlated_approval_events(&hooks.events(), "timeout", "error");
17253 }
17254
17255 #[tokio::test]
17257 async fn test_integration_yaml_to_chat_basic() {
17258 let mock = mock_with_response("Hello! How can I help you?");
17259 let agent = AgentBuilder::new()
17260 .system_prompt("You are a test assistant.")
17261 .llm(Arc::new(mock))
17262 .build()
17263 .unwrap();
17264
17265 let response = agent.chat("Hi").await.unwrap();
17266 assert!(!response.content.is_empty());
17267 assert_eq!(response.content, "Hello! How can I help you?");
17268 }
17269
17270 #[tokio::test]
17271 async fn stream_events_emit_one_authoritative_final_without_legacy_done() {
17272 let agent = AgentBuilder::new()
17273 .system_prompt("You are a test assistant.")
17274 .llm(Arc::new(mock_with_response(
17275 "Hello from the final response.",
17276 )))
17277 .build()
17278 .unwrap();
17279
17280 let mut stream = agent.chat_stream_events("Hi").await.unwrap();
17281 let mut final_responses = Vec::new();
17282 let mut legacy_done = 0;
17283 while let Some(event) = stream.next().await {
17284 match event {
17285 AgentStreamEvent::Chunk(StreamChunk::Done {}) => legacy_done += 1,
17286 AgentStreamEvent::Chunk(StreamChunk::Error { message }) => {
17287 panic!("unexpected stream error: {message}")
17288 }
17289 AgentStreamEvent::Final(response) => final_responses.push(response),
17290 AgentStreamEvent::Chunk(_) => {}
17291 }
17292 }
17293
17294 assert_eq!(legacy_done, 0);
17295 assert_eq!(final_responses.len(), 1);
17296 let response = final_responses.pop().unwrap();
17297 assert_eq!(response.content, "Hello from the final response.");
17298 assert!(
17299 response
17300 .metadata
17301 .as_ref()
17302 .is_some_and(|metadata| { metadata.contains_key("reasoning") })
17303 );
17304 }
17305
17306 #[tokio::test]
17307 async fn stream_final_content_includes_output_processing_after_provisional_chunks() {
17308 let yaml = r#"
17309name: ProcessedStreamAgent
17310system_prompt: "Answer directly."
17311process:
17312 output:
17313 - type: format
17314 config:
17315 template: "{{ response }} [finalized]"
17316streaming:
17317 enabled: true
17318"#;
17319 let agent = AgentBuilder::from_yaml(yaml)
17320 .unwrap()
17321 .llm(Arc::new(mock_with_response("provisional answer")))
17322 .auto_configure_features()
17323 .unwrap()
17324 .build()
17325 .unwrap();
17326
17327 let mut stream = agent.chat_stream_events("Hi").await.unwrap();
17328 let mut provisional = String::new();
17329 let mut final_content = None;
17330 while let Some(event) = stream.next().await {
17331 match event {
17332 AgentStreamEvent::Chunk(StreamChunk::Content { text }) => {
17333 provisional.push_str(&text)
17334 }
17335 AgentStreamEvent::Chunk(StreamChunk::Error { message }) => {
17336 panic!("unexpected stream error: {message}")
17337 }
17338 AgentStreamEvent::Final(response) => final_content = Some(response.content),
17339 AgentStreamEvent::Chunk(_) => {}
17340 }
17341 }
17342
17343 assert_eq!(provisional, "provisional answer");
17344 assert_eq!(
17345 final_content.as_deref(),
17346 Some("provisional answer [finalized]")
17347 );
17348 }
17349
17350 #[tokio::test]
17351 async fn stream_events_preserve_tool_progress_and_final_tool_calls() {
17352 let agent = AgentBuilder::new()
17353 .system_prompt("Use the echo tool once, then answer.")
17354 .llm(Arc::new(mock_with_responses(vec![
17355 r#"{"tool":"echo","arguments":{"message":"hello"}}"#,
17356 "Echo completed.",
17357 ])))
17358 .tool(Arc::new(ai_agents_tools::EchoTool::new()))
17359 .build()
17360 .unwrap();
17361
17362 let mut stream = agent.chat_stream_events("echo hello").await.unwrap();
17363 let mut starts = 0;
17364 let mut results = 0;
17365 let mut ends = 0;
17366 let mut final_response = None;
17367 while let Some(event) = stream.next().await {
17368 match event {
17369 AgentStreamEvent::Chunk(StreamChunk::ToolCallStart { name, .. }) => {
17370 assert_eq!(name, "echo");
17371 starts += 1;
17372 }
17373 AgentStreamEvent::Chunk(StreamChunk::ToolResult { name, success, .. }) => {
17374 assert_eq!(name, "echo");
17375 assert!(success);
17376 results += 1;
17377 }
17378 AgentStreamEvent::Chunk(StreamChunk::ToolCallEnd { .. }) => ends += 1,
17379 AgentStreamEvent::Chunk(StreamChunk::Error { message }) => {
17380 panic!("unexpected stream error: {message}")
17381 }
17382 AgentStreamEvent::Final(response) => final_response = Some(response),
17383 AgentStreamEvent::Chunk(_) => {}
17384 }
17385 }
17386
17387 assert_eq!((starts, results, ends), (1, 1, 1));
17388 let response = final_response.expect("tool stream must finalize");
17389 assert_eq!(response.content, "Echo completed.");
17390 assert_eq!(
17391 response.tool_calls.as_ref().map(|calls| calls
17392 .iter()
17393 .map(|call| call.name.as_str())
17394 .collect::<Vec<_>>()),
17395 Some(vec!["echo"])
17396 );
17397 }
17398
17399 #[tokio::test]
17400 async fn legacy_stream_still_emits_one_done_chunk() {
17401 let agent = AgentBuilder::new()
17402 .system_prompt("You are a test assistant.")
17403 .llm(Arc::new(mock_with_response(
17404 "Hello from the legacy stream.",
17405 )))
17406 .build()
17407 .unwrap();
17408
17409 let mut stream = agent.chat_stream("Hi").await.unwrap();
17410 let mut done = 0;
17411 while let Some(chunk) = stream.next().await {
17412 match chunk {
17413 StreamChunk::Done {} => done += 1,
17414 StreamChunk::Error { message } => panic!("unexpected stream error: {message}"),
17415 _ => {}
17416 }
17417 }
17418
17419 assert_eq!(done, 1);
17420 }
17421
17422 #[tokio::test]
17424 async fn test_integration_multi_turn_conversation() {
17425 let mock = mock_with_responses(vec![
17426 "Hello! I'm your assistant.",
17427 "The weather is sunny today.",
17428 "Goodbye!",
17429 ]);
17430 let agent = AgentBuilder::new()
17431 .system_prompt("You are helpful.")
17432 .llm(Arc::new(mock))
17433 .build()
17434 .unwrap();
17435
17436 let r1 = agent.chat("Hi").await.unwrap();
17437 assert_eq!(r1.content, "Hello! I'm your assistant.");
17438
17439 let r2 = agent.chat("What's the weather?").await.unwrap();
17440 assert_eq!(r2.content, "The weather is sunny today.");
17441
17442 let r3 = agent.chat("Bye").await.unwrap();
17443 assert_eq!(r3.content, "Goodbye!");
17444
17445 let messages = agent.memory.get_messages(None).await.unwrap();
17447 assert_eq!(messages.len(), 6);
17449 }
17450
17451 #[test]
17452 fn later_approval_preserves_modified_evidence() {
17453 let arguments = serde_json::json!({"dry_run": true});
17454 let mut record = Some(ToolApprovalRecord {
17455 status: ToolApprovalStatus::Modified,
17456 reason: None,
17457 modified_arguments: Some(arguments.clone()),
17458 });
17459
17460 merge_approved_record(&mut record);
17461
17462 let record = record.unwrap();
17463 assert!(matches!(record.status, ToolApprovalStatus::Modified));
17464 assert_eq!(record.modified_arguments, Some(arguments));
17465 }
17466
17467 #[test]
17468 fn approval_binding_rejects_replaced_tool_implementation() {
17469 let reviewed_tool: Arc<dyn ai_agents_core::Tool> = Arc::new(ContextEchoTool);
17470 let same_tool = Arc::clone(&reviewed_tool);
17471 let replacement_tool: Arc<dyn ai_agents_core::Tool> = Arc::new(ContextEchoTool);
17472 let arguments = serde_json::json!({"path": "."});
17473 let versions = ToolDecisionVersions {
17474 policy: 2,
17475 registry: 3,
17476 runtime_control: 4,
17477 state: Some(5),
17478 };
17479 let binding = ToolApprovalBinding {
17480 canonical_id: "context_echo".to_string(),
17481 arguments: arguments.clone(),
17482 confirmation_required: true,
17483 policy_version: versions.policy,
17484 runtime_control_version: versions.runtime_control,
17485 state_generation: versions.state,
17486 reviewed_tool,
17487 };
17488
17489 assert!(!binding.is_stale("context_echo", &arguments, true, versions, &same_tool,));
17490 assert!(binding.is_stale(
17491 "context_echo",
17492 &arguments,
17493 true,
17494 versions,
17495 &replacement_tool,
17496 ));
17497 }
17498
17499 #[tokio::test]
17500 async fn approved_mutation_to_dry_run_remains_executable() {
17501 use ai_agents_hitl::CallbackHandler;
17502
17503 let handler = CallbackHandler::new(|_| ApprovalResult::Modified {
17504 changes: HashMap::from([("dry_run".to_string(), serde_json::json!(true))]),
17505 });
17506 let agent = AgentBuilder::new()
17507 .system_prompt("Test safer approval modifications.")
17508 .llm(Arc::new(mock_with_response("done")))
17509 .tool(Arc::new(ai_agents_tools::FileWriteTool::new()))
17510 .hitl_engine(HITLEngine::new(ai_agents_hitl::HITLConfig::default()))
17511 .approval_handler(Arc::new(handler))
17512 .build()
17513 .unwrap();
17514
17515 let record = agent
17516 .invoke_tool(ToolExecutionRequest::new(
17517 "approved-dry-run",
17518 "file_write",
17519 serde_json::json!({
17520 "path": "./approval-dry-run.txt",
17521 "content": "not written"
17522 }),
17523 ToolCallSource::Manual,
17524 ))
17525 .await
17526 .unwrap();
17527
17528 assert!(record.executed);
17529 assert!(record.success);
17530 assert_eq!(record.executed_arguments["dry_run"], true);
17531 assert!(matches!(
17532 record.approval.as_ref().map(|approval| &approval.status),
17533 Some(ToolApprovalStatus::Modified)
17534 ));
17535 let output: Value = serde_json::from_str(&record.output).unwrap();
17536 assert_eq!(output["mutation_performed"], false);
17537 }
17538
17539 #[tokio::test]
17541 async fn shared_executor_approval_reaches_web_fetch_transport() {
17542 use ai_agents_hitl::{CallbackHandler, HITLConfig};
17543 use ai_agents_tools::{DomainPolicyConfig, ToolPolicyConfig};
17544
17545 let calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
17546 let tool = WebFetchTool::with_transport_and_resolver(
17547 Arc::new(RuntimeWebFetchTransport {
17548 calls: Arc::clone(&calls),
17549 }),
17550 Arc::new(RuntimeWebFetchResolver),
17551 );
17552 let mut security = ToolSecurityConfig {
17553 enabled: true,
17554 fail_closed: true,
17555 ..Default::default()
17556 };
17557 security.tools.insert(
17558 "web_fetch".to_string(),
17559 ToolPolicyConfig {
17560 domains: DomainPolicyConfig {
17561 requires_approval: vec!["approval.test".to_string()],
17562 ..Default::default()
17563 },
17564 allowed_schemes: vec!["https".to_string()],
17565 allowed_ports: vec![443],
17566 ..Default::default()
17567 },
17568 );
17569 let handler = CallbackHandler::new(|_| ApprovalResult::Approved);
17570 let agent = AgentBuilder::new()
17571 .system_prompt("Test approved web fetch execution.")
17572 .llm(Arc::new(mock_with_response("done")))
17573 .tool(Arc::new(tool))
17574 .tool_security(ToolSecurityEngine::new(security))
17575 .build()
17576 .unwrap()
17577 .with_hitl(HITLEngine::new(HITLConfig::default()), Arc::new(handler));
17578
17579 let record = agent
17580 .invoke_tool(ToolExecutionRequest::new(
17581 "approved-web-fetch",
17582 "web_fetch",
17583 serde_json::json!({
17584 "url": "https://approval.test/page",
17585 "cache_ttl_seconds": 0
17586 }),
17587 ToolCallSource::Manual,
17588 ))
17589 .await
17590 .unwrap();
17591
17592 assert!(record.success);
17593 assert!(
17594 record
17595 .approval
17596 .as_ref()
17597 .is_some_and(|approval| matches!(approval.status, ToolApprovalStatus::Approved))
17598 );
17599 assert_eq!(calls.load(Ordering::SeqCst), 1);
17600 }
17601
17602 #[tokio::test]
17603 async fn context_preserves_requested_and_canonical_identity() {
17604 let mock = mock_with_response("hello");
17605 let mut tools = ai_agents_tools::ToolRegistry::new();
17606 tools.register(Arc::new(ContextEchoTool)).unwrap();
17607
17608 let mut security = ToolSecurityConfig {
17609 enabled: true,
17610 fail_closed: true,
17611 ..Default::default()
17612 };
17613 let mut policy = ai_agents_tools::ToolPolicyConfig {
17614 read_paths: vec![".".to_string()],
17615 max_results: Some(7),
17616 ..Default::default()
17617 };
17618 policy
17619 .config
17620 .insert("backend".to_string(), serde_json::json!("memory"));
17621 security.tools.insert("context_echo".to_string(), policy);
17622
17623 let agent = AgentBuilder::new()
17624 .system_prompt("You are helpful.")
17625 .llm(Arc::new(mock))
17626 .tools(tools)
17627 .tool_security(ToolSecurityEngine::new(security))
17628 .build()
17629 .unwrap();
17630
17631 let record = agent
17632 .invoke_tool(ToolExecutionRequest::new(
17633 "ctx-call",
17634 "Context Echo",
17635 serde_json::json!({"path": ".", "max_results": 99}),
17636 ToolCallSource::Manual,
17637 ))
17638 .await
17639 .unwrap();
17640
17641 assert!(record.success);
17642 assert!(matches!(&record.source, ToolCallSource::Manual));
17643 assert_eq!(record.requested_name, "Context Echo");
17644 assert_eq!(record.canonical_id, "context_echo");
17645 assert_eq!(record.policy.outcome, PermissionOutcome::Allow);
17646 assert_eq!(record.executed_arguments["max_results"], 7);
17647 let output: Value = serde_json::from_str(&record.output).unwrap();
17648 assert_eq!(output["requested_name"], "Context Echo");
17649 assert_eq!(output["canonical_id"], "context_echo");
17650 assert_eq!(output["max_results"], 7);
17651 assert_eq!(output["custom_config"]["backend"], "memory");
17652 assert!(record.metadata.contains_key("effective_limits"));
17653 assert!(record.metadata.contains_key("policy_snapshot"));
17654 }
17655
17656 #[tokio::test]
17657 async fn test_runtime_control_cancels_active_tool_call() {
17658 let mock = mock_with_response("hello");
17659 let agent = Arc::new(
17660 AgentBuilder::new()
17661 .system_prompt("You are helpful.")
17662 .llm(Arc::new(mock))
17663 .tool(Arc::new(SlowTool))
17664 .build()
17665 .unwrap(),
17666 );
17667 let control = agent.runtime_control();
17668 let running_agent = Arc::clone(&agent);
17669 let handle = tokio::spawn(async move {
17670 running_agent
17671 .invoke_tool(ToolExecutionRequest::new(
17672 "slow-call",
17673 "slow",
17674 serde_json::json!({}),
17675 ToolCallSource::Manual,
17676 ))
17677 .await
17678 .unwrap()
17679 });
17680
17681 tokio::time::sleep(std::time::Duration::from_millis(100)).await;
17682 control.cancel_all();
17683 let record = handle.await.unwrap();
17684
17685 assert!(record.executed);
17686 assert!(record.cancelled);
17687 assert!(!record.success);
17688 assert!(record.cancellation_reason.is_some());
17689 }
17690
17691 #[tokio::test]
17693 async fn cancelled_tool_does_not_enter_fallback() {
17694 let fallback_calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
17695 let agent = Arc::new(
17696 AgentBuilder::new()
17697 .system_prompt("Test cancellation before fallback.")
17698 .llm(Arc::new(mock_with_response("done")))
17699 .tool(Arc::new(SlowTool))
17700 .tool(Arc::new(RecoveryTestTool {
17701 id: "fallback".to_string(),
17702 succeeds: true,
17703 calls: Arc::clone(&fallback_calls),
17704 max_output_chars: None,
17705 }))
17706 .recovery_manager(recovery_manager_with_fallbacks([(
17707 "slow".to_string(),
17708 "fallback".to_string(),
17709 )]))
17710 .build()
17711 .unwrap(),
17712 );
17713 let control = agent.runtime_control();
17714 let running_agent = Arc::clone(&agent);
17715 let handle = tokio::spawn(async move {
17716 running_agent
17717 .invoke_tool(ToolExecutionRequest::new(
17718 "cancelled-fallback-call",
17719 "slow",
17720 serde_json::json!({}),
17721 ToolCallSource::Manual,
17722 ))
17723 .await
17724 .unwrap()
17725 });
17726
17727 tokio::time::sleep(Duration::from_millis(100)).await;
17728 control.cancel_all();
17729 let record = handle.await.unwrap();
17730
17731 assert!(record.executed);
17732 assert!(record.cancelled);
17733 assert!(!record.success);
17734 assert_eq!(record.canonical_id, "slow");
17735 assert_eq!(fallback_calls.load(Ordering::SeqCst), 0);
17736 assert_eq!(agent.tool_call_history().len(), 1);
17737 }
17738
17739 #[tokio::test]
17740 async fn non_idempotent_tool_calls_are_not_retried() {
17741 use ai_agents_recovery::{ErrorRecoveryConfig, ToolRecoveryConfig, ToolRetryConfig};
17742
17743 let mock = mock_with_response("hello");
17744 let calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
17745 let agent = AgentBuilder::new()
17746 .system_prompt("You are helpful.")
17747 .llm(Arc::new(mock))
17748 .tool(Arc::new(FlakyWriteTool {
17749 calls: Arc::clone(&calls),
17750 }))
17751 .recovery_manager(RecoveryManager::new(ErrorRecoveryConfig {
17752 tools: ToolRecoveryConfig {
17753 default: ToolRetryConfig {
17754 max_retries: 2,
17755 ..Default::default()
17756 },
17757 ..Default::default()
17758 },
17759 ..Default::default()
17760 }))
17761 .build()
17762 .unwrap();
17763
17764 let record = agent
17765 .invoke_tool(ToolExecutionRequest::new(
17766 "flaky-call",
17767 "flaky_write",
17768 serde_json::json!({"path": "./tmp.txt"}),
17769 ToolCallSource::Manual,
17770 ))
17771 .await
17772 .unwrap();
17773
17774 assert!(!record.success);
17775 assert_eq!(calls.load(Ordering::SeqCst), 1);
17776 }
17777
17778 #[tokio::test]
17779 async fn safely_retryable_tool_receives_a_fresh_deadline_per_attempt() {
17780 use ai_agents_recovery::{ErrorRecoveryConfig, ToolRecoveryConfig, ToolRetryConfig};
17781
17782 let calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
17783 let deadlines = Arc::new(parking_lot::Mutex::new(Vec::new()));
17784 let remaining_ms = Arc::new(parking_lot::Mutex::new(Vec::new()));
17785 let agent = AgentBuilder::new()
17786 .system_prompt("Test retry deadlines.")
17787 .llm(Arc::new(mock_with_response("done")))
17788 .tool(Arc::new(RetryDeadlineTool {
17789 calls: Arc::clone(&calls),
17790 deadlines: Arc::clone(&deadlines),
17791 remaining_ms: Arc::clone(&remaining_ms),
17792 }))
17793 .recovery_manager(RecoveryManager::new(ErrorRecoveryConfig {
17794 tools: ToolRecoveryConfig {
17795 per_tool: HashMap::from([(
17796 "retry_deadline".to_string(),
17797 ToolRetryConfig {
17798 max_retries: 1,
17799 ..Default::default()
17800 },
17801 )]),
17802 ..Default::default()
17803 },
17804 ..Default::default()
17805 }))
17806 .build()
17807 .unwrap();
17808
17809 let record = agent
17810 .invoke_tool(ToolExecutionRequest::new(
17811 "retry-deadline-call",
17812 "retry_deadline",
17813 serde_json::json!({}),
17814 ToolCallSource::Manual,
17815 ))
17816 .await
17817 .unwrap();
17818
17819 assert!(record.executed);
17820 assert!(record.success);
17821 assert_eq!(calls.load(Ordering::SeqCst), 2);
17822 let deadlines = deadlines.lock();
17823 assert_eq!(deadlines.len(), 2);
17824 assert!(
17825 deadlines[1] > deadlines[0],
17826 "retry inherited the first invocation deadline"
17827 );
17828 let remaining_ms = remaining_ms.lock();
17829 assert_eq!(remaining_ms.len(), 2);
17830 assert!(
17831 remaining_ms
17832 .iter()
17833 .all(|remaining| (800..=1_000).contains(remaining))
17834 );
17835 }
17836
17837 #[tokio::test]
17839 async fn call_classification_timeout_controls_deadline_and_timer() {
17840 let calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
17841 let remaining_ms = Arc::new(parking_lot::Mutex::new(Vec::new()));
17842 let agent = AgentBuilder::new()
17843 .system_prompt("Test call-level timeout.")
17844 .llm(Arc::new(mock_with_response("done")))
17845 .tool(Arc::new(ClassifiedTimeoutTool {
17846 id: "classified_timeout",
17847 calls: Arc::clone(&calls),
17848 timeout_ms: 100,
17849 sleep_ms: 150,
17850 requires_approval: false,
17851 remaining_ms: Arc::clone(&remaining_ms),
17852 }))
17853 .build()
17854 .unwrap();
17855
17856 let started = Instant::now();
17857 let record = agent
17858 .invoke_tool(ToolExecutionRequest::new(
17859 "classified-timeout-call",
17860 "classified_timeout",
17861 serde_json::json!({}),
17862 ToolCallSource::Manual,
17863 ))
17864 .await
17865 .unwrap();
17866
17867 assert!(record.executed);
17868 assert!(record.timed_out);
17869 assert!(!record.success);
17870 assert_eq!(calls.load(Ordering::SeqCst), 1);
17871 assert!(started.elapsed() < Duration::from_secs(1));
17872 let remaining_ms = remaining_ms.lock();
17873 assert_eq!(remaining_ms.len(), 1);
17874 assert!((1..=100).contains(&remaining_ms[0]));
17875 }
17876
17877 #[tokio::test]
17879 async fn recovery_timeout_only_lowers_call_and_policy_timeouts() {
17880 use ai_agents_recovery::{ErrorRecoveryConfig, ToolRecoveryConfig, ToolRetryConfig};
17881
17882 let calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
17883 let remaining_ms = Arc::new(parking_lot::Mutex::new(Vec::new()));
17884 let agent = AgentBuilder::new()
17885 .system_prompt("Test recovery timeout.")
17886 .llm(Arc::new(mock_with_response("done")))
17887 .tool(Arc::new(ClassifiedTimeoutTool {
17888 id: "recovery_timeout",
17889 calls: Arc::clone(&calls),
17890 timeout_ms: 1_000,
17891 sleep_ms: 150,
17892 requires_approval: false,
17893 remaining_ms: Arc::clone(&remaining_ms),
17894 }))
17895 .recovery_manager(RecoveryManager::new(ErrorRecoveryConfig {
17896 tools: ToolRecoveryConfig {
17897 per_tool: HashMap::from([(
17898 "recovery_timeout".to_string(),
17899 ToolRetryConfig {
17900 timeout_ms: Some(100),
17901 ..Default::default()
17902 },
17903 )]),
17904 ..Default::default()
17905 },
17906 ..Default::default()
17907 }))
17908 .build()
17909 .unwrap();
17910
17911 let started = Instant::now();
17912 let record = agent
17913 .invoke_tool(ToolExecutionRequest::new(
17914 "recovery-timeout-call",
17915 "recovery_timeout",
17916 serde_json::json!({}),
17917 ToolCallSource::Manual,
17918 ))
17919 .await
17920 .unwrap();
17921
17922 assert!(record.executed);
17923 assert!(record.timed_out);
17924 assert!(!record.success);
17925 assert_eq!(calls.load(Ordering::SeqCst), 1);
17926 assert!(started.elapsed() < Duration::from_secs(1));
17927 assert_eq!(record.metadata["effective_limits"]["timeout_ms"], 100);
17928 let remaining_ms = remaining_ms.lock();
17929 assert_eq!(remaining_ms.len(), 1);
17930 assert!((1..=100).contains(&remaining_ms[0]));
17931 }
17932
17933 #[tokio::test]
17935 async fn recovery_default_timeout_controls_deadline_and_timer() {
17936 use ai_agents_recovery::{ErrorRecoveryConfig, ToolRecoveryConfig, ToolRetryConfig};
17937
17938 let calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
17939 let remaining_ms = Arc::new(parking_lot::Mutex::new(Vec::new()));
17940 let agent = AgentBuilder::new()
17941 .system_prompt("Test default recovery timeout.")
17942 .llm(Arc::new(mock_with_response("done")))
17943 .tool(Arc::new(ClassifiedTimeoutTool {
17944 id: "default_recovery_timeout",
17945 calls: Arc::clone(&calls),
17946 timeout_ms: 1_000,
17947 sleep_ms: 150,
17948 requires_approval: false,
17949 remaining_ms: Arc::clone(&remaining_ms),
17950 }))
17951 .recovery_manager(RecoveryManager::new(ErrorRecoveryConfig {
17952 tools: ToolRecoveryConfig {
17953 default: ToolRetryConfig {
17954 timeout_ms: Some(100),
17955 ..Default::default()
17956 },
17957 ..Default::default()
17958 },
17959 ..Default::default()
17960 }))
17961 .build()
17962 .unwrap();
17963
17964 let started = Instant::now();
17965 let record = agent
17966 .invoke_tool(ToolExecutionRequest::new(
17967 "default-recovery-timeout-call",
17968 "default_recovery_timeout",
17969 serde_json::json!({}),
17970 ToolCallSource::Manual,
17971 ))
17972 .await
17973 .unwrap();
17974
17975 assert!(record.executed);
17976 assert!(record.timed_out);
17977 assert!(!record.success);
17978 assert_eq!(calls.load(Ordering::SeqCst), 1);
17979 assert!(started.elapsed() < Duration::from_secs(1));
17980 assert_eq!(record.metadata["effective_limits"]["timeout_ms"], 100);
17981 let remaining_ms = remaining_ms.lock();
17982 assert_eq!(remaining_ms.len(), 1);
17983 assert!((1..=100).contains(&remaining_ms[0]));
17984 }
17985
17986 #[test]
17988 fn recovery_timeout_cannot_widen_security_baseline() {
17989 let security_engine = ToolSecurityEngine::new(ToolSecurityConfig {
17990 default_timeout_ms: 100,
17991 ..Default::default()
17992 });
17993 let safety = ToolSafetyMetadata::compute();
17994 let mut classification = ToolCallClassification::from_metadata(&safety);
17995 classification.timeout_ms = Some(500);
17996
17997 let (limits, timeout) = RuntimeAgent::effective_tool_limits(
17998 &security_engine,
17999 "recovery_cannot_widen",
18000 &safety,
18001 &classification,
18002 Some(1_000),
18003 )
18004 .unwrap();
18005
18006 assert_eq!(limits.timeout_ms, Some(100));
18007 assert_eq!(timeout.timer, Duration::from_millis(100));
18008 }
18009
18010 #[tokio::test]
18012 async fn invalid_call_timeout_stops_before_approval_or_tool_invocation() {
18013 let tool_calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
18014 let approval_calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
18015 let remaining_ms = Arc::new(parking_lot::Mutex::new(Vec::new()));
18016 let mut security = ToolSecurityConfig {
18017 enabled: true,
18018 ..Default::default()
18019 };
18020 security.tools.insert(
18021 "invalid_call_timeout".to_string(),
18022 ai_agents_tools::ToolPolicyConfig {
18023 require_confirmation: true,
18024 ..Default::default()
18025 },
18026 );
18027 let agent = AgentBuilder::new()
18028 .system_prompt("Test invalid call timeout.")
18029 .llm(Arc::new(mock_with_response("done")))
18030 .tool(Arc::new(ClassifiedTimeoutTool {
18031 id: "invalid_call_timeout",
18032 calls: Arc::clone(&tool_calls),
18033 timeout_ms: u64::MAX,
18034 sleep_ms: 0,
18035 requires_approval: false,
18036 remaining_ms,
18037 }))
18038 .tool_security(ToolSecurityEngine::new(security))
18039 .hitl_engine(HITLEngine::new(ai_agents_hitl::HITLConfig::default()))
18040 .approval_handler(Arc::new(CountingApprovalHandler {
18041 calls: Arc::clone(&approval_calls),
18042 }))
18043 .build()
18044 .unwrap();
18045
18046 let error = agent
18047 .invoke_tool(ToolExecutionRequest::new(
18048 "invalid-call-timeout",
18049 "invalid_call_timeout",
18050 serde_json::json!({}),
18051 ToolCallSource::Manual,
18052 ))
18053 .await
18054 .unwrap_err();
18055
18056 assert!(error.to_string().contains(
18057 "effective tool timeout_ms must be no greater than 3153600000000000 milliseconds"
18058 ));
18059 assert_eq!(approval_calls.load(Ordering::SeqCst), 0);
18060 assert_eq!(tool_calls.load(Ordering::SeqCst), 0);
18061 }
18062
18063 #[tokio::test]
18065 async fn invalid_modified_call_timeout_stops_before_lock_or_invocation() {
18066 use ai_agents_hitl::CallbackHandler;
18067
18068 let blocker_gate = PathMutationGate::new();
18069 let tool_calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
18070 let hooks = Arc::new(ToolLifecycleRecordingHooks::new());
18071 let handler = CallbackHandler::new(|_| ApprovalResult::Modified {
18072 changes: HashMap::from([("invalid_timeout".to_string(), Value::Bool(true))]),
18073 });
18074 let agent = Arc::new(
18075 AgentBuilder::new()
18076 .system_prompt("Test final call timeout validation.")
18077 .llm(Arc::new(mock_with_response("done")))
18078 .tool(Arc::new(BlockingPathMutationTool {
18079 id: "timeout_lock_blocker",
18080 path_fields: vec![ai_agents_core::PathPolicyBinding::write("path")],
18081 gate: blocker_gate.clone(),
18082 }))
18083 .tool(Arc::new(ApprovalModifiedTimeoutTool {
18084 calls: Arc::clone(&tool_calls),
18085 }))
18086 .hitl_engine(HITLEngine::new(ai_agents_hitl::HITLConfig::default()))
18087 .approval_handler(Arc::new(handler))
18088 .hooks(hooks.clone())
18089 .build()
18090 .unwrap(),
18091 );
18092 let blocking_agent = Arc::clone(&agent);
18093 let blocker = tokio::spawn(async move {
18094 blocking_agent
18095 .invoke_tool(ToolExecutionRequest::new(
18096 "timeout-lock-blocker",
18097 "timeout_lock_blocker",
18098 serde_json::json!({"path": "./shared-timeout.txt"}),
18099 ToolCallSource::Manual,
18100 ))
18101 .await
18102 .unwrap()
18103 });
18104 blocker_gate.wait_until_entered().await;
18105
18106 let record = tokio::time::timeout(
18107 Duration::from_millis(500),
18108 agent.invoke_tool(ToolExecutionRequest::new(
18109 "invalid-modified-timeout",
18110 "approval_modified_timeout",
18111 serde_json::json!({
18112 "path": "./shared-timeout.txt",
18113 "invalid_timeout": false
18114 }),
18115 ToolCallSource::Manual,
18116 )),
18117 )
18118 .await
18119 .expect("final timeout validation must not wait for the held path lock")
18120 .unwrap();
18121
18122 blocker_gate.release();
18123 assert!(blocker.await.unwrap().success);
18124 assert!(!record.executed);
18125 assert!(!record.success);
18126 assert_eq!(record.policy.outcome, PermissionOutcome::Deny);
18127 assert!(record.output.contains(
18128 "effective tool timeout_ms must be no greater than 3153600000000000 milliseconds"
18129 ));
18130 assert_eq!(tool_calls.load(Ordering::SeqCst), 0);
18131 let invalid_request_events = hooks
18132 .events()
18133 .into_iter()
18134 .filter(|event| event.contains("approval_modified_timeout") || event == "error")
18135 .collect::<Vec<_>>();
18136 assert_eq!(
18137 invalid_request_events,
18138 vec![
18139 "start:approval_modified_timeout",
18140 "complete:approval_modified_timeout:false",
18141 "record:approval_modified_timeout:false",
18142 "error"
18143 ]
18144 );
18145 }
18146
18147 #[tokio::test]
18148 async fn side_effecting_tools_are_serialized_per_resource() {
18149 let mock = mock_with_response("hello");
18150 let active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
18151 let max_active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
18152 let agent = Arc::new(
18153 AgentBuilder::new()
18154 .system_prompt("You are helpful.")
18155 .llm(Arc::new(mock))
18156 .tool(Arc::new(LockedWriteTool {
18157 active: Arc::clone(&active),
18158 max_active: Arc::clone(&max_active),
18159 }))
18160 .build()
18161 .unwrap(),
18162 );
18163
18164 let left = {
18165 let agent = Arc::clone(&agent);
18166 tokio::spawn(async move {
18167 agent
18168 .invoke_tool(ToolExecutionRequest::new(
18169 "lock-1",
18170 "locked_write",
18171 serde_json::json!({"path": "./same.txt"}),
18172 ToolCallSource::Manual,
18173 ))
18174 .await
18175 .unwrap()
18176 })
18177 };
18178 let right = {
18179 let agent = Arc::clone(&agent);
18180 tokio::spawn(async move {
18181 agent
18182 .invoke_tool(ToolExecutionRequest::new(
18183 "lock-2",
18184 "locked_write",
18185 serde_json::json!({"path": "./same.txt"}),
18186 ToolCallSource::Manual,
18187 ))
18188 .await
18189 .unwrap()
18190 })
18191 };
18192
18193 let left = left.await.unwrap();
18194 let right = right.await.unwrap();
18195 assert!(left.success);
18196 assert!(right.success);
18197 assert_eq!(max_active.load(Ordering::SeqCst), 1);
18198 }
18199
18200 #[tokio::test]
18201 async fn path_resources_use_shared_global_lock_and_cleanup() {
18202 let active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
18203 let max_active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
18204 let bindings = ai_agents_core::ToolPolicyBindings {
18205 path_fields: vec![
18206 ai_agents_core::PathPolicyBinding::read_write("source_path"),
18207 ai_agents_core::PathPolicyBinding::write("destination_path"),
18208 ],
18209 ..Default::default()
18210 };
18211 let classification = ai_agents_core::ToolCallClassification::from_metadata(
18212 &MultiResourceWriteTool {
18213 active: Arc::clone(&active),
18214 max_active: Arc::clone(&max_active),
18215 }
18216 .safety_metadata(),
18217 );
18218 let left_args = serde_json::json!({
18219 "source_path": "./a/../first.txt",
18220 "destination_path": "./second.txt"
18221 });
18222 let right_args = serde_json::json!({
18223 "source_path": "./second.txt",
18224 "destination_path": "./first.txt"
18225 });
18226 let left_keys = tool_resource_lock_keys(
18227 "multi_resource_write",
18228 &left_args,
18229 &bindings,
18230 &classification,
18231 );
18232 let right_keys = tool_resource_lock_keys(
18233 "multi_resource_write",
18234 &right_args,
18235 &bindings,
18236 &classification,
18237 );
18238 assert_eq!(left_keys, right_keys);
18239 assert_eq!(left_keys, vec!["path-mutation:global".to_string()]);
18240
18241 let locks = new_tool_resource_locks();
18242 let build_agent = || {
18243 AgentBuilder::new()
18244 .system_prompt("Test shared resource locks.")
18245 .llm(Arc::new(mock_with_response("done")))
18246 .tool(Arc::new(MultiResourceWriteTool {
18247 active: Arc::clone(&active),
18248 max_active: Arc::clone(&max_active),
18249 }))
18250 .build()
18251 .unwrap()
18252 .with_shared_resource_locks(Arc::clone(&locks))
18253 };
18254 let left_agent = Arc::new(build_agent());
18255 let right_agent = Arc::new(build_agent());
18256 let left = tokio::spawn(async move {
18257 left_agent
18258 .invoke_tool(ToolExecutionRequest::new(
18259 "multi-left",
18260 "multi_resource_write",
18261 left_args,
18262 ToolCallSource::Manual,
18263 ))
18264 .await
18265 .unwrap()
18266 });
18267 let right = tokio::spawn(async move {
18268 right_agent
18269 .invoke_tool(ToolExecutionRequest::new(
18270 "multi-right",
18271 "multi_resource_write",
18272 right_args,
18273 ToolCallSource::Manual,
18274 ))
18275 .await
18276 .unwrap()
18277 });
18278 let (left, right) = tokio::time::timeout(std::time::Duration::from_secs(2), async {
18279 tokio::join!(left, right)
18280 })
18281 .await
18282 .expect("reversed resource acquisition must not deadlock");
18283
18284 assert!(left.unwrap().success);
18285 assert!(right.unwrap().success);
18286 assert_eq!(max_active.load(Ordering::SeqCst), 1);
18287 assert!(locks.read().is_empty());
18288 }
18289
18290 #[tokio::test]
18291 async fn global_path_lock_serializes_copy_destination_with_file_write() {
18292 assert_path_mutation_pair_serialized(
18293 "copy_path",
18294 CopyPathTool::new().policy_bindings().path_fields,
18295 serde_json::json!({
18296 "source_path": "./source.txt",
18297 "destination_path": "./shared.txt"
18298 }),
18299 "file_write",
18300 FileWriteTool::new().policy_bindings().path_fields,
18301 serde_json::json!({"path": "./shared.txt"}),
18302 )
18303 .await;
18304 }
18305
18306 #[tokio::test]
18307 async fn parent_and_spawned_runtime_share_global_path_lock() {
18308 let workspace = MutationTestWorkspace::new();
18309 let destination = workspace.root.join("spawned.txt");
18310 let parent_gate = PathMutationGate::new();
18311 let parent = Arc::new(
18312 AgentBuilder::from_yaml(
18313 r#"
18314name: LockParent
18315system_prompt: parent
18316llm:
18317 default: default
18318tools:
18319 - parent_path_write
18320spawner:
18321 shared_llms: true
18322"#,
18323 )
18324 .unwrap()
18325 .llm(Arc::new(mock_with_response("done")))
18326 .auto_configure_spawner()
18327 .await
18328 .unwrap()
18329 .tool(Arc::new(BlockingPathMutationTool {
18330 id: "parent_path_write",
18331 path_fields: vec![ai_agents_core::PathPolicyBinding::write("path")],
18332 gate: parent_gate.clone(),
18333 }))
18334 .build()
18335 .unwrap(),
18336 );
18337
18338 let mut child_spec = crate::spec::AgentSpec {
18339 name: "LockChild".to_string(),
18340 system_prompt: "child".to_string(),
18341 tools: Some(vec![crate::spec::ToolEntry::Simple(
18342 "file_write".to_string(),
18343 )]),
18344 ..Default::default()
18345 };
18346 child_spec.tool_security.enabled = true;
18347 child_spec.tool_security.fail_closed = true;
18348 let file_write_policy = ai_agents_tools::ToolPolicyConfig {
18349 write_paths: vec![workspace.root.to_string_lossy().into_owned()],
18350 allow_without_confirmation: true,
18351 ..Default::default()
18352 };
18353 child_spec
18354 .tool_security
18355 .tools
18356 .insert("file_write".to_string(), file_write_policy);
18357 let spawned = parent
18358 .spawner()
18359 .unwrap()
18360 .spawn_from_spec(child_spec)
18361 .await
18362 .unwrap();
18363 assert!(Arc::ptr_eq(
18364 &parent.resource_locks,
18365 &spawned.agent.resource_locks
18366 ));
18367 assert!(!Arc::ptr_eq(
18368 &parent.runtime_control,
18369 &spawned.agent.runtime_control
18370 ));
18371
18372 let parent_call = {
18373 let parent = Arc::clone(&parent);
18374 let destination = destination.clone();
18375 tokio::spawn(async move {
18376 parent
18377 .invoke_tool(ToolExecutionRequest::new(
18378 "parent-lock-holder",
18379 "parent_path_write",
18380 serde_json::json!({"path": destination}),
18381 ToolCallSource::Manual,
18382 ))
18383 .await
18384 .unwrap()
18385 })
18386 };
18387 parent_gate.wait_until_entered().await;
18388
18389 let child_call = {
18390 let child = Arc::clone(&spawned.agent);
18391 let destination = destination.clone();
18392 tokio::spawn(async move {
18393 child
18394 .invoke_tool(ToolExecutionRequest::new(
18395 "spawned-file-write",
18396 "file_write",
18397 serde_json::json!({
18398 "path": destination,
18399 "content": "spawned",
18400 "dry_run": false
18401 }),
18402 ToolCallSource::Manual,
18403 ))
18404 .await
18405 .unwrap()
18406 })
18407 };
18408 wait_for_resource_lock_strong_count(&parent.resource_locks, 2).await;
18409 assert!(!child_call.is_finished());
18410
18411 parent_gate.release();
18412 let (parent_record, child_record) =
18413 tokio::time::timeout(std::time::Duration::from_secs(2), async {
18414 tokio::join!(parent_call, child_call)
18415 })
18416 .await
18417 .expect("parent and spawned path mutations did not finish");
18418 assert!(parent_record.unwrap().success);
18419 assert!(child_record.unwrap().success);
18420 assert_eq!(std::fs::read_to_string(destination).unwrap(), "spawned");
18421 assert!(parent.resource_locks.read().is_empty());
18422 }
18423
18424 #[tokio::test]
18425 async fn cancelled_global_path_lock_waiter_does_not_retain_weak_entry() {
18426 let locks = new_tool_resource_locks();
18427 let holder_gate = PathMutationGate::new();
18428 let waiter_gate = PathMutationGate::new();
18429 waiter_gate.release();
18430 let holder = Arc::new(
18431 AgentBuilder::new()
18432 .system_prompt("Hold the global path lock.")
18433 .llm(Arc::new(mock_with_response("done")))
18434 .tool(Arc::new(BlockingPathMutationTool {
18435 id: "holder_write",
18436 path_fields: vec![ai_agents_core::PathPolicyBinding::write("path")],
18437 gate: holder_gate.clone(),
18438 }))
18439 .build()
18440 .unwrap()
18441 .with_shared_resource_locks(Arc::clone(&locks)),
18442 );
18443 let waiter = Arc::new(
18444 AgentBuilder::new()
18445 .system_prompt("Wait for the global path lock.")
18446 .llm(Arc::new(mock_with_response("done")))
18447 .tool(Arc::new(BlockingPathMutationTool {
18448 id: "waiter_write",
18449 path_fields: vec![ai_agents_core::PathPolicyBinding::write("path")],
18450 gate: waiter_gate.clone(),
18451 }))
18452 .build()
18453 .unwrap()
18454 .with_shared_resource_locks(Arc::clone(&locks)),
18455 );
18456
18457 let holder_call = {
18458 let holder = Arc::clone(&holder);
18459 tokio::spawn(async move {
18460 holder
18461 .invoke_tool(ToolExecutionRequest::new(
18462 "holder-call",
18463 "holder_write",
18464 serde_json::json!({"path": "./shared.txt"}),
18465 ToolCallSource::Manual,
18466 ))
18467 .await
18468 .unwrap()
18469 })
18470 };
18471 holder_gate.wait_until_entered().await;
18472
18473 let waiter_call = {
18474 let waiter = Arc::clone(&waiter);
18475 tokio::spawn(async move {
18476 waiter
18477 .invoke_tool(ToolExecutionRequest::new(
18478 "waiter-call",
18479 "waiter_write",
18480 serde_json::json!({"path": "./shared.txt"}),
18481 ToolCallSource::Manual,
18482 ))
18483 .await
18484 .unwrap()
18485 })
18486 };
18487 wait_for_resource_lock_strong_count(&locks, 2).await;
18488 waiter.runtime_control().cancel_all();
18489
18490 let waiter_record = tokio::time::timeout(std::time::Duration::from_secs(2), waiter_call)
18491 .await
18492 .expect("cancelled lock waiter did not finish")
18493 .unwrap();
18494 assert!(!waiter_record.success);
18495 assert!(!waiter_record.executed);
18496 assert!(waiter_record.cancelled);
18497 assert_eq!(
18498 waiter_record.cancellation_reason.as_deref(),
18499 Some("runtime control cancellation")
18500 );
18501 assert!(!waiter_gate.entered.load(Ordering::SeqCst));
18502 assert_eq!(
18503 locks
18504 .read()
18505 .get("path-mutation:global")
18506 .map_or(0, |lock| lock.strong_count()),
18507 1
18508 );
18509
18510 holder_gate.release();
18511 let holder_record = tokio::time::timeout(std::time::Duration::from_secs(2), holder_call)
18512 .await
18513 .expect("lock holder did not finish")
18514 .unwrap();
18515 assert!(holder_record.success);
18516 assert!(locks.read().is_empty());
18517 }
18518
18519 #[tokio::test]
18520 async fn path_mutation_policy_and_approval_denials_do_not_invoke_tools() {
18521 for denial in [MutationDenial::Policy, MutationDenial::Approval] {
18522 let tools: [Arc<dyn Tool>; 3] = [
18523 Arc::new(CopyPathTool::new()),
18524 Arc::new(MovePathTool::new()),
18525 Arc::new(DeletePathTool::new()),
18526 ];
18527 for tool in tools {
18528 assert_path_mutation_denied(tool, denial).await;
18529 }
18530 }
18531 }
18532
18533 #[tokio::test]
18534 async fn policy_denial_keeps_executor_hook_lifecycle_and_record_authority() {
18535 let workspace = MutationTestWorkspace::new();
18536 let target = workspace.root.join("denied.txt");
18537 let hooks = Arc::new(ToolLifecycleRecordingHooks::new());
18538 let agent = AgentBuilder::new()
18539 .system_prompt("Test denied tool hooks.")
18540 .llm(Arc::new(mock_with_response("done")))
18541 .tool(Arc::new(FileWriteTool::new()))
18542 .tool_security(ToolSecurityEngine::new(mutation_denial_security_config(
18543 "file_write",
18544 &workspace.root,
18545 MutationDenial::Policy,
18546 )))
18547 .hooks(hooks.clone())
18548 .build()
18549 .unwrap();
18550
18551 let record = agent
18552 .invoke_tool(ToolExecutionRequest::new(
18553 "denied-hook-call",
18554 "file_write",
18555 serde_json::json!({
18556 "path": target.to_string_lossy(),
18557 "content": "blocked"
18558 }),
18559 ToolCallSource::Manual,
18560 ))
18561 .await
18562 .unwrap();
18563
18564 assert!(!record.executed);
18565 assert!(!record.success);
18566 assert_eq!(record.policy.outcome, PermissionOutcome::Deny);
18567 assert_eq!(
18568 hooks.events(),
18569 vec![
18570 "start:file_write",
18571 "complete:file_write:false",
18572 "record:file_write:false",
18573 "error"
18574 ]
18575 );
18576 assert!(!target.exists());
18577 }
18578
18579 #[tokio::test]
18580 async fn approval_argument_changes_are_rechecked_against_final_scope() {
18581 let active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
18582 let max_active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
18583 let entered = Arc::new(tokio::sync::Barrier::new(2));
18584 let release = Arc::new(tokio::sync::Notify::new());
18585 let handler = Arc::new(BlockingApprovalHandler {
18586 entered: Arc::clone(&entered),
18587 release: Arc::clone(&release),
18588 result: ApprovalResult::Modified {
18589 changes: HashMap::from([(
18590 "path".to_string(),
18591 Value::String("./after-approval.txt".to_string()),
18592 )]),
18593 },
18594 });
18595 let agent = Arc::new(
18596 AgentBuilder::new()
18597 .system_prompt("Test final scope validation.")
18598 .llm(Arc::new(mock_with_response("done")))
18599 .tool(Arc::new(LockedWriteTool {
18600 active: Arc::clone(&active),
18601 max_active: Arc::clone(&max_active),
18602 }))
18603 .tool_security(ToolSecurityEngine::new(approval_security_config(true)))
18604 .hitl_engine(HITLEngine::new(ai_agents_hitl::HITLConfig::default()))
18605 .approval_handler(handler)
18606 .build()
18607 .unwrap(),
18608 );
18609 let control = agent.runtime_control();
18610 let running = Arc::clone(&agent);
18611 let call = tokio::spawn(async move {
18612 running
18613 .invoke_tool(ToolExecutionRequest::new(
18614 "approval-scope",
18615 "locked_write",
18616 serde_json::json!({"path": "./before-approval.txt"}),
18617 ToolCallSource::Manual,
18618 ))
18619 .await
18620 .unwrap()
18621 });
18622 entered.wait().await;
18623 let expected_version = control.set_tool_scope(Vec::new());
18624 release.notify_one();
18625 let record = call.await.unwrap();
18626
18627 assert!(!record.executed);
18628 assert!(!record.success);
18629 assert_eq!(record.runtime_config_version, expected_version);
18630 assert_eq!(record.executed_arguments["path"], "./after-approval.txt");
18631 assert_eq!(max_active.load(Ordering::SeqCst), 0);
18632 assert_eq!(
18633 record.metadata["runtime_scope_snapshot"],
18634 serde_json::json!([])
18635 );
18636 }
18637
18638 #[tokio::test]
18639 async fn approval_is_rechecked_against_final_policy_snapshot() {
18640 let active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
18641 let max_active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
18642 let entered = Arc::new(tokio::sync::Barrier::new(2));
18643 let release = Arc::new(tokio::sync::Notify::new());
18644 let handler = Arc::new(BlockingApprovalHandler {
18645 entered: Arc::clone(&entered),
18646 release: Arc::clone(&release),
18647 result: ApprovalResult::Approved,
18648 });
18649 let agent = Arc::new(
18650 AgentBuilder::new()
18651 .system_prompt("Test final policy validation.")
18652 .llm(Arc::new(mock_with_response("done")))
18653 .tool(Arc::new(LockedWriteTool {
18654 active: Arc::clone(&active),
18655 max_active: Arc::clone(&max_active),
18656 }))
18657 .tool_security(ToolSecurityEngine::new(approval_security_config(true)))
18658 .hitl_engine(HITLEngine::new(ai_agents_hitl::HITLConfig::default()))
18659 .approval_handler(handler)
18660 .build()
18661 .unwrap(),
18662 );
18663 let control = agent.runtime_control();
18664 let running = Arc::clone(&agent);
18665 let call = tokio::spawn(async move {
18666 running
18667 .invoke_tool(ToolExecutionRequest::new(
18668 "approval-policy",
18669 "locked_write",
18670 serde_json::json!({"path": "./policy.txt"}),
18671 ToolCallSource::Manual,
18672 ))
18673 .await
18674 .unwrap()
18675 });
18676 entered.wait().await;
18677 let expected_version = control.set_tool_security(approval_security_config(false));
18678 release.notify_one();
18679 let record = call.await.unwrap();
18680
18681 assert!(!record.executed);
18682 assert!(!record.success);
18683 assert_eq!(record.runtime_config_version, expected_version);
18684 assert_eq!(record.policy.outcome, PermissionOutcome::Unavailable);
18685 assert_eq!(max_active.load(Ordering::SeqCst), 0);
18686 assert!(record.metadata.contains_key("policy_snapshot"));
18687 }
18688
18689 #[test]
18690 fn invalid_live_policy_does_not_replace_snapshot_or_generation() {
18691 let agent = AgentBuilder::new()
18692 .system_prompt("Test runtime policy validation.")
18693 .llm(Arc::new(mock_with_response("done")))
18694 .build()
18695 .unwrap();
18696 let control = agent.runtime_control();
18697 let mut valid = ToolSecurityConfig::default();
18698 valid.tools.insert(
18699 "web_search".to_string(),
18700 ai_agents_tools::ToolPolicyConfig {
18701 max_results: Some(5),
18702 ..Default::default()
18703 },
18704 );
18705 let generation = control.try_set_tool_security(valid).unwrap();
18706
18707 let mut invalid = ToolSecurityConfig::default();
18708 invalid.tools.insert(
18709 "web_search".to_string(),
18710 ai_agents_tools::ToolPolicyConfig {
18711 max_results: Some(0),
18712 ..Default::default()
18713 },
18714 );
18715 let error = control.try_set_tool_security(invalid).unwrap_err();
18716
18717 assert!(
18718 error
18719 .to_string()
18720 .contains("max_results must be greater than 0")
18721 );
18722 assert_eq!(control.version(), generation);
18723 assert_eq!(
18724 control
18725 .state
18726 .tool_security_override
18727 .read()
18728 .as_ref()
18729 .unwrap()
18730 .config()
18731 .tools["web_search"]
18732 .max_results,
18733 Some(5)
18734 );
18735 }
18736
18737 #[test]
18739 fn invalid_timeout_config_stops_before_approval_or_tool_invocation() {
18740 let tool_calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
18741 let approval_calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
18742 let spec = crate::spec::AgentSpec {
18743 tool_security: ToolSecurityConfig {
18744 enabled: true,
18745 default_timeout_ms: u64::MAX,
18746 ..Default::default()
18747 },
18748 ..Default::default()
18749 };
18750
18751 let result = AgentBuilder::from_spec(spec)
18752 .llm(Arc::new(mock_with_response("done")))
18753 .tool(Arc::new(FlakyWriteTool {
18754 calls: Arc::clone(&tool_calls),
18755 }))
18756 .hitl_engine(HITLEngine::new(ai_agents_hitl::HITLConfig::default()))
18757 .approval_handler(Arc::new(CountingApprovalHandler {
18758 calls: Arc::clone(&approval_calls),
18759 }))
18760 .build();
18761
18762 assert!(result.is_err());
18763 assert_eq!(approval_calls.load(Ordering::SeqCst), 0);
18764 assert_eq!(tool_calls.load(Ordering::SeqCst), 0);
18765 }
18766
18767 #[test]
18769 fn invalid_recovery_timeout_config_stops_before_approval_or_tool_invocation() {
18770 use ai_agents_recovery::{ErrorRecoveryConfig, ToolRecoveryConfig, ToolRetryConfig};
18771
18772 let tool_calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
18773 let approval_calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
18774 let spec = crate::spec::AgentSpec {
18775 error_recovery: ErrorRecoveryConfig {
18776 tools: ToolRecoveryConfig {
18777 default: ToolRetryConfig {
18778 timeout_ms: Some(u64::MAX),
18779 ..Default::default()
18780 },
18781 ..Default::default()
18782 },
18783 ..Default::default()
18784 },
18785 ..Default::default()
18786 };
18787
18788 let result = AgentBuilder::from_spec(spec)
18789 .llm(Arc::new(mock_with_response("done")))
18790 .tool(Arc::new(FlakyWriteTool {
18791 calls: Arc::clone(&tool_calls),
18792 }))
18793 .hitl_engine(HITLEngine::new(ai_agents_hitl::HITLConfig::default()))
18794 .approval_handler(Arc::new(CountingApprovalHandler {
18795 calls: Arc::clone(&approval_calls),
18796 }))
18797 .build();
18798
18799 assert!(result.is_err());
18800 assert_eq!(approval_calls.load(Ordering::SeqCst), 0);
18801 assert_eq!(tool_calls.load(Ordering::SeqCst), 0);
18802 }
18803
18804 #[test]
18806 fn invalid_timeout_policy_does_not_replace_snapshot_or_generation() {
18807 let agent = AgentBuilder::new()
18808 .system_prompt("Test runtime timeout policy validation.")
18809 .llm(Arc::new(mock_with_response("done")))
18810 .build()
18811 .unwrap();
18812 let control = agent.runtime_control();
18813 let valid = ToolSecurityConfig {
18814 default_timeout_ms: 5_000,
18815 ..Default::default()
18816 };
18817 let generation = control.try_set_tool_security(valid).unwrap();
18818
18819 let invalid = ToolSecurityConfig {
18820 default_timeout_ms: MAX_TOOL_TIMEOUT_MS + 1,
18821 ..Default::default()
18822 };
18823 let error = control.try_set_tool_security(invalid).unwrap_err();
18824
18825 assert!(error.to_string().contains(&format!(
18826 "tool_security.default_timeout_ms must be no greater than {MAX_TOOL_TIMEOUT_MS} milliseconds"
18827 )));
18828 assert_eq!(control.version(), generation);
18829 assert_eq!(
18830 control
18831 .state
18832 .tool_security_override
18833 .read()
18834 .as_ref()
18835 .unwrap()
18836 .config()
18837 .default_timeout_ms,
18838 5_000
18839 );
18840 }
18841
18842 #[test]
18844 fn runtime_tool_timeout_conversion_enforces_the_stable_boundary() {
18845 let timeout = RuntimeAgent::validated_tool_timeout(MAX_TOOL_TIMEOUT_MS).unwrap();
18846 assert_eq!(timeout.timer, Duration::from_millis(MAX_TOOL_TIMEOUT_MS));
18847 assert_eq!(
18848 timeout.deadline_delta,
18849 chrono::Duration::milliseconds(MAX_TOOL_TIMEOUT_MS as i64)
18850 );
18851
18852 for timeout_ms in [MAX_TOOL_TIMEOUT_MS + 1, u64::MAX] {
18853 let error = RuntimeAgent::validated_tool_timeout(timeout_ms).unwrap_err();
18854 assert!(error.to_string().contains(&format!(
18855 "effective tool timeout_ms must be no greater than {MAX_TOOL_TIMEOUT_MS} milliseconds"
18856 )));
18857 }
18858 }
18859
18860 #[tokio::test]
18861 async fn persistent_override_preserves_rate_history_within_generation() {
18862 let calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
18863 let agent = AgentBuilder::new()
18864 .system_prompt("Test persistent policy overrides.")
18865 .llm(Arc::new(mock_with_response("done")))
18866 .tool(Arc::new(RecoveryTestTool {
18867 id: "limited_override".to_string(),
18868 succeeds: true,
18869 calls: Arc::clone(&calls),
18870 max_output_chars: None,
18871 }))
18872 .build()
18873 .unwrap();
18874 let mut security = ToolSecurityConfig {
18875 enabled: true,
18876 fail_closed: true,
18877 ..Default::default()
18878 };
18879 let policy = ai_agents_tools::ToolPolicyConfig {
18880 write_paths: vec![".".to_string()],
18881 rate_limit: Some(1),
18882 ..Default::default()
18883 };
18884 security
18885 .tools
18886 .insert("limited_override".to_string(), policy);
18887 let generation = agent.runtime_control().set_tool_security(security);
18888
18889 let first = agent
18890 .invoke_tool(ToolExecutionRequest::new(
18891 "limited-first",
18892 "limited_override",
18893 serde_json::json!({"path": "./limited.txt"}),
18894 ToolCallSource::Manual,
18895 ))
18896 .await
18897 .unwrap();
18898 let second = agent
18899 .invoke_tool(ToolExecutionRequest::new(
18900 "limited-second",
18901 "limited_override",
18902 serde_json::json!({"path": "./limited.txt"}),
18903 ToolCallSource::Manual,
18904 ))
18905 .await
18906 .unwrap();
18907
18908 assert!(first.success);
18909 assert_eq!(first.policy_version, generation);
18910 assert!(!second.executed);
18911 assert!(second.output.contains("Rate limit exceeded"));
18912 assert_eq!(second.policy_version, generation);
18913 assert_eq!(calls.load(Ordering::SeqCst), 1);
18914 }
18915
18916 #[tokio::test]
18917 async fn concurrent_rate_admission_consumes_capacity_atomically() {
18918 let calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
18919 let tool = Arc::new(RecoveryTestTool {
18920 id: "atomic_rate".to_string(),
18921 succeeds: true,
18922 calls: Arc::clone(&calls),
18923 max_output_chars: None,
18924 });
18925 let arguments = serde_json::json!({"path": "./atomic-rate.txt"});
18926 let bindings = tool.policy_bindings();
18927 let classification = tool.classify_call(&arguments);
18928 let resource_keys =
18929 tool_resource_lock_keys(tool.id(), &arguments, &bindings, &classification);
18930 let mut security = ToolSecurityConfig {
18931 enabled: true,
18932 fail_closed: true,
18933 ..Default::default()
18934 };
18935 let policy = ai_agents_tools::ToolPolicyConfig {
18936 write_paths: vec![".".to_string()],
18937 rate_limit: Some(1),
18938 ..Default::default()
18939 };
18940 security.tools.insert(tool.id().to_string(), policy);
18941 let agent = Arc::new(
18942 AgentBuilder::new()
18943 .system_prompt("Test atomic rate admission.")
18944 .llm(Arc::new(mock_with_response("done")))
18945 .tool(tool)
18946 .tool_security(ToolSecurityEngine::new(security))
18947 .build()
18948 .unwrap(),
18949 );
18950 let held = agent
18951 .acquire_tool_resource_locks(&resource_keys)
18952 .await
18953 .unwrap();
18954 let left = {
18955 let agent = Arc::clone(&agent);
18956 let arguments = arguments.clone();
18957 tokio::spawn(async move {
18958 agent
18959 .invoke_tool(ToolExecutionRequest::new(
18960 "atomic-rate-left",
18961 "atomic_rate",
18962 arguments,
18963 ToolCallSource::Manual,
18964 ))
18965 .await
18966 .unwrap()
18967 })
18968 };
18969 let right = {
18970 let agent = Arc::clone(&agent);
18971 tokio::spawn(async move {
18972 agent
18973 .invoke_tool(ToolExecutionRequest::new(
18974 "atomic-rate-right",
18975 "atomic_rate",
18976 arguments,
18977 ToolCallSource::Manual,
18978 ))
18979 .await
18980 .unwrap()
18981 })
18982 };
18983 tokio::time::sleep(std::time::Duration::from_millis(25)).await;
18984 drop(held);
18985 let (left, right) = tokio::join!(left, right);
18986 let records = [left.unwrap(), right.unwrap()];
18987
18988 assert_eq!(records.iter().filter(|record| record.success).count(), 1);
18989 assert_eq!(records.iter().filter(|record| record.executed).count(), 1);
18990 assert!(
18991 records.iter().any(|record| {
18992 !record.executed && record.output.contains("Rate limit exceeded")
18993 })
18994 );
18995 assert_eq!(calls.load(Ordering::SeqCst), 1);
18996 }
18997
18998 #[tokio::test]
18999 async fn changed_policy_generation_invalidates_pending_approval() {
19000 let active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19001 let max_active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19002 let entered = Arc::new(tokio::sync::Barrier::new(2));
19003 let release = Arc::new(tokio::sync::Notify::new());
19004 let handler = Arc::new(BlockingApprovalHandler {
19005 entered: Arc::clone(&entered),
19006 release: Arc::clone(&release),
19007 result: ApprovalResult::Approved,
19008 });
19009 let agent = Arc::new(
19010 AgentBuilder::new()
19011 .system_prompt("Test stale approval denial.")
19012 .llm(Arc::new(mock_with_response("done")))
19013 .tool(Arc::new(LockedWriteTool {
19014 active: Arc::clone(&active),
19015 max_active: Arc::clone(&max_active),
19016 }))
19017 .tool_security(ToolSecurityEngine::new(approval_security_config(true)))
19018 .hitl_engine(HITLEngine::new(ai_agents_hitl::HITLConfig::default()))
19019 .approval_handler(handler)
19020 .build()
19021 .unwrap(),
19022 );
19023 let running = Arc::clone(&agent);
19024 let call = tokio::spawn(async move {
19025 running
19026 .invoke_tool(ToolExecutionRequest::new(
19027 "stale-approval",
19028 "locked_write",
19029 serde_json::json!({"path": "./stale.txt"}),
19030 ToolCallSource::Manual,
19031 ))
19032 .await
19033 .unwrap()
19034 });
19035 entered.wait().await;
19036 let generation = agent
19037 .runtime_control()
19038 .set_tool_security(approval_security_config(true));
19039 release.notify_one();
19040 let record = call.await.unwrap();
19041
19042 assert!(!record.executed);
19043 assert!(record.output.contains("Approval became stale"));
19044 assert_eq!(record.policy_version, generation);
19045 assert_eq!(max_active.load(Ordering::SeqCst), 0);
19046 }
19047
19048 #[tokio::test]
19049 async fn final_policy_reapplies_argument_caps_after_approval_changes() {
19050 use ai_agents_hitl::CallbackHandler;
19051
19052 let mut security = ToolSecurityConfig {
19053 enabled: true,
19054 fail_closed: true,
19055 ..Default::default()
19056 };
19057 let policy = ai_agents_tools::ToolPolicyConfig {
19058 read_paths: vec![".".to_string()],
19059 max_results: Some(5),
19060 require_confirmation: true,
19061 ..Default::default()
19062 };
19063 security.tools.insert("context_echo".to_string(), policy);
19064 let handler = CallbackHandler::new(|_| ApprovalResult::Modified {
19065 changes: HashMap::from([("max_results".to_string(), serde_json::json!(99))]),
19066 });
19067 let agent = AgentBuilder::new()
19068 .system_prompt("Test final argument caps.")
19069 .llm(Arc::new(mock_with_response("done")))
19070 .tool(Arc::new(ContextEchoTool))
19071 .tool_security(ToolSecurityEngine::new(security))
19072 .hitl_engine(HITLEngine::new(ai_agents_hitl::HITLConfig::default()))
19073 .approval_handler(Arc::new(handler))
19074 .build()
19075 .unwrap();
19076
19077 let record = agent
19078 .invoke_tool(ToolExecutionRequest::new(
19079 "final-cap",
19080 "context_echo",
19081 serde_json::json!({"path": ".", "max_results": 1}),
19082 ToolCallSource::Manual,
19083 ))
19084 .await
19085 .unwrap();
19086
19087 assert!(record.success);
19088 assert_eq!(record.executed_arguments["max_results"], 5);
19089 assert_eq!(
19090 record.approval.unwrap().modified_arguments.unwrap()["max_results"],
19091 5
19092 );
19093 }
19094
19095 #[tokio::test]
19096 async fn no_binding_writes_use_canonical_fallback_lock() {
19097 let active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19098 let max_active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19099 let agent = Arc::new(
19100 AgentBuilder::new()
19101 .system_prompt("Test fallback resource locks.")
19102 .llm(Arc::new(mock_with_response("done")))
19103 .tool(Arc::new(NoBindingWriteTool {
19104 active: Arc::clone(&active),
19105 max_active: Arc::clone(&max_active),
19106 }))
19107 .build()
19108 .unwrap(),
19109 );
19110 let left = {
19111 let agent = Arc::clone(&agent);
19112 tokio::spawn(async move {
19113 agent
19114 .invoke_tool(ToolExecutionRequest::new(
19115 "no-binding-left",
19116 "no_binding_write",
19117 serde_json::json!({}),
19118 ToolCallSource::Manual,
19119 ))
19120 .await
19121 .unwrap()
19122 })
19123 };
19124 let right = {
19125 let agent = Arc::clone(&agent);
19126 tokio::spawn(async move {
19127 agent
19128 .invoke_tool(ToolExecutionRequest::new(
19129 "no-binding-right",
19130 "no_binding_write",
19131 serde_json::json!({}),
19132 ToolCallSource::Manual,
19133 ))
19134 .await
19135 .unwrap()
19136 })
19137 };
19138 let (left, right) = tokio::join!(left, right);
19139
19140 assert!(left.unwrap().success);
19141 assert!(right.unwrap().success);
19142 assert_eq!(max_active.load(Ordering::SeqCst), 1);
19143 }
19144
19145 #[tokio::test]
19146 async fn parent_and_child_paths_share_a_resource_lock() {
19147 let active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19148 let max_active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19149 let agent = Arc::new(
19150 AgentBuilder::new()
19151 .system_prompt("Test parent child resource locks.")
19152 .llm(Arc::new(mock_with_response("done")))
19153 .tool(Arc::new(LockedWriteTool {
19154 active: Arc::clone(&active),
19155 max_active: Arc::clone(&max_active),
19156 }))
19157 .build()
19158 .unwrap(),
19159 );
19160 let parent = format!("./lock-parent-{}", uuid::Uuid::new_v4());
19161 let child = format!("{}/child.txt", parent);
19162 let left = {
19163 let agent = Arc::clone(&agent);
19164 tokio::spawn(async move {
19165 agent
19166 .invoke_tool(ToolExecutionRequest::new(
19167 "parent-lock",
19168 "locked_write",
19169 serde_json::json!({"path": parent}),
19170 ToolCallSource::Manual,
19171 ))
19172 .await
19173 .unwrap()
19174 })
19175 };
19176 let right = {
19177 let agent = Arc::clone(&agent);
19178 tokio::spawn(async move {
19179 agent
19180 .invoke_tool(ToolExecutionRequest::new(
19181 "child-lock",
19182 "locked_write",
19183 serde_json::json!({"path": child}),
19184 ToolCallSource::Manual,
19185 ))
19186 .await
19187 .unwrap()
19188 })
19189 };
19190 let (left, right) = tokio::join!(left, right);
19191
19192 assert!(left.unwrap().success);
19193 assert!(right.unwrap().success);
19194 assert_eq!(max_active.load(Ordering::SeqCst), 1);
19195 }
19196
19197 #[tokio::test]
19198 async fn tool_hooks_can_reenter_after_resource_guards_are_dropped() {
19199 let calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19200 let hooks = Arc::new(ReentrantToolHooks {
19201 agent: parking_lot::Mutex::new(None),
19202 invoked: AtomicBool::new(false),
19203 nested_success: AtomicBool::new(false),
19204 });
19205 let agent = Arc::new(
19206 AgentBuilder::new()
19207 .system_prompt("Test hook reentrancy.")
19208 .llm(Arc::new(mock_with_response("done")))
19209 .tool(Arc::new(RecoveryTestTool {
19210 id: "reentrant_write".to_string(),
19211 succeeds: true,
19212 calls: Arc::clone(&calls),
19213 max_output_chars: None,
19214 }))
19215 .hooks(hooks.clone())
19216 .build()
19217 .unwrap(),
19218 );
19219 *hooks.agent.lock() = Some(Arc::downgrade(&agent));
19220 let record = tokio::time::timeout(
19221 std::time::Duration::from_secs(2),
19222 agent.invoke_tool(ToolExecutionRequest::new(
19223 "outer-hook-call",
19224 "reentrant_write",
19225 serde_json::json!({"path": "./hook.txt"}),
19226 ToolCallSource::Manual,
19227 )),
19228 )
19229 .await
19230 .expect("tool completion hook must not retain resource guards")
19231 .unwrap();
19232
19233 assert!(record.success);
19234 assert!(hooks.nested_success.load(Ordering::SeqCst));
19235 assert_eq!(calls.load(Ordering::SeqCst), 2);
19236 }
19237
19238 #[tokio::test]
19240 async fn fallback_finalizes_original_record_before_shared_execution() {
19241 let primary_calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19242 let fallback_calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19243 let hooks = Arc::new(ToolLifecycleRecordingHooks::new());
19244 let agent = AgentBuilder::new()
19245 .system_prompt("Test fallback execution.")
19246 .llm(Arc::new(mock_with_response("done")))
19247 .tool(Arc::new(RecoveryTestTool {
19248 id: "primary".to_string(),
19249 succeeds: false,
19250 calls: Arc::clone(&primary_calls),
19251 max_output_chars: None,
19252 }))
19253 .tool(Arc::new(RecoveryTestTool {
19254 id: "fallback".to_string(),
19255 succeeds: true,
19256 calls: Arc::clone(&fallback_calls),
19257 max_output_chars: None,
19258 }))
19259 .recovery_manager(recovery_manager_with_fallbacks([(
19260 "primary".to_string(),
19261 "fallback".to_string(),
19262 )]))
19263 .hooks(hooks.clone())
19264 .build()
19265 .unwrap();
19266 let record = tokio::time::timeout(
19267 std::time::Duration::from_secs(2),
19268 agent.invoke_tool(ToolExecutionRequest::new(
19269 "fallback-call",
19270 "primary",
19271 serde_json::json!({"path": "./shared.txt"}),
19272 ToolCallSource::Manual,
19273 )),
19274 )
19275 .await
19276 .expect("fallback must not retain the primary resource guard")
19277 .unwrap();
19278
19279 assert_eq!(
19280 hooks.events(),
19281 vec![
19282 "start:primary",
19283 "complete:primary:false",
19284 "record:primary:true",
19285 "error",
19286 "start:fallback",
19287 "complete:fallback:true",
19288 "record:fallback:true",
19289 ]
19290 );
19291 let records = hooks.records();
19292 assert_eq!(records.len(), 2);
19293 let original = &records[0];
19294 assert_eq!(original.canonical_id, "primary");
19295 assert!(matches!(original.source, ToolCallSource::Manual));
19296 assert!(original.executed);
19297 assert!(!original.success);
19298
19299 let fallback = &records[1];
19300 assert_eq!(fallback.canonical_id, "fallback");
19301 assert_eq!(fallback.call_id, "fallback-call");
19302 assert!(matches!(
19303 &fallback.source,
19304 ToolCallSource::Fallback { original_tool } if original_tool == "primary"
19305 ));
19306 assert!(fallback.executed);
19307 assert!(fallback.success);
19308 assert_eq!(record.canonical_id, fallback.canonical_id);
19309 assert_eq!(record.output, fallback.output);
19310
19311 let history = agent.tool_call_history();
19312 assert_eq!(
19313 history
19314 .iter()
19315 .map(|entry| entry.tool_id.as_str())
19316 .collect::<Vec<_>>(),
19317 vec!["primary", "fallback"]
19318 );
19319 assert_eq!(history[0].result.get("success"), Some(&Value::Bool(false)));
19320 assert_eq!(primary_calls.load(Ordering::SeqCst), 1);
19321 assert_eq!(fallback_calls.load(Ordering::SeqCst), 1);
19322 }
19323
19324 #[tokio::test]
19326 async fn self_fallback_cycle_is_denied_before_reinvocation() {
19327 let calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19328 let hooks = Arc::new(ToolLifecycleRecordingHooks::new());
19329 let agent = AgentBuilder::new()
19330 .system_prompt("Test self-fallback cycle admission.")
19331 .llm(Arc::new(mock_with_response("done")))
19332 .tool(Arc::new(RecoveryTestTool {
19333 id: "primary".to_string(),
19334 succeeds: false,
19335 calls: Arc::clone(&calls),
19336 max_output_chars: None,
19337 }))
19338 .recovery_manager(recovery_manager_with_fallbacks([(
19339 "primary".to_string(),
19340 "primary".to_string(),
19341 )]))
19342 .hooks(hooks.clone())
19343 .build()
19344 .unwrap();
19345
19346 let record = tokio::time::timeout(
19347 std::time::Duration::from_secs(2),
19348 agent.invoke_tool(ToolExecutionRequest::new(
19349 "self-fallback-call",
19350 "primary",
19351 serde_json::json!({"path": "./shared.txt"}),
19352 ToolCallSource::Manual,
19353 )),
19354 )
19355 .await
19356 .expect("self fallback must terminate without recursive execution")
19357 .unwrap();
19358
19359 assert_eq!(calls.load(Ordering::SeqCst), 1);
19360 assert_eq!(record.canonical_id, "primary");
19361 assert!(!record.executed);
19362 assert!(!record.success);
19363 assert_eq!(record.policy.outcome, PermissionOutcome::Deny);
19364 assert!(record.output.contains("fallback cycle"));
19365 assert!(matches!(
19366 record.source,
19367 ToolCallSource::Fallback { ref original_tool } if original_tool == "primary"
19368 ));
19369 assert_eq!(
19370 record.metadata.get("fallback_chain"),
19371 Some(&serde_json::json!(["primary"]))
19372 );
19373 assert_eq!(
19374 hooks.events(),
19375 vec![
19376 "start:primary",
19377 "complete:primary:false",
19378 "record:primary:true",
19379 "error",
19380 "complete:primary:false",
19381 "record:primary:false",
19382 "error",
19383 ]
19384 );
19385 assert_eq!(agent.tool_call_history().len(), 2);
19386 }
19387
19388 #[tokio::test]
19390 async fn alias_mediated_fallback_cycle_is_denied_canonically() {
19391 let primary_calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19392 let secondary_calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19393 let hooks = Arc::new(ToolLifecycleRecordingHooks::new());
19394 let agent = AgentBuilder::new()
19395 .system_prompt("Test canonical fallback cycle admission.")
19396 .llm(Arc::new(mock_with_response("done")))
19397 .tool(Arc::new(RecoveryTestTool {
19398 id: "primary".to_string(),
19399 succeeds: false,
19400 calls: Arc::clone(&primary_calls),
19401 max_output_chars: None,
19402 }))
19403 .tool(Arc::new(RecoveryTestTool {
19404 id: "secondary".to_string(),
19405 succeeds: false,
19406 calls: Arc::clone(&secondary_calls),
19407 max_output_chars: None,
19408 }))
19409 .recovery_manager(recovery_manager_with_fallbacks([
19410 ("primary".to_string(), "secondary".to_string()),
19411 ("secondary".to_string(), "primary alias".to_string()),
19412 ]))
19413 .hooks(hooks.clone())
19414 .build()
19415 .unwrap();
19416 agent.tools.set_tool_aliases(
19417 "primary",
19418 ToolAliases::new().with_name("en", "primary alias"),
19419 );
19420
19421 let record = tokio::time::timeout(
19422 std::time::Duration::from_secs(2),
19423 agent.invoke_tool(ToolExecutionRequest::new(
19424 "alias-fallback-call",
19425 "primary",
19426 serde_json::json!({"path": "./shared.txt"}),
19427 ToolCallSource::Manual,
19428 )),
19429 )
19430 .await
19431 .expect("alias-mediated fallback cycle must terminate")
19432 .unwrap();
19433
19434 assert_eq!(primary_calls.load(Ordering::SeqCst), 1);
19435 assert_eq!(secondary_calls.load(Ordering::SeqCst), 1);
19436 assert_eq!(record.requested_name, "primary alias");
19437 assert_eq!(record.canonical_id, "primary");
19438 assert!(!record.executed);
19439 assert!(record.output.contains("fallback cycle"));
19440 assert_eq!(
19441 record.metadata.get("fallback_chain"),
19442 Some(&serde_json::json!(["primary", "secondary"]))
19443 );
19444 assert_eq!(hooks.records().len(), 3);
19445 assert_eq!(agent.tool_call_history().len(), 3);
19446 }
19447
19448 #[tokio::test]
19450 async fn final_canonical_drift_cannot_bypass_fallback_ancestry() {
19451 let primary_calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19452 let secondary_calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19453 let provider = Arc::new(DriftingFallbackProvider {
19454 refreshed: AtomicBool::new(false),
19455 primary_calls: Arc::clone(&primary_calls),
19456 secondary_calls: Arc::clone(&secondary_calls),
19457 });
19458 let registry = ToolRegistry::new();
19459 registry.register_provider(provider).await.unwrap();
19460 let lifecycle = Arc::new(ToolLifecycleRecordingHooks::new());
19461 let hooks = Arc::new(RefreshFallbackProviderHooks {
19462 agent: parking_lot::Mutex::new(None),
19463 lifecycle: Arc::clone(&lifecycle),
19464 });
19465 let agent = Arc::new(
19466 AgentBuilder::new()
19467 .system_prompt("Test final canonical fallback admission.")
19468 .llm(Arc::new(mock_with_response("done")))
19469 .tools(registry)
19470 .recovery_manager(recovery_manager_with_fallbacks([(
19471 "primary".to_string(),
19472 "fallback alias".to_string(),
19473 )]))
19474 .hooks(hooks.clone())
19475 .build()
19476 .unwrap(),
19477 );
19478 *hooks.agent.lock() = Some(Arc::downgrade(&agent));
19479
19480 let record = agent
19481 .invoke_tool(ToolExecutionRequest::new(
19482 "drifting-fallback-call",
19483 "primary",
19484 serde_json::json!({"path": "./shared.txt"}),
19485 ToolCallSource::Manual,
19486 ))
19487 .await
19488 .unwrap();
19489
19490 assert_eq!(primary_calls.load(Ordering::SeqCst), 1);
19491 assert_eq!(secondary_calls.load(Ordering::SeqCst), 0);
19492 assert_eq!(record.canonical_id, "secondary");
19493 assert!(!record.executed);
19494 assert_eq!(record.policy.outcome, PermissionOutcome::Deny);
19495 assert!(record.output.contains("fallback cycle"));
19496 assert_eq!(
19497 record.metadata.get("fallback_chain"),
19498 Some(&serde_json::json!(["primary", "secondary"]))
19499 );
19500 assert_eq!(
19501 record.metadata.get("final_resolved_canonical_id"),
19502 Some(&serde_json::json!("primary"))
19503 );
19504 assert_eq!(
19505 lifecycle.events(),
19506 vec![
19507 "start:primary",
19508 "complete:primary:false",
19509 "record:primary:true",
19510 "error",
19511 "start:secondary",
19512 "complete:secondary:false",
19513 "record:secondary:false",
19514 "error",
19515 ]
19516 );
19517 let records = lifecycle.records();
19518 assert_eq!(records.len(), 2);
19519 assert_eq!(records[1].canonical_id, "secondary");
19520 assert_eq!(
19521 records[1].metadata.get("final_resolved_canonical_id"),
19522 Some(&serde_json::json!("primary"))
19523 );
19524 let history = agent.tool_call_history();
19525 assert_eq!(
19526 history
19527 .iter()
19528 .map(|entry| entry.tool_id.as_str())
19529 .collect::<Vec<_>>(),
19530 vec!["primary", "secondary"]
19531 );
19532 }
19533
19534 #[tokio::test]
19536 async fn acyclic_fallback_chain_is_denied_after_the_hop_limit() {
19537 let tool_count = MAX_TOOL_FALLBACK_HOPS + 2;
19538 let calls = (0..tool_count)
19539 .map(|_| Arc::new(std::sync::atomic::AtomicUsize::new(0)))
19540 .collect::<Vec<_>>();
19541 let mut builder = AgentBuilder::new()
19542 .system_prompt("Test bounded acyclic fallback admission.")
19543 .llm(Arc::new(mock_with_response("done")));
19544 for (index, counter) in calls.iter().enumerate() {
19545 builder = builder.tool(Arc::new(RecoveryTestTool {
19546 id: format!("fallback_{index}"),
19547 succeeds: false,
19548 calls: Arc::clone(counter),
19549 max_output_chars: None,
19550 }));
19551 }
19552 let fallbacks = (0..tool_count - 1).map(|index| {
19553 (
19554 format!("fallback_{index}"),
19555 format!("fallback_{}", index + 1),
19556 )
19557 });
19558 let agent = builder
19559 .recovery_manager(recovery_manager_with_fallbacks(fallbacks))
19560 .build()
19561 .unwrap();
19562
19563 let record = tokio::time::timeout(
19564 std::time::Duration::from_secs(2),
19565 agent.invoke_tool(ToolExecutionRequest::new(
19566 "bounded-fallback-call",
19567 "fallback_0",
19568 serde_json::json!({"path": "./shared.txt"}),
19569 ToolCallSource::Manual,
19570 )),
19571 )
19572 .await
19573 .expect("bounded fallback chain must terminate")
19574 .unwrap();
19575
19576 for counter in calls.iter().take(MAX_TOOL_FALLBACK_HOPS + 1) {
19577 assert_eq!(counter.load(Ordering::SeqCst), 1);
19578 }
19579 assert_eq!(calls[MAX_TOOL_FALLBACK_HOPS + 1].load(Ordering::SeqCst), 0);
19580 assert_eq!(
19581 record.canonical_id,
19582 format!("fallback_{}", MAX_TOOL_FALLBACK_HOPS + 1)
19583 );
19584 assert!(!record.executed);
19585 assert!(record.output.contains("maximum of 16 hops"));
19586 assert_eq!(agent.tool_call_history().len(), tool_count);
19587 }
19588
19589 #[tokio::test]
19590 async fn diagnostics_without_provider_records_unavailable_without_execution() {
19591 let mock = mock_with_response("hello");
19592 let yaml = r#"
19593name: DiagnosticsNoProviderAgent
19594system_prompt: "Review diagnostics."
19595tools: [diagnostics]
19596"#;
19597 let agent = AgentBuilder::from_yaml(yaml)
19598 .unwrap()
19599 .llm(Arc::new(mock))
19600 .auto_configure_features()
19601 .unwrap()
19602 .build()
19603 .unwrap();
19604
19605 let record = agent
19606 .invoke_tool(ToolExecutionRequest::new(
19607 "diagnostics-call",
19608 "diagnostics",
19609 serde_json::json!({}),
19610 ToolCallSource::Manual,
19611 ))
19612 .await
19613 .unwrap();
19614
19615 assert!(!record.executed);
19616 assert!(!record.success);
19617 assert_eq!(record.policy.outcome, PermissionOutcome::Unavailable);
19618 }
19619
19620 #[tokio::test]
19621 async fn web_search_without_provider_records_unavailable_without_execution() {
19622 let mock = mock_with_response("hello");
19623 let yaml = r#"
19624name: WebSearchNoProviderAgent
19625system_prompt: "You search the web."
19626tools: [web_search]
19627"#;
19628 let agent = AgentBuilder::from_yaml(yaml)
19629 .unwrap()
19630 .llm(Arc::new(mock))
19631 .auto_configure_features()
19632 .unwrap()
19633 .build()
19634 .unwrap();
19635
19636 let record = agent
19637 .invoke_tool(ToolExecutionRequest::new(
19638 "web-search-call",
19639 "web_search",
19640 serde_json::json!({"query": "rust async"}),
19641 ToolCallSource::Manual,
19642 ))
19643 .await
19644 .unwrap();
19645
19646 assert!(!record.executed);
19647 assert!(!record.success);
19648 assert_eq!(record.policy.outcome, PermissionOutcome::Unavailable);
19649 }
19650
19651 #[tokio::test]
19652 async fn unavailable_host_tool_does_not_request_approval() {
19653 let approvals = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19654 let handler = Arc::new(CountingApprovalHandler {
19655 calls: Arc::clone(&approvals),
19656 });
19657 let mut security = ToolSecurityConfig {
19658 enabled: true,
19659 fail_closed: true,
19660 ..Default::default()
19661 };
19662 security.tools.insert(
19663 "web_search".to_string(),
19664 ai_agents_tools::ToolPolicyConfig {
19665 enabled: true,
19666 require_confirmation: true,
19667 ..Default::default()
19668 },
19669 );
19670 let yaml = r#"
19671name: UnavailableApprovalAgent
19672system_prompt: "Search only with approval."
19673tools: [web_search]
19674"#;
19675 let agent = AgentBuilder::from_yaml(yaml)
19676 .unwrap()
19677 .llm(Arc::new(mock_with_response("done")))
19678 .auto_configure_features()
19679 .unwrap()
19680 .tool_security(ToolSecurityEngine::new(security))
19681 .hitl_engine(HITLEngine::new(ai_agents_hitl::HITLConfig::default()))
19682 .approval_handler(handler)
19683 .build()
19684 .unwrap();
19685
19686 let record = agent
19687 .invoke_tool(ToolExecutionRequest::new(
19688 "unavailable-before-approval",
19689 "web_search",
19690 serde_json::json!({"query": "rust async"}),
19691 ToolCallSource::Manual,
19692 ))
19693 .await
19694 .unwrap();
19695
19696 assert_eq!(approvals.load(Ordering::SeqCst), 0);
19697 assert!(!record.executed);
19698 assert!(!record.success);
19699 assert_eq!(record.policy.outcome, PermissionOutcome::Unavailable);
19700 assert!(
19701 record
19702 .approval
19703 .as_ref()
19704 .is_some_and(|approval| matches!(approval.status, ToolApprovalStatus::Unavailable))
19705 );
19706 }
19707
19708 #[tokio::test]
19709 async fn test_spawner_section_does_not_grant_core_tools_when_top_level_tools_omitted() {
19710 let mock = mock_with_response("hello");
19711 let yaml = r#"
19712name: SpawnerNoGrantAgent
19713system_prompt: "You manage agents."
19714spawner:
19715 max_agents: 2
19716"#;
19717 let agent = AgentBuilder::from_yaml(yaml)
19718 .unwrap()
19719 .llm(Arc::new(mock))
19720 .auto_configure_features()
19721 .unwrap()
19722 .auto_configure_spawner()
19723 .await
19724 .unwrap()
19725 .build()
19726 .unwrap();
19727
19728 let available = agent.get_available_tool_ids().await.unwrap();
19729 assert!(available.is_empty());
19730 }
19731
19732 #[tokio::test]
19733 async fn test_spawner_section_does_not_grant_core_tools_when_top_level_tools_empty() {
19734 let mock = mock_with_response("hello");
19735 let yaml = r#"
19736name: EmptySpawnerNoGrantAgent
19737system_prompt: "You manage agents."
19738tools: []
19739spawner:
19740 max_agents: 2
19741"#;
19742 let agent = AgentBuilder::from_yaml(yaml)
19743 .unwrap()
19744 .llm(Arc::new(mock))
19745 .auto_configure_features()
19746 .unwrap()
19747 .auto_configure_spawner()
19748 .await
19749 .unwrap()
19750 .build()
19751 .unwrap();
19752
19753 let available = agent.get_available_tool_ids().await.unwrap();
19754 assert!(available.is_empty());
19755 }
19756
19757 #[tokio::test]
19758 async fn test_management_tools_flag_grants_core_tools_when_top_level_tools_empty() {
19759 let mock = mock_with_response("hello");
19760 let yaml = r#"
19761name: ManagementGrantAgent
19762system_prompt: "You manage agents."
19763tools: []
19764spawner:
19765 management_tools: true
19766"#;
19767 let agent = AgentBuilder::from_yaml(yaml)
19768 .unwrap()
19769 .llm(Arc::new(mock))
19770 .auto_configure_features()
19771 .unwrap()
19772 .auto_configure_spawner()
19773 .await
19774 .unwrap()
19775 .build()
19776 .unwrap();
19777
19778 let available = agent.get_available_tool_ids().await.unwrap();
19779 assert_eq!(available.len(), 4);
19780 assert!(available.contains(&"spawn_agent".to_string()));
19781 assert!(available.contains(&"send_agent_message".to_string()));
19782 assert!(available.contains(&"list_agents".to_string()));
19783 assert!(available.contains(&"remove_agent".to_string()));
19784 }
19785
19786 #[tokio::test]
19787 async fn test_management_tools_flag_grants_core_tools_when_top_level_tools_omitted() {
19788 let mock = mock_with_response("hello");
19789 let yaml = r#"
19790name: ManagementOmittedToolsGrantAgent
19791system_prompt: "You manage agents."
19792spawner:
19793 management_tools: true
19794"#;
19795 let agent = AgentBuilder::from_yaml(yaml)
19796 .unwrap()
19797 .llm(Arc::new(mock))
19798 .auto_configure_features()
19799 .unwrap()
19800 .auto_configure_spawner()
19801 .await
19802 .unwrap()
19803 .build()
19804 .unwrap();
19805
19806 let available = agent.get_available_tool_ids().await.unwrap();
19807 assert_eq!(available.len(), 4);
19808 assert!(available.contains(&"spawn_agent".to_string()));
19809 assert!(available.contains(&"send_agent_message".to_string()));
19810 assert!(available.contains(&"list_agents".to_string()));
19811 assert!(available.contains(&"remove_agent".to_string()));
19812 }
19813
19814 #[tokio::test]
19815 async fn test_management_tools_selected_grants_only_selected_tools() {
19816 let mock = mock_with_response("hello");
19817 let yaml = r#"
19818name: ManagementSelectedGrantAgent
19819system_prompt: "You manage agents."
19820tools: []
19821spawner:
19822 management_tools:
19823 - spawn_agent
19824 - send_agent_message
19825 - list_agents
19826"#;
19827 let agent = AgentBuilder::from_yaml(yaml)
19828 .unwrap()
19829 .llm(Arc::new(mock))
19830 .auto_configure_features()
19831 .unwrap()
19832 .auto_configure_spawner()
19833 .await
19834 .unwrap()
19835 .build()
19836 .unwrap();
19837
19838 let available = agent.get_available_tool_ids().await.unwrap();
19839 assert_eq!(available.len(), 3);
19840 assert!(available.contains(&"spawn_agent".to_string()));
19841 assert!(available.contains(&"send_agent_message".to_string()));
19842 assert!(available.contains(&"list_agents".to_string()));
19843 assert!(!available.contains(&"remove_agent".to_string()));
19844 }
19845
19846 #[tokio::test]
19847 async fn test_orchestration_tools_flag_grants_tools_when_top_level_tools_empty() {
19848 let mock = mock_with_response("hello");
19849 let yaml = r#"
19850name: OrchestrationGrantAgent
19851system_prompt: "You coordinate agents."
19852llms:
19853 default:
19854 provider: openai
19855 model: gpt-4
19856 router:
19857 provider: openai
19858 model: gpt-4
19859llm:
19860 default: default
19861 router: router
19862tools: []
19863spawner:
19864 orchestration_tools: true
19865"#;
19866 let agent = AgentBuilder::from_yaml(yaml)
19867 .unwrap()
19868 .llm(Arc::new(mock))
19869 .auto_configure_features()
19870 .unwrap()
19871 .auto_configure_spawner()
19872 .await
19873 .unwrap()
19874 .build()
19875 .unwrap();
19876
19877 let available = agent.get_available_tool_ids().await.unwrap();
19878 assert_eq!(available.len(), 5);
19879 assert!(available.contains(&"route_to_agent".to_string()));
19880 assert!(available.contains(&"pipeline_process".to_string()));
19881 assert!(available.contains(&"concurrent_ask".to_string()));
19882 assert!(available.contains(&"group_discussion".to_string()));
19883 assert!(available.contains(&"handoff_conversation".to_string()));
19884 }
19885
19886 #[tokio::test]
19887 async fn test_persona_evolve_flag_grants_tool_when_top_level_tools_empty() {
19888 let mock = mock_with_response("hello");
19889 let yaml = r#"
19890name: PersonaGrantAgent
19891system_prompt: "You can evolve persona."
19892llm:
19893 provider: openai
19894 model: gpt-4
19895tools: []
19896persona:
19897 identity:
19898 name: "Guide"
19899 role: "Helper"
19900 evolution:
19901 enabled: true
19902 allow_llm_evolve: true
19903 mutable_fields:
19904 - traits.personality
19905"#;
19906 let agent = AgentBuilder::from_yaml(yaml)
19907 .unwrap()
19908 .llm(Arc::new(mock))
19909 .build()
19910 .unwrap();
19911
19912 let available = agent.get_available_tool_ids().await.unwrap();
19913 assert_eq!(available, vec!["persona_evolve".to_string()]);
19914 }
19915
19916 #[tokio::test]
19917 async fn test_persona_evolve_flag_grants_tool_when_top_level_tools_omitted() {
19918 let mock = mock_with_response("hello");
19919 let yaml = r#"
19920name: PersonaOmittedToolsGrantAgent
19921system_prompt: "You can evolve persona."
19922llm:
19923 provider: openai
19924 model: gpt-4
19925persona:
19926 identity:
19927 name: "Guide"
19928 role: "Helper"
19929 evolution:
19930 enabled: true
19931 allow_llm_evolve: true
19932 mutable_fields:
19933 - traits.personality
19934"#;
19935 let agent = AgentBuilder::from_yaml(yaml)
19936 .unwrap()
19937 .llm(Arc::new(mock))
19938 .build()
19939 .unwrap();
19940
19941 let available = agent.get_available_tool_ids().await.unwrap();
19942 assert_eq!(available, vec!["persona_evolve".to_string()]);
19943 }
19944
19945 #[tokio::test]
19946 async fn test_omitted_yaml_tools_exposes_no_tools() {
19947 let mock = mock_with_response("hello");
19948 let yaml = r#"
19949name: NoToolsAgent
19950system_prompt: "You are helpful."
19951"#;
19952 let agent = AgentBuilder::from_yaml(yaml)
19953 .unwrap()
19954 .llm(Arc::new(mock))
19955 .auto_configure_features()
19956 .unwrap()
19957 .build()
19958 .unwrap();
19959
19960 let available = agent.get_available_tool_ids().await.unwrap();
19961 assert!(available.is_empty());
19962 }
19963
19964 #[tokio::test]
19965 async fn runtime_scope_cannot_widen_omitted_or_empty_yaml_grants() {
19966 for tools in ["", "tools: []"] {
19967 let yaml = format!(
19968 r#"
19969name: RuntimeScopeNoGrantAgent
19970system_prompt: "No ordinary tools are granted."
19971{tools}
19972"#
19973 );
19974 let agent = AgentBuilder::from_yaml(&yaml)
19975 .unwrap()
19976 .llm(Arc::new(mock_with_response("done")))
19977 .auto_configure_features()
19978 .unwrap()
19979 .build()
19980 .unwrap();
19981
19982 agent
19983 .runtime_control()
19984 .set_tool_scope(vec!["calculator".to_string()]);
19985
19986 assert!(agent.get_available_tool_ids().await.unwrap().is_empty());
19987 }
19988 }
19989
19990 #[tokio::test]
19991 async fn runtime_scope_widening_attempt_keeps_only_declared_tools() {
19992 let yaml = r#"
19993name: RuntimeScopeWideningAgent
19994system_prompt: "Runtime scope cannot add authority."
19995tools: [calculator]
19996"#;
19997 let agent = AgentBuilder::from_yaml(yaml)
19998 .unwrap()
19999 .llm(Arc::new(mock_with_response("done")))
20000 .auto_configure_features()
20001 .unwrap()
20002 .build()
20003 .unwrap();
20004
20005 agent
20006 .runtime_control()
20007 .set_tool_scope(vec!["calculator".to_string(), "datetime".to_string()]);
20008
20009 assert_eq!(
20010 agent.get_available_tool_ids().await.unwrap(),
20011 vec!["calculator".to_string()]
20012 );
20013 }
20014
20015 #[tokio::test]
20016 async fn runtime_scope_is_canonical_unique_ordered_and_clear_restores_declared_grant() {
20017 let yaml = r#"
20018name: RuntimeScopeIntersectionAgent
20019system_prompt: "Use only declared tools."
20020tools: [calculator, datetime]
20021"#;
20022 let agent = AgentBuilder::from_yaml(yaml)
20023 .unwrap()
20024 .llm(Arc::new(mock_with_response("done")))
20025 .auto_configure_features()
20026 .unwrap()
20027 .build()
20028 .unwrap();
20029 let mut aliases = ai_agents_tools::ToolAliases::default();
20030 aliases
20031 .names
20032 .insert("en".to_string(), "calculate_alias".to_string());
20033 agent.tools.set_tool_aliases("calculator", aliases);
20034 let control = agent.runtime_control();
20035
20036 control.set_tool_scope(vec![
20037 "datetime".to_string(),
20038 "calculate_alias".to_string(),
20039 "calculator".to_string(),
20040 "unknown".to_string(),
20041 "datetime".to_string(),
20042 ]);
20043 assert_eq!(
20044 agent.get_available_tool_ids().await.unwrap(),
20045 vec!["calculator".to_string(), "datetime".to_string()]
20046 );
20047
20048 control.set_tool_scope(vec!["datetime".to_string()]);
20049 assert_eq!(
20050 agent.get_available_tool_ids().await.unwrap(),
20051 vec!["datetime".to_string()]
20052 );
20053
20054 control.clear_tool_scope_override();
20055 assert_eq!(
20056 agent.get_available_tool_ids().await.unwrap(),
20057 vec!["calculator".to_string(), "datetime".to_string()]
20058 );
20059 }
20060
20061 #[tokio::test]
20062 async fn runtime_scope_preserves_programmatic_registration_as_declared_grant() {
20063 let agent = AgentBuilder::new()
20064 .system_prompt("Use registered tools.")
20065 .llm(Arc::new(mock_with_response("done")))
20066 .tool(Arc::new(ContextEchoTool))
20067 .tool(Arc::new(SlowTool))
20068 .build()
20069 .unwrap();
20070
20071 agent.runtime_control().set_tool_scope(vec![
20072 "Context Echo".to_string(),
20073 "context_echo".to_string(),
20074 "unknown".to_string(),
20075 ]);
20076
20077 assert_eq!(
20078 agent.get_available_tool_ids().await.unwrap(),
20079 vec!["context_echo".to_string()]
20080 );
20081 }
20082
20083 #[tokio::test]
20084 async fn nested_state_scopes_intersect_every_ancestor_with_aliases() {
20085 let yaml = r#"
20086name: NestedStateScopeAgent
20087system_prompt: "Honor every state scope."
20088tools: [calculator, datetime, echo]
20089states:
20090 initial: root
20091 states:
20092 root:
20093 tools: [calculate_alias, datetime]
20094 initial: middle
20095 states:
20096 middle:
20097 initial: leaf
20098 states:
20099 leaf:
20100 tools: [datetime_alias, echo]
20101"#;
20102 let agent = AgentBuilder::from_yaml(yaml)
20103 .unwrap()
20104 .llm(Arc::new(mock_with_response("done")))
20105 .auto_configure_features()
20106 .unwrap()
20107 .build()
20108 .unwrap();
20109 let mut calculator_aliases = ai_agents_tools::ToolAliases::default();
20110 calculator_aliases
20111 .names
20112 .insert("en".to_string(), "calculate_alias".to_string());
20113 agent
20114 .tools
20115 .set_tool_aliases("calculator", calculator_aliases);
20116 let mut datetime_aliases = ai_agents_tools::ToolAliases::default();
20117 datetime_aliases
20118 .names
20119 .insert("en".to_string(), "datetime_alias".to_string());
20120 agent.tools.set_tool_aliases("datetime", datetime_aliases);
20121 agent.runtime_control().set_tool_scope(vec![
20122 "unknown".to_string(),
20123 "datetime_alias".to_string(),
20124 "calculate_alias".to_string(),
20125 "datetime".to_string(),
20126 ]);
20127
20128 assert_eq!(agent.current_state().as_deref(), Some("root.middle.leaf"));
20129 assert_eq!(
20130 agent.get_available_tool_ids().await.unwrap(),
20131 vec!["datetime".to_string()]
20132 );
20133 }
20134
20135 #[tokio::test]
20136 async fn ancestor_empty_state_scope_denies_omitted_descendants() {
20137 let yaml = r#"
20138name: NestedEmptyStateScopeAgent
20139system_prompt: "An empty ancestor scope denies all tools."
20140tools: [calculator]
20141states:
20142 initial: root
20143 states:
20144 root:
20145 tools: []
20146 initial: middle
20147 states:
20148 middle:
20149 initial: leaf
20150 states:
20151 leaf: {}
20152"#;
20153 let agent = AgentBuilder::from_yaml(yaml)
20154 .unwrap()
20155 .llm(Arc::new(mock_with_response("done")))
20156 .auto_configure_features()
20157 .unwrap()
20158 .build()
20159 .unwrap();
20160
20161 assert!(agent.get_available_tool_ids().await.unwrap().is_empty());
20162 }
20163
20164 #[tokio::test]
20165 async fn state_change_during_approval_invalidates_the_reviewed_authority() {
20166 let active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
20167 let max_active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
20168 let entered = Arc::new(tokio::sync::Barrier::new(2));
20169 let release = Arc::new(tokio::sync::Notify::new());
20170 let handler = Arc::new(BlockingApprovalHandler {
20171 entered: Arc::clone(&entered),
20172 release: Arc::clone(&release),
20173 result: ApprovalResult::Approved,
20174 });
20175 let yaml = r#"
20176name: ApprovalStateGenerationAgent
20177system_prompt: "State authority may change during approval."
20178tools: [locked_write]
20179states:
20180 initial: first
20181 states:
20182 first:
20183 tools: [locked_write]
20184 second:
20185 tools: [locked_write]
20186"#;
20187 let agent = Arc::new(
20188 AgentBuilder::from_yaml(yaml)
20189 .unwrap()
20190 .llm(Arc::new(mock_with_response("done")))
20191 .tool(Arc::new(LockedWriteTool {
20192 active: Arc::clone(&active),
20193 max_active: Arc::clone(&max_active),
20194 }))
20195 .tool_security(ToolSecurityEngine::new(approval_security_config(true)))
20196 .hitl_engine(HITLEngine::new(ai_agents_hitl::HITLConfig::default()))
20197 .approval_handler(handler)
20198 .build()
20199 .unwrap(),
20200 );
20201 let running = Arc::clone(&agent);
20202 let call = tokio::spawn(async move {
20203 running
20204 .invoke_tool(ToolExecutionRequest::new(
20205 "approval-state-generation",
20206 "locked_write",
20207 serde_json::json!({"path": "./state-generation.txt"}),
20208 ToolCallSource::Manual,
20209 ))
20210 .await
20211 .unwrap()
20212 });
20213
20214 entered.wait().await;
20215 agent.transition_to("second").await.unwrap();
20216 release.notify_one();
20217 let record = call.await.unwrap();
20218
20219 assert!(!record.executed);
20220 assert!(record.output.contains("Approval became stale"));
20221 assert_eq!(max_active.load(Ordering::SeqCst), 0);
20222 }
20223
20224 #[tokio::test]
20225 async fn state_change_while_waiting_for_resource_lock_fails_final_admission() {
20226 let holder_gate = PathMutationGate::new();
20227 let waiter_gate = PathMutationGate::new();
20228 let yaml = r#"
20229name: LockedStateGenerationAgent
20230system_prompt: "State authority must remain stable through admission."
20231tools: [state_lock_holder, state_lock_waiter]
20232states:
20233 initial: first
20234 states:
20235 first:
20236 tools: [state_lock_holder, state_lock_waiter]
20237 second:
20238 tools: [state_lock_holder, state_lock_waiter]
20239"#;
20240 let agent = Arc::new(
20241 AgentBuilder::from_yaml(yaml)
20242 .unwrap()
20243 .llm(Arc::new(mock_with_response("done")))
20244 .tool(Arc::new(BlockingPathMutationTool {
20245 id: "state_lock_holder",
20246 path_fields: vec![ai_agents_core::PathPolicyBinding::write("path")],
20247 gate: holder_gate.clone(),
20248 }))
20249 .tool(Arc::new(BlockingPathMutationTool {
20250 id: "state_lock_waiter",
20251 path_fields: vec![ai_agents_core::PathPolicyBinding::write("path")],
20252 gate: waiter_gate.clone(),
20253 }))
20254 .build()
20255 .unwrap(),
20256 );
20257 let holder_call = {
20258 let agent = Arc::clone(&agent);
20259 tokio::spawn(async move {
20260 agent
20261 .invoke_tool(ToolExecutionRequest::new(
20262 "state-lock-holder",
20263 "state_lock_holder",
20264 serde_json::json!({"path": "./shared-state-path.txt"}),
20265 ToolCallSource::Manual,
20266 ))
20267 .await
20268 .unwrap()
20269 })
20270 };
20271 holder_gate.wait_until_entered().await;
20272 let waiter_call = {
20273 let agent = Arc::clone(&agent);
20274 tokio::spawn(async move {
20275 agent
20276 .invoke_tool(ToolExecutionRequest::new(
20277 "state-lock-waiter",
20278 "state_lock_waiter",
20279 serde_json::json!({"path": "./shared-state-path.txt"}),
20280 ToolCallSource::Manual,
20281 ))
20282 .await
20283 .unwrap()
20284 })
20285 };
20286
20287 wait_for_resource_lock_strong_count(&agent.resource_locks, 2).await;
20288 agent.transition_to("second").await.unwrap();
20289 holder_gate.release();
20290 let holder_record = holder_call.await.unwrap();
20291 let waiter_record = waiter_call.await.unwrap();
20292
20293 assert!(holder_record.success);
20294 assert!(!waiter_record.executed);
20295 assert!(
20296 waiter_record
20297 .output
20298 .contains("state scope changed before admission")
20299 );
20300 assert!(!waiter_gate.entered.load(Ordering::SeqCst));
20301 }
20302
20303 #[tokio::test]
20304 async fn test_state_tools_cannot_widen_top_level_grant() {
20305 let mock = mock_with_response("hello");
20306 let yaml = r#"
20307name: NarrowToolsAgent
20308system_prompt: "You are helpful."
20309tools:
20310 - calculator
20311states:
20312 initial: current
20313 states:
20314 current:
20315 tools: [datetime]
20316"#;
20317 let agent = AgentBuilder::from_yaml(yaml)
20318 .unwrap()
20319 .llm(Arc::new(mock))
20320 .auto_configure_features()
20321 .unwrap()
20322 .build()
20323 .unwrap();
20324
20325 let available = agent.get_available_tool_ids().await.unwrap();
20326 assert!(available.is_empty());
20327 }
20328
20329 #[tokio::test]
20331 async fn test_integration_tool_execution() {
20332 let mock = mock_with_responses(vec![
20334 r#"I'll calculate that for you.
20336[TOOL_CALL: {"name": "calculator", "arguments": {"expression": "2+2"}}]"#,
20337 "The answer is 4.",
20339 ]);
20340 let mut tools = ai_agents_tools::ToolRegistry::new();
20341 tools
20342 .register(Arc::new(ai_agents_tools::CalculatorTool))
20343 .unwrap();
20344
20345 let agent = AgentBuilder::new()
20346 .system_prompt("You are a calculator assistant.")
20347 .llm(Arc::new(mock))
20348 .tools(tools)
20349 .build()
20350 .unwrap();
20351
20352 let response = agent.chat("What is 2+2?").await.unwrap();
20353 assert!(!response.content.is_empty());
20355 }
20356
20357 #[tokio::test]
20358 async fn test_tool_hitl_rejection_finalizes_blocking_turn() {
20359 let responses = Arc::new(std::sync::atomic::AtomicUsize::new(0));
20360 let hooks = Arc::new(ResponseCountingHooks {
20361 responses: Arc::clone(&responses),
20362 });
20363 let mock = mock_with_response(r#"{"tool":"echo","arguments":{"message":"hello"}}"#);
20364 let yaml = r#"
20365name: ToolRejectAgent
20366system_prompt: "You use tools when requested."
20367tools:
20368 - echo
20369hitl:
20370 tools:
20371 echo:
20372 require_approval: true
20373 approval_message: "Approve echo?"
20374"#;
20375 let agent = AgentBuilder::from_yaml(yaml)
20376 .unwrap()
20377 .llm(Arc::new(mock))
20378 .auto_configure_features()
20379 .unwrap()
20380 .hooks(hooks)
20381 .build()
20382 .unwrap();
20383
20384 let response = agent.chat("echo hello").await.unwrap();
20385
20386 assert!(
20387 response.content.contains("Operation cancelled"),
20388 "unexpected response: {}",
20389 response.content
20390 );
20391 assert_eq!(responses.load(Ordering::SeqCst), 1);
20392 let messages = agent.memory.get_messages(None).await.unwrap();
20393 assert_eq!(messages.len(), 3);
20394 assert_eq!(messages[0].content, "echo hello");
20395 assert!(messages[1].content.contains("\"tool\":\"echo\""));
20396 assert!(messages[2].content.contains("rejected by the approver"));
20397 }
20398
20399 #[tokio::test]
20400 async fn test_tool_hitl_rejection_finalizes_streaming_turn() {
20401 use futures::StreamExt;
20402
20403 let responses = Arc::new(std::sync::atomic::AtomicUsize::new(0));
20404 let hooks = Arc::new(ResponseCountingHooks {
20405 responses: Arc::clone(&responses),
20406 });
20407 let mock = mock_with_response(r#"{"tool":"echo","arguments":{"message":"hello"}}"#);
20408 let yaml = r#"
20409name: ToolRejectStreamingAgent
20410system_prompt: "You use tools when requested."
20411tools:
20412 - echo
20413streaming:
20414 enabled: true
20415hitl:
20416 tools:
20417 echo:
20418 require_approval: true
20419 approval_message: "Approve echo?"
20420"#;
20421 let agent = AgentBuilder::from_yaml(yaml)
20422 .unwrap()
20423 .llm(Arc::new(mock))
20424 .auto_configure_features()
20425 .unwrap()
20426 .hooks(hooks)
20427 .build()
20428 .unwrap();
20429
20430 let mut stream = agent.chat_stream("echo hello").await.unwrap();
20431 let mut terminal_error = String::new();
20432 let mut done = false;
20433 while let Some(chunk) = stream.next().await {
20434 match chunk {
20435 StreamChunk::Error { message } => terminal_error = message,
20436 StreamChunk::Done {} => {
20437 done = true;
20438 break;
20439 }
20440 _ => {}
20441 }
20442 }
20443
20444 assert!(done);
20445 assert!(
20446 terminal_error.contains("Operation cancelled"),
20447 "unexpected terminal error: {}",
20448 terminal_error
20449 );
20450 assert_eq!(responses.load(Ordering::SeqCst), 1);
20451 let messages = agent.memory.get_messages(None).await.unwrap();
20452 assert_eq!(messages.len(), 3);
20453 assert_eq!(messages[0].content, "echo hello");
20454 assert!(messages[1].content.contains("\"tool\":\"echo\""));
20455 assert!(messages[2].content.contains("rejected by the approver"));
20456 }
20457
20458 #[tokio::test]
20459 async fn tool_hitl_rejection_preserves_legacy_error_but_finalizes_event_stream() {
20460 let mock = mock_with_response(r#"{"tool":"echo","arguments":{"message":"hello"}}"#);
20461 let yaml = r#"
20462name: ToolRejectEventAgent
20463system_prompt: "You use tools when requested."
20464tools:
20465 - echo
20466streaming:
20467 enabled: true
20468hitl:
20469 tools:
20470 echo:
20471 require_approval: true
20472 approval_message: "Approve echo?"
20473"#;
20474 let agent = AgentBuilder::from_yaml(yaml)
20475 .unwrap()
20476 .llm(Arc::new(mock))
20477 .auto_configure_features()
20478 .unwrap()
20479 .build()
20480 .unwrap();
20481
20482 let mut stream = agent.chat_stream_events("echo hello").await.unwrap();
20483 let mut error_seen = false;
20484 let mut final_response = None;
20485 while let Some(event) = stream.next().await {
20486 match event {
20487 AgentStreamEvent::Chunk(StreamChunk::Error { .. }) => error_seen = true,
20488 AgentStreamEvent::Final(response) => final_response = Some(response),
20489 AgentStreamEvent::Chunk(_) => {}
20490 }
20491 }
20492
20493 assert!(!error_seen);
20494 assert!(
20495 final_response
20496 .is_some_and(|response| { response.content.contains("Operation cancelled") })
20497 );
20498 }
20499
20500 #[tokio::test]
20501 async fn test_pre_response_guard_transition_skips_old_state_llm() {
20502 let mock = mock_with_response("Billing state response");
20503 let call_counter = mock.clone();
20504 let yaml = r#"
20505name: OptimizedStateAgent
20506system_prompt: "You route before answering."
20507runtime:
20508 optimization:
20509 enabled: true
20510 pre_response_deterministic_transitions: true
20511states:
20512 initial: greeting
20513 states:
20514 greeting:
20515 prompt: "Old state prompt that should be skipped."
20516 transitions:
20517 - to: billing
20518 guard:
20519 context:
20520 topic:
20521 eq: billing
20522 timing: pre_response
20523 billing:
20524 prompt: "Answer from the billing state."
20525"#;
20526 let agent = AgentBuilder::from_yaml(yaml)
20527 .unwrap()
20528 .llm(Arc::new(mock))
20529 .build()
20530 .unwrap();
20531 agent
20532 .set_context("topic", serde_json::json!("billing"))
20533 .unwrap();
20534
20535 let response = agent.chat("I need billing help").await.unwrap();
20536
20537 assert_eq!(agent.current_state().as_deref(), Some("billing"));
20538 assert_eq!(response.content, "Billing state response");
20539 assert_eq!(call_counter.call_count(), 1);
20540 assert_eq!(agent.actor_facts().len(), 0);
20541 }
20542
20543 #[tokio::test]
20544 async fn test_set_context_supports_dotted_paths_for_pre_response_guards() {
20545 let mock = mock_with_response("Billing state response");
20546 let call_counter = mock.clone();
20547 let yaml = r#"
20548name: OptimizedStateAgent
20549system_prompt: "You route before answering."
20550runtime:
20551 optimization:
20552 enabled: true
20553 pre_response_deterministic_transitions: true
20554context:
20555 request:
20556 type: runtime
20557 default:
20558 topic: general
20559states:
20560 initial: greeting
20561 states:
20562 greeting:
20563 prompt: "Old state prompt that should be skipped."
20564 transitions:
20565 - to: billing
20566 guard:
20567 context:
20568 request.topic:
20569 eq: billing
20570 timing: pre_response
20571 billing:
20572 prompt: "Answer from the billing state."
20573"#;
20574 let agent = AgentBuilder::from_yaml(yaml)
20575 .unwrap()
20576 .llm(Arc::new(mock))
20577 .build()
20578 .unwrap();
20579 agent
20580 .set_context("request.topic", serde_json::json!("billing"))
20581 .unwrap();
20582
20583 let response = agent.chat("I need billing help").await.unwrap();
20584
20585 assert_eq!(agent.current_state().as_deref(), Some("billing"));
20586 assert_eq!(response.content, "Billing state response");
20587 assert_eq!(call_counter.call_count(), 1);
20588 assert_eq!(
20589 agent.get_context().get("request"),
20590 Some(&serde_json::json!({"topic": "billing"}))
20591 );
20592 }
20593
20594 #[tokio::test]
20595 async fn test_pre_response_rejection_does_not_commit_staged_context_or_user() {
20596 let mock = mock_with_response("billing");
20597 let yaml = r#"
20598name: OptimizedStateAgent
20599system_prompt: "You route before answering."
20600runtime:
20601 optimization:
20602 enabled: true
20603 pre_response_deterministic_transitions: true
20604hitl:
20605 states:
20606 billing:
20607 on_enter: require_approval
20608 approval_message: "Approve billing route?"
20609states:
20610 initial: greeting
20611 states:
20612 greeting:
20613 prompt: "Old state prompt."
20614 extract:
20615 - key: topic
20616 description: "Support topic"
20617 transitions:
20618 - to: billing
20619 guard:
20620 context:
20621 topic:
20622 eq: billing
20623 timing: pre_response
20624 run_extractors: true
20625 billing:
20626 prompt: "Billing state."
20627"#;
20628 let agent = AgentBuilder::from_yaml(yaml)
20629 .unwrap()
20630 .llm(Arc::new(mock))
20631 .build()
20632 .unwrap();
20633
20634 let response = agent
20635 .try_pre_response_transition("billing please")
20636 .await
20637 .unwrap();
20638
20639 assert!(response.is_none());
20640 assert_eq!(agent.current_state().as_deref(), Some("greeting"));
20641 assert!(!agent.get_context().contains_key("topic"));
20642 assert_eq!(agent.memory.get_messages(None).await.unwrap().len(), 0);
20643 }
20644
20645 #[tokio::test]
20646 async fn test_pre_response_extractor_commits_context_on_winning_path() {
20647 let mock = mock_with_responses(vec!["billing", "Billing response"]);
20648 let yaml = r#"
20649name: OptimizedStateAgent
20650system_prompt: "You route before answering."
20651runtime:
20652 optimization:
20653 enabled: true
20654 pre_response_deterministic_transitions: true
20655states:
20656 initial: greeting
20657 states:
20658 greeting:
20659 prompt: "Old state prompt."
20660 extract:
20661 - key: topic
20662 description: "Support topic"
20663 transitions:
20664 - to: billing
20665 guard:
20666 context:
20667 topic:
20668 eq: billing
20669 timing: pre_response
20670 run_extractors: true
20671 billing:
20672 prompt: "Billing state."
20673"#;
20674 let agent = AgentBuilder::from_yaml(yaml)
20675 .unwrap()
20676 .llm(Arc::new(mock))
20677 .build()
20678 .unwrap();
20679
20680 let response = agent.chat("billing please").await.unwrap();
20681
20682 assert_eq!(agent.current_state().as_deref(), Some("billing"));
20683 assert_eq!(response.content, "Billing response");
20684 assert_eq!(
20685 agent.get_context().get("topic"),
20686 Some(&serde_json::json!("billing"))
20687 );
20688 }
20689
20690 #[tokio::test]
20691 async fn test_pre_response_extractor_miss_does_not_mutate_context() {
20692 let mock = mock_with_response("__NONE__");
20693 let yaml = r#"
20694name: OptimizedStateAgent
20695system_prompt: "You route before answering."
20696runtime:
20697 optimization:
20698 enabled: true
20699 pre_response_deterministic_transitions: true
20700states:
20701 initial: greeting
20702 states:
20703 greeting:
20704 prompt: "Old state prompt."
20705 extract:
20706 - key: topic
20707 description: "Support topic"
20708 transitions:
20709 - to: billing
20710 guard:
20711 context:
20712 topic:
20713 eq: billing
20714 timing: pre_response
20715 run_extractors: true
20716 billing:
20717 prompt: "Billing state."
20718"#;
20719 let agent = AgentBuilder::from_yaml(yaml)
20720 .unwrap()
20721 .llm(Arc::new(mock))
20722 .build()
20723 .unwrap();
20724
20725 let response = agent.try_pre_response_transition("hello").await.unwrap();
20726
20727 assert!(response.is_none());
20728 assert_eq!(agent.current_state().as_deref(), Some("greeting"));
20729 assert!(!agent.get_context().contains_key("topic"));
20730 }
20731
20732 #[tokio::test]
20733 async fn test_default_guard_transition_stays_post_response() {
20734 let mock = mock_with_responses(vec!["Greeting response", "Billing response"]);
20735 let call_counter = mock.clone();
20736 let yaml = r#"
20737name: TimingAgent
20738system_prompt: "You route carefully."
20739runtime:
20740 optimization:
20741 enabled: true
20742 pre_response_deterministic_transitions: true
20743states:
20744 initial: greeting
20745 states:
20746 greeting:
20747 prompt: "Old state prompt."
20748 transitions:
20749 - to: billing
20750 guard:
20751 context:
20752 topic:
20753 eq: billing
20754 billing:
20755 prompt: "Billing state."
20756"#;
20757 let agent = AgentBuilder::from_yaml(yaml)
20758 .unwrap()
20759 .llm(Arc::new(mock))
20760 .build()
20761 .unwrap();
20762 agent
20763 .set_context("topic", serde_json::json!("billing"))
20764 .unwrap();
20765
20766 let response = agent.chat("billing please").await.unwrap();
20767
20768 assert_eq!(agent.current_state().as_deref(), Some("billing"));
20769 assert_eq!(response.content, "Billing response");
20770 assert_eq!(call_counter.call_count(), 2);
20771 }
20772
20773 #[tokio::test]
20774 async fn test_explicit_post_response_guard_transition_stays_post_response() {
20775 let mock = mock_with_responses(vec!["Greeting response", "Billing response"]);
20776 let call_counter = mock.clone();
20777 let yaml = r#"
20778name: TimingAgent
20779system_prompt: "You route carefully."
20780runtime:
20781 optimization:
20782 enabled: true
20783 pre_response_deterministic_transitions: true
20784states:
20785 initial: greeting
20786 states:
20787 greeting:
20788 prompt: "Old state prompt."
20789 transitions:
20790 - to: billing
20791 guard:
20792 context:
20793 topic:
20794 eq: billing
20795 timing: post_response
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 agent
20805 .set_context("topic", serde_json::json!("billing"))
20806 .unwrap();
20807
20808 let response = agent.chat("billing please").await.unwrap();
20809
20810 assert_eq!(agent.current_state().as_deref(), Some("billing"));
20811 assert_eq!(response.content, "Billing response");
20812 assert_eq!(call_counter.call_count(), 2);
20813 }
20814
20815 #[tokio::test]
20816 async fn test_pre_response_extractors_are_transition_scoped() {
20817 let mock = mock_with_responses(vec!["billing", "Billing response"]);
20818 let yaml = r#"
20819name: ScopedExtractorAgent
20820system_prompt: "You route carefully."
20821runtime:
20822 optimization:
20823 enabled: true
20824 pre_response_deterministic_transitions: true
20825states:
20826 initial: greeting
20827 states:
20828 greeting:
20829 prompt: "Old state prompt."
20830 extract:
20831 - key: topic
20832 description: "Support topic"
20833 transitions:
20834 - to: wrong
20835 guard:
20836 context:
20837 topic:
20838 eq: billing
20839 timing: pre_response
20840 - to: billing
20841 guard:
20842 context:
20843 topic:
20844 eq: billing
20845 timing: pre_response
20846 run_extractors: true
20847 wrong:
20848 prompt: "Wrong state."
20849 billing:
20850 prompt: "Billing state."
20851"#;
20852 let agent = AgentBuilder::from_yaml(yaml)
20853 .unwrap()
20854 .llm(Arc::new(mock))
20855 .build()
20856 .unwrap();
20857
20858 let response = agent.chat("billing please").await.unwrap();
20859
20860 assert_eq!(agent.current_state().as_deref(), Some("billing"));
20861 assert_eq!(response.content, "Billing response");
20862 }
20863
20864 #[tokio::test]
20865 async fn test_pre_response_resolved_intent_routes_early() {
20866 let mock = mock_with_response("Billing response");
20867 let yaml = r#"
20868name: IntentAgent
20869system_prompt: "You route carefully."
20870runtime:
20871 optimization:
20872 enabled: true
20873 pre_response_deterministic_transitions: true
20874states:
20875 initial: greeting
20876 states:
20877 greeting:
20878 prompt: "Old state prompt."
20879 transitions:
20880 - to: billing
20881 intent: billing
20882 timing: pre_response
20883 billing:
20884 prompt: "Billing state."
20885"#;
20886 let agent = AgentBuilder::from_yaml(yaml)
20887 .unwrap()
20888 .llm(Arc::new(mock))
20889 .build()
20890 .unwrap();
20891 agent
20892 .set_context("resolved_intent", serde_json::json!("billing"))
20893 .unwrap();
20894
20895 let response = agent
20896 .try_pre_response_transition("I need billing help")
20897 .await
20898 .unwrap()
20899 .unwrap();
20900
20901 assert_eq!(agent.current_state().as_deref(), Some("billing"));
20902 assert_eq!(response.content, "Billing response");
20903 }
20904
20905 #[tokio::test]
20906 async fn test_background_overflow_error_surfaces() {
20907 let mut config = RuntimeConfig::default();
20908 config.optimization.enabled = true;
20909 config.optimization.post_turn.max_background_tasks = 1;
20910 config.optimization.post_turn.on_background_overflow = BackgroundOverflowPolicy::Error;
20911 let policy = crate::optimization::MaintenanceTaskPolicy {
20912 mode: MaintenanceMode::Background,
20913 await_before_next_turn: AwaitBeforeNextTurn::Always,
20914 };
20915 let agent = AgentBuilder::new()
20916 .system_prompt("You are helpful.")
20917 .llm(Arc::new(mock_with_response("ok")))
20918 .build()
20919 .unwrap()
20920 .with_runtime_config(config);
20921 agent
20922 .background_maintenance
20923 .spawn(None, async { std::future::pending::<Result<()>>().await })
20924 .unwrap();
20925
20926 let result = agent
20927 .spawn_or_handle_background(None, async { Ok(()) }, "facts", &policy)
20928 .await;
20929
20930 assert!(result.is_err());
20931 }
20932
20933 #[tokio::test]
20934 async fn test_speculative_reasoning_low_cap_uses_serial_reasoning() {
20935 let default_mock = mock_with_response("Plain draft response");
20936 let router_mock = mock_with_response("cot");
20937 let router_counter = router_mock.clone();
20938 let yaml = r#"
20939name: ReasoningReservationAgent
20940system_prompt: "You answer plainly unless reasoning wins."
20941llm:
20942 default: default
20943 router: router
20944observability:
20945 enabled: true
20946 export:
20947 write_raw_events: true
20948reasoning:
20949 mode: auto
20950 judge_llm: router
20951runtime:
20952 optimization:
20953 enabled: true
20954 max_speculative_llm_calls_per_turn: 1
20955 speculative_reasoning_auto: true
20956 max_parallel_runtime_tasks: 2
20957"#;
20958 let agent = AgentBuilder::from_yaml(yaml)
20959 .unwrap()
20960 .llm_alias("default", Arc::new(default_mock))
20961 .llm_alias("router", Arc::new(router_mock))
20962 .build()
20963 .unwrap();
20964
20965 let response = agent.chat("hello").await.unwrap();
20966
20967 assert_eq!(response.content, "Plain draft response");
20968 assert_eq!(router_counter.call_count(), 1);
20969 let events = agent.observability().unwrap().raw_events();
20970 assert!(!events.iter().any(|event| {
20971 event.dimensions.get("commit_behavior") == Some(&"reasoning_decision".to_string())
20972 }));
20973 }
20974
20975 #[tokio::test]
20976 async fn test_forced_reasoning_skips_plain_speculative_draft() {
20977 let mock = mock_with_response("Reasoned response");
20978 let yaml = r#"
20979name: ForcedReasoningAgent
20980system_prompt: "You reason before answering."
20981observability:
20982 enabled: true
20983 export:
20984 write_raw_events: true
20985reasoning:
20986 mode: cot
20987runtime:
20988 optimization:
20989 enabled: true
20990 max_speculative_llm_calls_per_turn: 2
20991 speculative_state_transitions: true
20992 max_parallel_runtime_tasks: 2
20993states:
20994 initial: triage
20995 states:
20996 triage:
20997 prompt: "Answer from triage."
20998 transitions:
20999 - to: billing
21000 guard:
21001 context:
21002 route:
21003 eq: billing
21004 timing: parallel
21005 billing:
21006 prompt: "Billing state."
21007"#;
21008 let agent = AgentBuilder::from_yaml(yaml)
21009 .unwrap()
21010 .llm(Arc::new(mock))
21011 .build()
21012 .unwrap();
21013
21014 let response = agent.chat("hello").await.unwrap();
21015
21016 assert_eq!(response.content, "Reasoned response");
21017 let events = agent.observability().unwrap().raw_events();
21018 assert!(
21019 !events
21020 .iter()
21021 .any(|event| event.dimensions.contains_key("branch_status"))
21022 );
21023 }
21024
21025 #[tokio::test]
21026 async fn test_speculative_skill_low_cap_uses_serial_skill_route() {
21027 let default_mock = mock_with_response("Skill committed response");
21028 let router_mock = mock_with_response("helper");
21029 let router_counter = router_mock.clone();
21030 let yaml = r#"
21031name: SkillReservationAgent
21032system_prompt: "Use skills when they match."
21033llm:
21034 default: default
21035 router: router
21036observability:
21037 enabled: true
21038 export:
21039 write_raw_events: true
21040runtime:
21041 optimization:
21042 enabled: true
21043 max_speculative_llm_calls_per_turn: 1
21044 speculative_skill_routing: true
21045 max_parallel_runtime_tasks: 2
21046skills:
21047 - id: helper
21048 description: "Answer helper requests"
21049 trigger: "User asks for helper"
21050 steps:
21051 - prompt: "Answer the helper request: {{ user_input }}"
21052"#;
21053 let agent = AgentBuilder::from_yaml(yaml)
21054 .unwrap()
21055 .llm_alias("default", Arc::new(default_mock))
21056 .llm_alias("router", Arc::new(router_mock))
21057 .build()
21058 .unwrap();
21059
21060 let response = agent.chat("please use helper").await.unwrap();
21061
21062 assert_eq!(response.content, "Skill committed response");
21063 assert_eq!(router_counter.call_count(), 1);
21064 let events = agent.observability().unwrap().raw_events();
21065 assert!(
21066 !events
21067 .iter()
21068 .any(|event| event.dimensions.contains_key("branch_status"))
21069 );
21070 }
21071
21072 #[tokio::test]
21073 async fn test_parallel_transition_low_cap_allows_deterministic_route() {
21074 let mock = mock_with_response("unused");
21075 let call_counter = mock.clone();
21076 let yaml = r#"
21077name: ParallelTransitionLowCapAgent
21078system_prompt: "Route before stale responses when safe."
21079runtime:
21080 optimization:
21081 enabled: true
21082 max_speculative_llm_calls_per_turn: 1
21083 speculative_state_transitions: true
21084 max_parallel_runtime_tasks: 2
21085states:
21086 initial: triage
21087 states:
21088 triage:
21089 prompt: "Triage state."
21090 transitions:
21091 - to: billing
21092 guard:
21093 context:
21094 route:
21095 eq: billing
21096 timing: parallel
21097 billing:
21098 prompt: "Billing state."
21099"#;
21100 let agent = AgentBuilder::from_yaml(yaml)
21101 .unwrap()
21102 .llm(Arc::new(mock))
21103 .build()
21104 .unwrap();
21105 agent
21106 .set_context("route", serde_json::json!("billing"))
21107 .unwrap();
21108 agent.update_active_turn_context("billing help", HashMap::new());
21109 assert!(
21110 agent.reserve_active_speculative_llm_call(
21111 RuntimeOptimizationKind::ParallelStateTransition
21112 )
21113 );
21114
21115 let selection = agent
21116 .select_parallel_transition_candidate("billing help")
21117 .await
21118 .unwrap();
21119 agent.end_root_turn();
21120
21121 match selection {
21122 ParallelTransitionSelection::Candidate(candidate) => {
21123 assert_eq!(candidate.target(), "billing");
21124 }
21125 ParallelTransitionSelection::NoMatch => panic!("deterministic route did not match"),
21126 ParallelTransitionSelection::ReservationExhausted => {
21127 panic!("deterministic route consumed LLM budget")
21128 }
21129 }
21130 assert_eq!(call_counter.call_count(), 0);
21131 }
21132
21133 #[tokio::test]
21134 async fn speculative_transition_drops_loser_before_state_actions() {
21135 let lock = Arc::new(tokio::sync::Mutex::new(()));
21136 let first_started = Arc::new(tokio::sync::Notify::new());
21137 let first_dropped = Arc::new(AtomicBool::new(false));
21138 let committed_after_drop = Arc::new(AtomicBool::new(false));
21139 let default = Arc::new(FirstCallLockingProvider {
21140 lock,
21141 first_started: Arc::clone(&first_started),
21142 first_dropped: Arc::clone(&first_dropped),
21143 committed_after_drop: Arc::clone(&committed_after_drop),
21144 calls: AtomicU64::new(0),
21145 });
21146 let router = Arc::new(RoutingAfterProviderStart {
21147 provider_started: first_started,
21148 });
21149 let yaml = r#"
21150name: SpeculativeCancellationAgent
21151system_prompt: "Route before committed work."
21152llm:
21153 default: default
21154 router: router
21155runtime:
21156 optimization:
21157 enabled: true
21158 max_speculative_llm_calls_per_turn: 2
21159 speculative_state_transitions: true
21160 max_parallel_runtime_tasks: 2
21161states:
21162 initial: triage
21163 states:
21164 triage:
21165 prompt: "Triage state."
21166 transitions:
21167 - to: technical
21168 when: "The request needs technical support"
21169 timing: parallel
21170 technical:
21171 prompt: "Technical state."
21172 on_enter:
21173 - prompt: "Prepare technical context."
21174 llm: default
21175 store_as: preparation
21176"#;
21177 let agent = AgentBuilder::from_yaml(yaml)
21178 .unwrap()
21179 .llm_alias("default", default)
21180 .llm_alias("router", router)
21181 .build()
21182 .unwrap();
21183
21184 let response = tokio::time::timeout(
21185 std::time::Duration::from_secs(2),
21186 agent.chat("I cannot log in because of AUTH-17."),
21187 )
21188 .await
21189 .expect("committed work must not wait on the losing provider future")
21190 .unwrap();
21191
21192 assert_eq!(response.content, "Committed technical response.");
21193 assert_eq!(agent.current_state().as_deref(), Some("technical"));
21194 assert!(first_dropped.load(Ordering::SeqCst));
21195 assert!(committed_after_drop.load(Ordering::SeqCst));
21196 }
21197
21198 #[tokio::test]
21199 async fn buffered_transition_drops_stale_stream_before_redispatch() {
21200 use futures::StreamExt;
21201
21202 let lock = Arc::new(tokio::sync::Mutex::new(()));
21203 let stream_started = Arc::new(tokio::sync::Notify::new());
21204 let stream_dropped = Arc::new(AtomicBool::new(false));
21205 let committed_after_drop = Arc::new(AtomicBool::new(false));
21206 let default = Arc::new(BufferedLockingProvider {
21207 lock,
21208 stream_started: Arc::clone(&stream_started),
21209 stream_dropped: Arc::clone(&stream_dropped),
21210 committed_after_drop: Arc::clone(&committed_after_drop),
21211 });
21212 let router = Arc::new(RoutingAfterProviderStart {
21213 provider_started: stream_started,
21214 });
21215 let yaml = r#"
21216name: BufferedCancellationAgent
21217system_prompt: "Hide stale streamed output."
21218llm:
21219 default: default
21220 router: router
21221streaming:
21222 enabled: true
21223 buffer_size: 8
21224runtime:
21225 optimization:
21226 enabled: true
21227 max_speculative_llm_calls_per_turn: 2
21228 speculative_state_transitions: true
21229 streaming_policy: buffer_until_routing_done
21230 max_parallel_runtime_tasks: 2
21231states:
21232 initial: triage
21233 states:
21234 triage:
21235 prompt: "Triage state."
21236 transitions:
21237 - to: technical
21238 when: "The request needs technical support"
21239 timing: parallel
21240 technical:
21241 prompt: "Technical state."
21242"#;
21243 let agent = AgentBuilder::from_yaml(yaml)
21244 .unwrap()
21245 .llm_alias("default", default)
21246 .llm_alias("router", router)
21247 .build()
21248 .unwrap();
21249
21250 let content = tokio::time::timeout(std::time::Duration::from_secs(2), async {
21251 let mut stream = agent
21252 .chat_stream("AUTH-17 needs technical help.")
21253 .await
21254 .unwrap();
21255 let mut content = String::new();
21256 while let Some(chunk) = stream.next().await {
21257 match chunk {
21258 StreamChunk::Content { text } => content.push_str(&text),
21259 StreamChunk::Done {} => break,
21260 StreamChunk::Error { message } => panic!("unexpected stream error: {message}"),
21261 _ => {}
21262 }
21263 }
21264 content
21265 })
21266 .await
21267 .expect("redispatch must not wait on the stale streaming future");
21268
21269 assert_eq!(content, "Committed technical response.");
21270 assert_eq!(agent.current_state().as_deref(), Some("technical"));
21271 assert!(stream_dropped.load(Ordering::SeqCst));
21272 assert!(committed_after_drop.load(Ordering::SeqCst));
21273 }
21274
21275 #[tokio::test]
21276 async fn buffered_transition_drops_established_stream_before_redispatch() {
21277 use futures::StreamExt;
21278
21279 let stream_started = Arc::new(tokio::sync::Notify::new());
21280 let stream_dropped = Arc::new(AtomicBool::new(false));
21281 let stream_dropped_notify = Arc::new(tokio::sync::Notify::new());
21282 let committed_after_drop = Arc::new(AtomicBool::new(false));
21283 let default = Arc::new(EstablishedStreamProvider {
21284 stream_started: Arc::clone(&stream_started),
21285 stream_dropped: Arc::clone(&stream_dropped),
21286 stream_dropped_notify,
21287 committed_after_drop: Arc::clone(&committed_after_drop),
21288 });
21289 let router = Arc::new(RoutingAfterProviderStart {
21290 provider_started: stream_started,
21291 });
21292 let yaml = r#"
21293name: EstablishedStreamCancellationAgent
21294system_prompt: "Hide stale streamed output."
21295llm:
21296 default: default
21297 router: router
21298streaming:
21299 enabled: true
21300 buffer_size: 8
21301runtime:
21302 optimization:
21303 enabled: true
21304 max_speculative_llm_calls_per_turn: 2
21305 speculative_state_transitions: true
21306 streaming_policy: buffer_until_routing_done
21307 max_parallel_runtime_tasks: 2
21308states:
21309 initial: triage
21310 states:
21311 triage:
21312 prompt: "Triage state."
21313 transitions:
21314 - to: technical
21315 when: "The request needs technical support"
21316 timing: parallel
21317 technical:
21318 prompt: "Technical state."
21319"#;
21320 let agent = AgentBuilder::from_yaml(yaml)
21321 .unwrap()
21322 .llm_alias("default", default)
21323 .llm_alias("router", router)
21324 .build()
21325 .unwrap();
21326
21327 let content = tokio::time::timeout(std::time::Duration::from_secs(2), async {
21328 let mut stream = agent
21329 .chat_stream("AUTH-17 needs technical help.")
21330 .await
21331 .unwrap();
21332 let mut content = String::new();
21333 while let Some(chunk) = stream.next().await {
21334 match chunk {
21335 StreamChunk::Content { text } => content.push_str(&text),
21336 StreamChunk::Done {} => break,
21337 StreamChunk::Error { message } => panic!("unexpected stream error: {message}"),
21338 _ => {}
21339 }
21340 }
21341 content
21342 })
21343 .await
21344 .expect("redispatch must wait for the established stale stream to be dropped");
21345
21346 assert_eq!(content, "Committed technical response.");
21347 assert_eq!(agent.current_state().as_deref(), Some("technical"));
21348 assert!(stream_dropped.load(Ordering::SeqCst));
21349 assert!(committed_after_drop.load(Ordering::SeqCst));
21350 }
21351
21352 #[tokio::test]
21353 async fn test_buffered_streaming_transition_reservation_falls_back() {
21354 use futures::StreamExt;
21355
21356 let mock = mock_with_responses(vec![
21357 "Serial streaming response",
21358 "Serial streaming response",
21359 ]);
21360 let router_mock = mock_with_response("1");
21361 let router_counter = router_mock.clone();
21362 let yaml = r#"
21363name: BufferedReservationFallbackAgent
21364system_prompt: "Stream normally if speculative routing cannot be evaluated."
21365llm:
21366 default: default
21367 router: router
21368observability:
21369 enabled: true
21370 export:
21371 write_raw_events: true
21372streaming:
21373 enabled: true
21374 buffer_size: 8
21375runtime:
21376 optimization:
21377 enabled: true
21378 max_speculative_llm_calls_per_turn: 1
21379 speculative_state_transitions: true
21380 streaming_policy: buffer_until_routing_done
21381 max_parallel_runtime_tasks: 2
21382states:
21383 initial: triage
21384 states:
21385 triage:
21386 prompt: "Triage state."
21387 transitions:
21388 - to: billing
21389 guard:
21390 context:
21391 route:
21392 eq: billing
21393 when: "User asks about billing"
21394 timing: parallel
21395 billing:
21396 prompt: "Billing state."
21397"#;
21398 let agent = AgentBuilder::from_yaml(yaml)
21399 .unwrap()
21400 .llm_alias("default", Arc::new(mock))
21401 .llm_alias("router", Arc::new(router_mock))
21402 .build()
21403 .unwrap();
21404
21405 let mut stream = agent.chat_stream("hello").await.unwrap();
21406 let mut content = String::new();
21407 let mut error = None;
21408 while let Some(chunk) = stream.next().await {
21409 match chunk {
21410 StreamChunk::Content { text } => content.push_str(&text),
21411 StreamChunk::Error { message } => error = Some(message),
21412 StreamChunk::Done {} => break,
21413 _ => {}
21414 }
21415 }
21416
21417 assert_eq!(error, None);
21418 assert_eq!(content, "Serial streaming response");
21419 assert_eq!(router_counter.call_count(), 0);
21420 let events = agent.observability().unwrap().raw_events();
21421 assert!(events.iter().any(|event| {
21422 event.dimensions.get("branch_status") == Some(&"cancelled".to_string())
21423 && event.dimensions.get("commit_behavior")
21424 == Some(&"transition_decision".to_string())
21425 }));
21426 }
21427
21428 #[tokio::test]
21429 async fn test_blocking_error_cleanup_resets_root_turn_for_next_chat() {
21430 let mut mock = mock_with_response("Recovered response");
21431 mock.set_error("boom");
21432 let mut handle = mock.clone();
21433 let agent = AgentBuilder::new()
21434 .system_prompt("You are helpful.")
21435 .llm(Arc::new(mock))
21436 .build()
21437 .unwrap();
21438
21439 assert!(agent.chat("first").await.is_err());
21440 handle.clear_error();
21441 let response = agent.chat("second").await.unwrap();
21442
21443 assert_eq!(response.content, "Recovered response");
21444 let messages = agent.memory.get_messages(None).await.unwrap();
21445 let user_count = messages
21446 .iter()
21447 .filter(|message| message.role == ai_agents_core::Role::User)
21448 .count();
21449 assert_eq!(user_count, 2);
21450 }
21451
21452 #[tokio::test]
21453 async fn test_streaming_error_cleanup_resets_root_turn_for_next_chat() {
21454 use futures::StreamExt;
21455
21456 let mut mock = mock_with_response("Recovered response");
21457 mock.set_error("stream boom");
21458 let mut handle = mock.clone();
21459 let agent = AgentBuilder::new()
21460 .system_prompt("You are helpful.")
21461 .llm(Arc::new(mock))
21462 .build()
21463 .unwrap();
21464
21465 let mut stream = agent.chat_stream("first").await.unwrap();
21466 let mut saw_error = false;
21467 while let Some(chunk) = stream.next().await {
21468 if matches!(chunk, StreamChunk::Error { .. }) {
21469 saw_error = true;
21470 }
21471 }
21472 assert!(saw_error);
21473
21474 handle.clear_error();
21475 let response = agent.chat("second").await.unwrap();
21476
21477 assert_eq!(response.content, "Recovered response");
21478 let messages = agent.memory.get_messages(None).await.unwrap();
21479 let user_count = messages
21480 .iter()
21481 .filter(|message| message.role == ai_agents_core::Role::User)
21482 .count();
21483 assert_eq!(user_count, 2);
21484 }
21485
21486 #[tokio::test]
21487 async fn test_buffered_streaming_route_miss_releases_buffer_limit() {
21488 use futures::StreamExt;
21489
21490 let mut mock = mock_with_response("one two three");
21491 mock.set_latency(10);
21492 let yaml = r#"
21493name: BufferedMissAgent
21494system_prompt: "You stream safely."
21495llm:
21496 default: default
21497streaming:
21498 enabled: true
21499 buffer_size: 1
21500runtime:
21501 optimization:
21502 enabled: true
21503 max_speculative_llm_calls_per_turn: 2
21504 speculative_state_transitions: true
21505 streaming_policy: buffer_until_routing_done
21506 max_parallel_runtime_tasks: 2
21507states:
21508 initial: triage
21509 states:
21510 triage:
21511 prompt: "Answer from triage."
21512 transitions:
21513 - to: billing
21514 guard:
21515 context:
21516 route:
21517 eq: billing
21518 timing: parallel
21519 billing:
21520 prompt: "Billing state."
21521"#;
21522 let agent = AgentBuilder::from_yaml(yaml)
21523 .unwrap()
21524 .llm_alias("default", Arc::new(mock))
21525 .build()
21526 .unwrap();
21527
21528 let mut stream = agent.chat_stream("hello").await.unwrap();
21529 let mut content = String::new();
21530 let mut error = None;
21531 while let Some(chunk) = stream.next().await {
21532 match chunk {
21533 StreamChunk::Content { text } => content.push_str(&text),
21534 StreamChunk::Error { message } => error = Some(message),
21535 StreamChunk::Done {} => break,
21536 _ => {}
21537 }
21538 }
21539
21540 assert_eq!(error, None);
21541 assert_eq!(content, "one two three");
21542 }
21543
21544 #[tokio::test]
21545 async fn test_buffered_streaming_main_failure_finalizes_branch() {
21546 use futures::StreamExt;
21547
21548 let mock = mock_with_response("one two");
21549 let mut router_mock = mock_with_response("0");
21550 router_mock.set_latency(50);
21551 let yaml = r#"
21552name: BufferedFailureAgent
21553system_prompt: "You stream safely."
21554llm:
21555 default: default
21556 router: router
21557observability:
21558 enabled: true
21559 export:
21560 write_raw_events: true
21561streaming:
21562 enabled: true
21563 buffer_size: 1
21564runtime:
21565 optimization:
21566 enabled: true
21567 max_speculative_llm_calls_per_turn: 2
21568 speculative_state_transitions: true
21569 streaming_policy: buffer_until_routing_done
21570 max_parallel_runtime_tasks: 2
21571states:
21572 initial: triage
21573 states:
21574 triage:
21575 prompt: "Ask for the category."
21576 transitions:
21577 - to: billing
21578 when: "User asks about billing"
21579 timing: parallel
21580 billing:
21581 prompt: "Billing state."
21582"#;
21583 let agent = AgentBuilder::from_yaml(yaml)
21584 .unwrap()
21585 .llm_alias("default", Arc::new(mock))
21586 .llm_alias("router", Arc::new(router_mock))
21587 .build()
21588 .unwrap();
21589
21590 let mut stream = agent.chat_stream("hello").await.unwrap();
21591 let mut error = String::new();
21592 while let Some(chunk) = stream.next().await {
21593 if let StreamChunk::Error { message } = chunk {
21594 error = message;
21595 }
21596 }
21597
21598 assert!(
21599 error.contains("stream buffer filled"),
21600 "unexpected stream error: {}",
21601 error
21602 );
21603 let events = agent.observability().unwrap().raw_events();
21604 assert!(events.iter().any(|event| {
21605 event.dimensions.get("branch_status") == Some(&"failed".to_string())
21606 && event.dimensions.get("commit_behavior") == Some(&"final_response".to_string())
21607 && event.dimensions.get("optimization")
21608 == Some(&"buffered_streaming_routing".to_string())
21609 }));
21610 }
21611
21612 #[tokio::test]
21613 async fn test_streaming_preflight_does_not_emit_old_state_content() {
21614 use futures::StreamExt;
21615
21616 let mock = mock_with_response("Billing streamed response");
21617 let yaml = r#"
21618name: StreamingOptimizedAgent
21619system_prompt: "You route before streaming."
21620runtime:
21621 optimization:
21622 enabled: true
21623 pre_response_deterministic_transitions: true
21624streaming:
21625 enabled: true
21626states:
21627 initial: greeting
21628 states:
21629 greeting:
21630 prompt: "OLD_STATE_SENTINEL"
21631 transitions:
21632 - to: billing
21633 guard:
21634 context:
21635 topic:
21636 eq: billing
21637 timing: pre_response
21638 billing:
21639 prompt: "Billing state."
21640"#;
21641 let agent = AgentBuilder::from_yaml(yaml)
21642 .unwrap()
21643 .llm(Arc::new(mock))
21644 .build()
21645 .unwrap();
21646 agent
21647 .set_context("topic", serde_json::json!("billing"))
21648 .unwrap();
21649
21650 let mut stream = agent.chat_stream("billing please").await.unwrap();
21651 let mut content = String::new();
21652 while let Some(chunk) = stream.next().await {
21653 match chunk {
21654 StreamChunk::Content { text } => content.push_str(&text),
21655 StreamChunk::Error { message } => panic!("stream error: {}", message),
21656 StreamChunk::Done {} => break,
21657 _ => {}
21658 }
21659 }
21660
21661 assert_eq!(agent.current_state().as_deref(), Some("billing"));
21662 assert!(content.contains("Billing streamed response"));
21663 assert!(!content.contains("OLD_STATE_SENTINEL"));
21664 }
21665
21666 #[tokio::test]
21668 async fn test_integration_state_machine_basic() {
21669 let yaml = r#"
21670name: StateAgent
21671system_prompt: "You are a support agent."
21672states:
21673 initial: greeting
21674 states:
21675 greeting:
21676 prompt: "Welcome the user warmly."
21677 transitions:
21678 - to: support
21679 when: "User needs help"
21680 auto: true
21681 support:
21682 prompt: "Help solve the user's problem."
21683"#;
21684 let mock = mock_with_responses(vec![
21685 "Welcome! How can I help?", "1", "I'll help you with that.", ]);
21689 let builder = AgentBuilder::from_yaml(yaml).unwrap();
21690 let agent = builder.llm(Arc::new(mock)).build().unwrap();
21691
21692 assert_eq!(agent.current_state(), Some("greeting".to_string()));
21693 let _ = agent.chat("I need help").await.unwrap();
21694 }
21697
21698 #[tokio::test]
21700 async fn test_integration_state_on_enter_set_context() {
21701 let yaml = r#"
21702name: ActionAgent
21703system_prompt: "You are helpful."
21704states:
21705 initial: step1
21706 states:
21707 step1:
21708 prompt: "Step 1"
21709 on_exit:
21710 - set_context:
21711 step1_exited: true
21712 transitions:
21713 - to: step2
21714 when: "always"
21715 auto: true
21716 step2:
21717 prompt: "Step 2"
21718 on_enter:
21719 - set_context:
21720 step2_entered: true
21721"#;
21722 let mock = mock_with_responses(vec![
21724 "Processing step 1.",
21725 "0", ]);
21727 let builder = AgentBuilder::from_yaml(yaml).unwrap();
21728 let agent = builder.llm(Arc::new(mock)).build().unwrap();
21729
21730 assert_eq!(agent.current_state(), Some("step1".to_string()));
21731
21732 agent.transition_to("step2").await.unwrap();
21734
21735 assert_eq!(agent.current_state(), Some("step2".to_string()));
21736
21737 let ctx = agent.get_context();
21739 assert_eq!(ctx.get("step1_exited"), Some(&serde_json::json!(true)));
21740 assert_eq!(ctx.get("step2_entered"), Some(&serde_json::json!(true)));
21741 }
21742
21743 #[tokio::test]
21744 async fn state_action_tool_preserves_source_in_stored_record() {
21745 let yaml = r#"
21746name: StateActionToolAgent
21747system_prompt: "You are helpful."
21748tools:
21749 - context_echo
21750states:
21751 initial: idle
21752 states:
21753 idle:
21754 prompt: "Idle"
21755 active:
21756 prompt: "Active"
21757 on_enter:
21758 - set_context:
21759 action_started: true
21760 - tool: context_echo
21761 args: {}
21762"#;
21763 let agent = AgentBuilder::from_yaml(yaml)
21764 .unwrap()
21765 .llm(Arc::new(mock_with_response("unused")))
21766 .tool(Arc::new(ContextEchoTool))
21767 .build()
21768 .unwrap();
21769
21770 agent.transition_to("active").await.unwrap();
21771
21772 let record: ToolExecutionRecord = serde_json::from_value(
21773 agent
21774 .get_context()
21775 .get("last_tool_record")
21776 .cloned()
21777 .expect("successful state action must store its execution record"),
21778 )
21779 .unwrap();
21780 assert!(record.executed);
21781 assert!(record.success);
21782 assert_eq!(record.canonical_id, "context_echo");
21783 assert!(matches!(
21784 &record.source,
21785 ToolCallSource::StateAction {
21786 state: Some(state),
21787 action_index: 1,
21788 } if state == "active"
21789 ));
21790 }
21791
21792 #[tokio::test]
21793 async fn test_ordinary_transition_uses_on_enter_then_on_reenter() {
21794 let yaml = r#"
21795name: OrdinaryLifecycleAgent
21796system_prompt: "You are helpful."
21797states:
21798 initial: intake
21799 regenerate_on_transition: false
21800 states:
21801 intake:
21802 prompt: "Intake"
21803 transitions:
21804 - to: drafting
21805 guard:
21806 context:
21807 route:
21808 eq: drafting
21809 drafting:
21810 prompt: "Drafting"
21811 on_enter:
21812 - set_context:
21813 draft_version: 1
21814 on_reenter:
21815 - set_context:
21816 draft_version: 2
21817 transitions:
21818 - to: review
21819 guard:
21820 context:
21821 route:
21822 eq: review
21823 review:
21824 prompt: "Review"
21825 on_enter:
21826 - set_context:
21827 review_entry: first
21828 transitions:
21829 - to: drafting
21830 guard:
21831 context:
21832 route:
21833 eq: drafting
21834"#;
21835 let agent = AgentBuilder::from_yaml(yaml)
21836 .unwrap()
21837 .llm(Arc::new(mock_with_responses(vec![
21838 "Intake response",
21839 "Draft response",
21840 "Review response",
21841 ])))
21842 .build()
21843 .unwrap();
21844
21845 agent
21846 .set_context("route", serde_json::json!("drafting"))
21847 .unwrap();
21848 agent.chat("Start a draft").await.unwrap();
21849 assert_eq!(agent.current_state().as_deref(), Some("drafting"));
21850 assert_eq!(
21851 agent.get_context().get("draft_version"),
21852 Some(&serde_json::json!(1))
21853 );
21854
21855 agent
21856 .set_context("route", serde_json::json!("review"))
21857 .unwrap();
21858 agent.chat("Review this").await.unwrap();
21859 assert_eq!(agent.current_state().as_deref(), Some("review"));
21860 assert_eq!(
21861 agent.get_context().get("review_entry"),
21862 Some(&serde_json::json!("first"))
21863 );
21864
21865 agent
21866 .set_context("route", serde_json::json!("drafting"))
21867 .unwrap();
21868 agent.chat("Revise this").await.unwrap();
21869 assert_eq!(agent.current_state().as_deref(), Some("drafting"));
21870 assert_eq!(
21871 agent.get_context().get("draft_version"),
21872 Some(&serde_json::json!(2))
21873 );
21874 }
21875
21876 #[tokio::test]
21877 async fn test_manual_transition_uses_on_enter_then_on_reenter() {
21878 let yaml = r#"
21879name: ManualLifecycleAgent
21880system_prompt: "You are helpful."
21881states:
21882 initial: intake
21883 states:
21884 intake:
21885 prompt: "Intake"
21886 drafting:
21887 prompt: "Drafting"
21888 on_enter:
21889 - set_context:
21890 draft_version: 1
21891 on_reenter:
21892 - set_context:
21893 draft_version: 2
21894 review:
21895 prompt: "Review"
21896"#;
21897 let agent = AgentBuilder::from_yaml(yaml)
21898 .unwrap()
21899 .llm(Arc::new(mock_with_response("unused")))
21900 .build()
21901 .unwrap();
21902
21903 assert!(!agent.get_context().contains_key("draft_version"));
21904 agent.transition_to("drafting").await.unwrap();
21905 assert_eq!(agent.current_state().as_deref(), Some("drafting"));
21906 assert_eq!(
21907 agent.get_context().get("draft_version"),
21908 Some(&serde_json::json!(1))
21909 );
21910
21911 agent.transition_to("review").await.unwrap();
21912 agent.transition_to("drafting").await.unwrap();
21913 assert_eq!(agent.current_state().as_deref(), Some("drafting"));
21914 assert_eq!(
21915 agent.get_context().get("draft_version"),
21916 Some(&serde_json::json!(2))
21917 );
21918 }
21919
21920 #[tokio::test]
21921 async fn test_timeout_transition_uses_on_enter_then_on_reenter() {
21922 let yaml = r#"
21923name: TimeoutLifecycleAgent
21924system_prompt: "You are helpful."
21925states:
21926 initial: intake
21927 regenerate_on_transition: false
21928 states:
21929 intake:
21930 prompt: "Intake"
21931 max_turns: 1
21932 timeout_to: drafting
21933 drafting:
21934 prompt: "Drafting"
21935 max_turns: 1
21936 timeout_to: review
21937 on_enter:
21938 - set_context:
21939 draft_version: 1
21940 on_reenter:
21941 - set_context:
21942 draft_version: 2
21943 review:
21944 prompt: "Review"
21945 max_turns: 1
21946 timeout_to: drafting
21947 on_enter:
21948 - set_context:
21949 review_entry: first
21950"#;
21951 let agent = AgentBuilder::from_yaml(yaml)
21952 .unwrap()
21953 .llm(Arc::new(mock_with_responses(vec![
21954 "Intake",
21955 "First draft",
21956 "Review",
21957 "Revised draft",
21958 ])))
21959 .build()
21960 .unwrap();
21961
21962 agent.chat("First turn").await.unwrap();
21963 assert_eq!(agent.current_state().as_deref(), Some("intake"));
21964 assert!(!agent.get_context().contains_key("draft_version"));
21965
21966 agent.chat("Second turn").await.unwrap();
21967 assert_eq!(agent.current_state().as_deref(), Some("drafting"));
21968 assert_eq!(
21969 agent.get_context().get("draft_version"),
21970 Some(&serde_json::json!(1))
21971 );
21972
21973 agent.chat("Third turn").await.unwrap();
21974 assert_eq!(agent.current_state().as_deref(), Some("review"));
21975 assert_eq!(
21976 agent.get_context().get("review_entry"),
21977 Some(&serde_json::json!("first"))
21978 );
21979
21980 agent.chat("Fourth turn").await.unwrap();
21981 assert_eq!(agent.current_state().as_deref(), Some("drafting"));
21982 assert_eq!(
21983 agent.get_context().get("draft_version"),
21984 Some(&serde_json::json!(2))
21985 );
21986 }
21987
21988 #[tokio::test]
21990 async fn test_integration_process_normalize() {
21991 let yaml = r#"
21992name: ProcessAgent
21993system_prompt: "You are helpful."
21994process:
21995 input:
21996 - type: normalize
21997 config:
21998 trim: true
21999 collapse_whitespace: true
22000"#;
22001 let mock = mock_with_response("Got your message.");
22002 let builder = AgentBuilder::from_yaml(yaml).unwrap();
22003 let agent = builder.llm(Arc::new(mock.clone())).build().unwrap();
22004
22005 let _ = agent.chat(" hello world ").await.unwrap();
22006
22007 let history = mock.call_history();
22009 assert!(!history.is_empty());
22010 let last_call = history.last().unwrap();
22012 let user_msg = last_call
22013 .messages
22014 .iter()
22015 .find(|m| m.role == ai_agents_core::Role::User)
22016 .unwrap();
22017 assert_eq!(user_msg.content, "hello world");
22018 }
22019
22020 #[tokio::test]
22024 async fn test_integration_memory_compression() {
22025 let yaml = r#"
22026name: MemoryAgent
22027system_prompt: "You are helpful."
22028memory:
22029 type: compacting
22030 max_messages: 100
22031 compress_threshold: 5
22032 max_recent_messages: 3
22033 summarize_batch_size: 2
22034"#;
22035 let responses: Vec<&str> = (0..8).map(|_| "Response from assistant.").collect();
22037 let mock = mock_with_responses(responses);
22038 let builder = AgentBuilder::from_yaml(yaml).unwrap();
22039 let agent = builder.llm(Arc::new(mock)).build().unwrap();
22040
22041 for i in 0..6 {
22043 let _ = agent.chat(&format!("Message {}", i)).await.unwrap();
22044 }
22045
22046 let messages = agent.memory.get_messages(None).await.unwrap();
22049 assert!(messages.len() <= 12); }
22053
22054 #[tokio::test]
22056 async fn test_integration_multi_llm_registry() {
22057 let mut mock_default = MockLLMProvider::new("default");
22058 mock_default.set_response("Default LLM response.");
22059 let mut mock_router = MockLLMProvider::new("router");
22060 mock_router.set_response("Router response.");
22061
22062 let agent = AgentBuilder::new()
22063 .system_prompt("You are helpful.")
22064 .llm_alias("default", Arc::new(mock_default))
22065 .llm_alias("router", Arc::new(mock_router))
22066 .build()
22067 .unwrap();
22068
22069 let response = agent.chat("Hello").await.unwrap();
22070 assert_eq!(response.content, "Default LLM response.");
22071 }
22072
22073 #[tokio::test]
22075 async fn test_integration_agent_reset() {
22076 let mock = mock_with_responses(vec!["Hello!", "Hello again!"]);
22077 let agent = AgentBuilder::new()
22078 .system_prompt("You are helpful.")
22079 .llm(Arc::new(mock))
22080 .build()
22081 .unwrap();
22082
22083 let _ = agent.chat("Hi").await.unwrap();
22084 let messages = agent.memory.get_messages(None).await.unwrap();
22085 assert_eq!(messages.len(), 2); agent.reset().await.unwrap();
22088 let messages = agent.memory.get_messages(None).await.unwrap();
22089 assert_eq!(messages.len(), 0);
22090 }
22091
22092 #[tokio::test]
22094 async fn test_integration_process_validate_reject() {
22095 use ai_agents_process::{ProcessConfig, ProcessProcessor};
22096
22097 let validate_config = ai_agents_process::ValidateStage {
22098 id: Some("length_check".to_string()),
22099 condition: None,
22100 config: ai_agents_process::ValidateConfig {
22101 rules: vec![ai_agents_process::ValidationRule::MinLength {
22102 min_length: 10,
22103 on_fail: ai_agents_process::ValidationAction {
22104 action: ai_agents_process::ValidationActionType::Reject,
22105 message: None,
22106 },
22107 }],
22108 ..Default::default()
22109 },
22110 };
22111 let process_config = ProcessConfig {
22112 input: vec![ai_agents_process::ProcessStage::Validate(validate_config)],
22113 ..Default::default()
22114 };
22115 let processor = ProcessProcessor::new(process_config);
22116
22117 let mock = mock_with_response("Should not reach here.");
22118 let agent = AgentBuilder::new()
22119 .system_prompt("You are helpful.")
22120 .llm(Arc::new(mock))
22121 .process_processor(processor)
22122 .build()
22123 .unwrap();
22124
22125 let response = agent.chat("Hi").await.unwrap();
22126 assert!(
22128 response.content.contains("rejected")
22129 || response.content.contains("Input rejected")
22130 || response.content.contains("too short")
22131 || response.content.contains("Too short")
22132 || response.content.len() < 50, "Expected rejection response, got: {}",
22134 response.content
22135 );
22136 }
22137
22138 #[tokio::test]
22140 async fn test_llm_fallback_on_failure() {
22141 use ai_agents_recovery::{ErrorRecoveryConfig, LLMFailureAction, LLMRecoveryConfig};
22142
22143 let mut primary = MockLLMProvider::new("primary");
22144 primary.set_error("Primary LLM is unavailable");
22145
22146 let mut fallback = MockLLMProvider::new("fallback");
22147 fallback.set_response("Fallback response works!");
22148
22149 let agent = AgentBuilder::new()
22150 .system_prompt("You are helpful.")
22151 .llm_alias("default", Arc::new(primary))
22152 .llm_alias("backup", Arc::new(fallback))
22153 .recovery_manager(RecoveryManager::new(ErrorRecoveryConfig {
22154 llm: LLMRecoveryConfig {
22155 on_failure: LLMFailureAction::FallbackLlm {
22156 fallback_llm: "backup".to_string(),
22157 },
22158 ..Default::default()
22159 },
22160 ..Default::default()
22161 }))
22162 .build()
22163 .unwrap();
22164
22165 let response = agent.chat("Hello").await.unwrap();
22166 assert!(
22167 response.content.contains("Fallback response"),
22168 "Expected fallback response, got: {}",
22169 response.content
22170 );
22171 }
22172
22173 #[tokio::test]
22175 async fn test_llm_fallback_response_static_message() {
22176 use ai_agents_recovery::{ErrorRecoveryConfig, LLMFailureAction, LLMRecoveryConfig};
22177
22178 let mut primary = MockLLMProvider::new("primary");
22179 primary.set_error("Primary LLM is unavailable");
22180
22181 let agent = AgentBuilder::new()
22182 .system_prompt("You are helpful.")
22183 .llm(Arc::new(primary))
22184 .recovery_manager(RecoveryManager::new(ErrorRecoveryConfig {
22185 llm: LLMRecoveryConfig {
22186 on_failure: LLMFailureAction::FallbackResponse {
22187 message: "I am temporarily unavailable. Please try again later."
22188 .to_string(),
22189 },
22190 ..Default::default()
22191 },
22192 ..Default::default()
22193 }))
22194 .build()
22195 .unwrap();
22196
22197 let response = agent.chat("Hello").await.unwrap();
22198 assert!(
22199 response.content.contains("temporarily unavailable"),
22200 "Expected static fallback message, got: {}",
22201 response.content
22202 );
22203 }
22204
22205 #[tokio::test]
22207 async fn test_tool_failure_skip() {
22208 use ai_agents_recovery::{
22209 ErrorRecoveryConfig, ToolFailureAction, ToolRecoveryConfig, ToolRetryConfig,
22210 };
22211
22212 let mock = mock_with_responses(vec![
22214 r#"I'll use the nonexistent tool.
22215[TOOL_CALL: {"name": "nonexistent_tool", "arguments": {}}]"#,
22216 "The tool was unavailable, but I can still help you.",
22217 ]);
22218
22219 let agent = AgentBuilder::new()
22220 .system_prompt("You are helpful.")
22221 .llm(Arc::new(mock))
22222 .recovery_manager(RecoveryManager::new(ErrorRecoveryConfig {
22223 tools: ToolRecoveryConfig {
22224 default: ToolRetryConfig {
22225 max_retries: 0,
22226 timeout_ms: None,
22227 on_failure: ToolFailureAction::Skip,
22228 },
22229 ..Default::default()
22230 },
22231 ..Default::default()
22232 }))
22233 .build()
22234 .unwrap();
22235
22236 let response = agent.chat("Use the nonexistent tool").await;
22238 assert!(
22239 response.is_ok(),
22240 "Expected Ok with skip policy, got: {:?}",
22241 response
22242 );
22243 }
22244}