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
16#[cfg(test)]
17#[path = "routing_tests.rs"]
18mod routing_tests;
19
20pub(crate) type RootTurnGate = Arc<tokio::sync::Mutex<()>>;
22
23pub(crate) type RootTurnGateIdentityStack = Arc<[RootTurnGate]>;
25
26tokio::task_local! {
27 static RUNTIME_GATE_IDENTITY_STACK: RootTurnGateIdentityStack;
28}
29
30pub(crate) fn current_runtime_gate_identity_stack() -> RootTurnGateIdentityStack {
32 RUNTIME_GATE_IDENTITY_STACK
33 .try_with(Arc::clone)
34 .unwrap_or_default()
35}
36
37pub(crate) async fn scope_runtime_gate_identity_stack<F, T>(
39 identity_stack: &RootTurnGateIdentityStack,
40 future: F,
41) -> T
42where
43 F: Future<Output = T>,
44{
45 RUNTIME_GATE_IDENTITY_STACK
46 .scope(Arc::clone(identity_stack), future)
47 .await
48}
49
50pub(crate) type ToolResourceLocks = Arc<RwLock<HashMap<String, Weak<tokio::sync::Mutex<()>>>>>;
52
53struct ToolResourceGuards {
57 guards: Vec<tokio::sync::OwnedMutexGuard<()>>,
58 locks: ToolResourceLocks,
59}
60
61struct RootTurnAdmission {
65 guard: tokio::sync::OwnedMutexGuard<()>,
66 identity_stack: RootTurnGateIdentityStack,
67}
68
69#[derive(Clone)]
70struct StoredSessionRestore {
71 snapshot: AgentSnapshot,
72 metadata: Option<ai_agents_core::SessionMetadata>,
73}
74
75struct RuntimeSessionRestorePoint {
76 snapshot: AgentSnapshot,
77 metadata: ai_agents_core::SessionMetadata,
78 actor_id: Option<String>,
79 session_id: Option<String>,
80}
81
82impl Drop for ToolResourceGuards {
83 fn drop(&mut self) {
84 self.guards.clear();
85 self.locks.write().retain(|_, lock| lock.strong_count() > 0);
86 }
87}
88
89#[derive(Clone)]
93struct RuntimeSafetySnapshot {
94 version: u64,
95 emergency_deny: bool,
96 tool_security: ToolSecurityEngine,
97 tool_scope_override: Option<Vec<String>>,
98}
99
100#[derive(Clone, Copy)]
104struct ToolDecisionVersions {
105 policy: u64,
106 registry: u64,
107 runtime_control: u64,
108 state: Option<u64>,
109}
110
111#[derive(Clone, Debug, Default)]
115struct ToolFallbackState {
116 visited_canonical_ids: Vec<String>,
117}
118
119impl ToolFallbackState {
120 fn rejection_reason(&self, canonical_id: &str) -> Option<String> {
124 if self
125 .visited_canonical_ids
126 .iter()
127 .any(|visited| visited == canonical_id)
128 {
129 return Some(format!(
130 "Tool fallback cycle detected at '{canonical_id}' after [{}]",
131 self.visited_canonical_ids.join(" -> ")
132 ));
133 }
134 if self.visited_canonical_ids.len() > MAX_TOOL_FALLBACK_HOPS {
135 return Some(format!(
136 "Tool fallback chain exceeds the maximum of {MAX_TOOL_FALLBACK_HOPS} hops"
137 ));
138 }
139 None
140 }
141
142 fn with_current(mut self, canonical_id: String) -> Self {
146 self.visited_canonical_ids.push(canonical_id);
147 self
148 }
149
150 fn final_rejection_reason(
154 &self,
155 admitted_canonical_id: &str,
156 final_canonical_id: &str,
157 ) -> Option<String> {
158 if admitted_canonical_id == final_canonical_id {
159 return None;
160 }
161 if self
162 .visited_canonical_ids
163 .iter()
164 .any(|visited| visited == final_canonical_id)
165 {
166 return Some(format!(
167 "Tool fallback cycle detected after final resolution changed '{admitted_canonical_id}' to '{final_canonical_id}'"
168 ));
169 }
170 Some(format!(
171 "Tool canonical target changed after initial admission from '{admitted_canonical_id}' to '{final_canonical_id}'"
172 ))
173 }
174}
175
176#[derive(Clone, Copy, Debug)]
180struct ValidatedToolTimeout {
181 timer: Duration,
182 deadline_delta: chrono::Duration,
183}
184
185struct AvailableToolIdsSnapshot {
189 tool_ids: Vec<String>,
190 state_generation: Option<u64>,
191}
192
193#[derive(Clone)]
197struct ToolApprovalBinding {
198 canonical_id: String,
199 arguments: Value,
200 confirmation_required: bool,
201 policy_version: u64,
202 runtime_control_version: u64,
203 state_generation: Option<u64>,
204 reviewed_tool: Arc<dyn ai_agents_core::Tool>,
205}
206
207fn merge_approved_record(record: &mut Option<ToolApprovalRecord>) {
211 if record
212 .as_ref()
213 .is_some_and(|record| matches!(record.status, ToolApprovalStatus::Modified))
214 {
215 return;
216 }
217 *record = Some(ToolApprovalRecord {
218 status: ToolApprovalStatus::Approved,
219 reason: None,
220 modified_arguments: None,
221 });
222}
223
224impl ToolApprovalBinding {
225 fn is_stale(
227 &self,
228 canonical_id: &str,
229 arguments: &Value,
230 confirmation_required: bool,
231 versions: ToolDecisionVersions,
232 resolved_tool: &Arc<dyn ai_agents_core::Tool>,
233 ) -> bool {
234 self.canonical_id != canonical_id
235 || self.arguments != *arguments
236 || self.confirmation_required != confirmation_required
237 || self.policy_version != versions.policy
238 || self.runtime_control_version != versions.runtime_control
239 || self.state_generation != versions.state
240 || !Arc::ptr_eq(&self.reviewed_tool, resolved_tool)
241 }
242}
243
244use crate::turn_context::{current_turn_actor_context, scope_actor_context};
245
246use ai_agents_context::{ContextManager, ContextProvider, TemplateRenderer};
247use ai_agents_core::traits::storage::StorageCapability;
248use ai_agents_core::{
249 AgentError, AgentSnapshot, AgentStorage, ChatMessage, FinishReason, LLMChunk, LLMError,
250 LLMFeature, LLMProvider, LLMResponse, LLMToolDefinition, LLMToolRequest, PermissionOutcome,
251 Result, ToolActorContext, ToolApprovalRecord, ToolApprovalStatus, ToolCallClassification,
252 ToolCallSource, ToolCancellationToken, ToolChoice, ToolExecutionContext, ToolExecutionLimits,
253 ToolExecutionRecord, ToolExecutionRequest, ToolInvoker, ToolPolicyDecisionRecord, ToolResult,
254 ToolSafetyMetadata, decode_native_tool_call_markers, encode_native_tool_call_markers,
255 encode_native_tool_result_marker, inspect_native_history, native_readable_projection,
256};
257use ai_agents_disambiguation::{
258 AmbiguityDetectionResult, ClarificationObserver, ClarificationParseFuture,
259 ClarificationQuestion, ClarificationQuestionFuture, ConfirmationParseFuture,
260 DisambiguationConfig, DisambiguationContext, DisambiguationManager, DisambiguationResult,
261};
262use ai_agents_hitl::{
263 ApprovalHandler, ApprovalResolvedOutcome, ApprovalResult, ApprovalTrigger, HITLCheckResult,
264 HITLEngine, RejectAllHandler, TimeoutAction,
265};
266use ai_agents_hooks::{AgentHooks, NoopHooks};
267use ai_agents_llm::LLMRegistry;
268use ai_agents_memory::{
269 CompressResult, EvictionReason, Memory, MemoryBudgetEvent, MemoryCompressEvent,
270 MemoryEvictEvent, MemoryTokenBudget, OverflowStrategy,
271};
272use ai_agents_observability::{
273 EventStatus, EventType, ObservabilityManager, ObservationPurpose, SpanContext,
274 current_observation_context, new_session_id as new_observation_session_id,
275 resolve_language_from_context, with_observation_context, with_observation_purpose,
276};
277use ai_agents_process::{
278 ProcessData, ProcessProcessor, ProcessPurposeHint, ProcessStageFuture, ProcessStageObserver,
279};
280use ai_agents_reasoning::{
281 CriterionResult, EvaluationResult, Plan, PlanAction, PlanStatus, PlanStep, ReasoningConfig,
282 ReasoningMetadata, ReasoningMode, ReasoningOutput, ReflectionAttempt, ReflectionConfig,
283 ReflectionMetadata, ReflectionMode, StepFailureAction,
284};
285use ai_agents_recovery::{
286 ByRoleFilter, ContextOverflowAction, FilterConfig, KeepRecentFilter, LLMFailureAction,
287 MessageFilter, RecoveryManager, SkipPatternFilter, ToolFailureAction,
288};
289use ai_agents_relationships::RelationshipManager;
290use ai_agents_skills::{SkillDefinition, SkillExecutor, SkillRouter};
291use ai_agents_state::{
292 PromptMode, StateAction, StateMachine, StateMachineSnapshot, StateTransitionEvent, Transition,
293 TransitionContext, TransitionEvaluator, TransitionTiming, evaluate_guard,
294};
295use ai_agents_storage::{StorageConfig as StorageStorageConfig, create_storage};
296use ai_agents_tools::{
297 CommandRunner, ConditionEvaluator, DiagnosticsProvider, EvaluationContext, LLMGetter,
298 MAX_TOOL_TIMEOUT_MS, QuestionHandler, SecurityCheckResult, TodoItem, ToolCallRecord,
299 ToolRegistry, ToolSecurityConfig, ToolSecurityEngine,
300};
301
302use super::{
303 Agent, AgentInfo, AgentResponse, AgentStreamEvent, ParallelToolsConfig, StreamChunk,
304 StreamingConfig, ToolCall,
305};
306use crate::optimization::{
307 AwaitBeforeNextTurn, BackgroundMaintenanceQueue, BackgroundOverflowPolicy, MainResponseDraft,
308 MaintenanceMode, MaintenanceSequenceKey, RuntimeBranch, RuntimeBranchResult,
309 RuntimeBranchStatus, RuntimeCommitBehavior, RuntimeConfig, RuntimeOptimizationKind,
310 RuntimeTaskPriority, RuntimeTaskPurpose, ScheduledBranchSet, SkillCandidate,
311 StreamingDraftResult, TransitionCandidate, TurnBranchScheduler, TurnOptimizationContext,
312};
313use crate::spec::StorageConfig;
314
315enum ToolCallOutcome {
317 Continue,
319 TransitionFired,
321 Rejected(AgentResponse),
323}
324
325#[derive(Clone)]
326struct MainToolProtocol {
327 choice: Option<ToolChoice>,
328 tool_ids: Vec<String>,
329 definitions: Vec<LLMToolDefinition>,
330}
331
332struct MainProviderResponse {
333 response: LLMResponse,
334 used_native_tools: bool,
335}
336
337enum MainStreamSource {
342 Stream(Box<dyn Stream<Item = std::result::Result<LLMChunk, LLMError>> + Unpin + Send>),
343 StaticResponse(String),
344}
345
346#[derive(Clone)]
348struct ActiveNativeExchange {
349 exchange_id: String,
350 call_ids: Vec<String>,
351}
352
353struct CommittedTextResponse<'a> {
357 processed_input: &'a str,
358 input_context: &'a HashMap<String, Value>,
359 answer: String,
360 reasoning_mode: ReasoningMode,
361 auto_detected: bool,
362 iterations: u32,
363 thinking_content: Option<String>,
364 all_tool_calls: Vec<ToolCall>,
365}
366
367struct AgentResponseParts {
371 content: String,
372 all_tool_calls: Vec<ToolCall>,
373 reasoning_mode: ReasoningMode,
374 auto_detected: bool,
375 iterations: u32,
376 thinking: Option<String>,
377 reflection_metadata: Option<ReflectionMetadata>,
378}
379
380type RuntimeStreamTerminalSlot = Arc<RwLock<Option<AgentResponse>>>;
384
385fn new_runtime_stream_terminal_slot() -> RuntimeStreamTerminalSlot {
389 Arc::new(RwLock::new(None))
390}
391
392fn record_runtime_stream_final(slot: &RuntimeStreamTerminalSlot, response: AgentResponse) {
396 *slot.write() = Some(response);
397}
398
399#[derive(Clone, Copy)]
400struct DisambiguationOwnership {
401 epoch: u64,
402 state_generation: Option<u64>,
403}
404
405enum SkillRouteResult {
407 NoMatch,
409 Response { skill_id: String, content: String },
411 NeedsClarification {
413 response: AgentResponse,
414 ownership: Option<DisambiguationOwnership>,
415 },
416}
417
418enum ParallelTransitionSelection {
420 Candidate(TransitionCandidate),
422 NoMatch,
424 ReservationExhausted,
426}
427
428enum DisambiguationDispatch {
434 Proceed(String),
436 Terminal(AgentResponse),
438 RecheckSkill {
440 skill_id: String,
441 enriched_input: String,
442 disambiguation_epoch: u64,
443 state_generation: Option<u64>,
444 },
445}
446
447enum PostLoopResult {
448 NoTransition(String),
450 Transitioned { content: String, regenerated: bool },
453 NeedsRedispatch,
456}
457
458struct AppliedPostLoop {
460 content: String,
461 transitioned: bool,
462 regenerated: bool,
464}
465
466struct StateTransitionReservation<'a> {
467 reserved: &'a AtomicBool,
468}
469
470impl Drop for StateTransitionReservation<'_> {
471 fn drop(&mut self) {
472 self.reserved.store(false, Ordering::SeqCst);
473 }
474}
475
476struct RootTurnCleanup<'a> {
477 agent: &'a RuntimeAgent,
478}
479
480impl<'a> RootTurnCleanup<'a> {
481 fn new(agent: &'a RuntimeAgent) -> Self {
482 Self { agent }
483 }
484}
485
486impl Drop for RootTurnCleanup<'_> {
487 fn drop(&mut self) {
488 self.agent.end_root_turn();
489 }
490}
491
492#[derive(Debug)]
494struct RuntimeControlState {
495 snapshot_guard: RwLock<()>,
497 version: AtomicU64,
499 emergency_deny: Arc<AtomicBool>,
501 tool_security_override: RwLock<Option<ToolSecurityEngine>>,
503 tool_scope_override: RwLock<Option<Vec<String>>>,
505}
506
507impl Default for RuntimeControlState {
508 fn default() -> Self {
509 Self {
510 snapshot_guard: RwLock::new(()),
511 version: AtomicU64::new(1),
512 emergency_deny: Arc::new(AtomicBool::new(false)),
513 tool_security_override: RwLock::new(None),
514 tool_scope_override: RwLock::new(None),
515 }
516 }
517}
518
519#[derive(Clone)]
521pub struct RuntimeControlHandle {
522 state: Arc<RuntimeControlState>,
523}
524
525impl RuntimeControlHandle {
526 pub fn version(&self) -> u64 {
528 self.state.version.load(Ordering::SeqCst)
529 }
530
531 fn bump(&self) -> u64 {
532 self.state.version.fetch_add(1, Ordering::SeqCst) + 1
533 }
534
535 pub fn set_tool_security(&self, config: ToolSecurityConfig) -> u64 {
537 self.try_set_tool_security(config)
538 .expect("invalid tool security configuration")
539 }
540
541 pub fn try_set_tool_security(&self, config: ToolSecurityConfig) -> Result<u64> {
543 config.validate()?;
544 let _guard = self.state.snapshot_guard.write();
545 let generation = self.bump();
546 *self.state.tool_security_override.write() = Some(
547 ToolSecurityEngine::new_with_policy_version(config, generation),
548 );
549 Ok(generation)
550 }
551
552 pub fn clear_tool_security_override(&self) -> u64 {
554 let _guard = self.state.snapshot_guard.write();
555 *self.state.tool_security_override.write() = None;
556 self.bump()
557 }
558
559 pub fn set_tool_scope(&self, tool_ids: Vec<String>) -> u64 {
561 let _guard = self.state.snapshot_guard.write();
562 *self.state.tool_scope_override.write() = Some(tool_ids);
563 self.bump()
564 }
565
566 pub fn clear_tool_scope_override(&self) -> u64 {
568 let _guard = self.state.snapshot_guard.write();
569 *self.state.tool_scope_override.write() = None;
570 self.bump()
571 }
572
573 pub fn set_emergency_deny(&self, enabled: bool) -> u64 {
575 let _guard = self.state.snapshot_guard.write();
576 self.state.emergency_deny.store(enabled, Ordering::SeqCst);
577 self.bump()
578 }
579
580 pub fn cancel_all(&self) -> u64 {
582 self.set_emergency_deny(true)
583 }
584}
585
586pub struct RuntimeAgent {
587 info: AgentInfo,
588 llm_registry: Arc<LLMRegistry>,
589 memory: Arc<dyn Memory>,
590 tools: Arc<ToolRegistry>,
591 skills: Vec<SkillDefinition>,
592 skill_router: Option<SkillRouter>,
593 skill_executor: Option<SkillExecutor>,
594 base_system_prompt: String,
595 max_iterations: u32,
596 iteration_count: RwLock<u32>,
597 max_context_tokens: u32,
598 memory_token_budget: Option<MemoryTokenBudget>,
599 recovery_manager: RecoveryManager,
600 tool_security: ToolSecurityEngine,
601 process_processor: Option<ProcessProcessor>,
602 message_filters: RwLock<HashMap<String, Arc<dyn MessageFilter>>>,
603 state_machine: Option<Arc<StateMachine>>,
604 transition_evaluator: Option<Arc<dyn TransitionEvaluator>>,
605 context_manager: Arc<ContextManager>,
606 template_renderer: TemplateRenderer,
607 tool_call_history: RwLock<Vec<ToolCallRecord>>,
608 parallel_tools: ParallelToolsConfig,
609 streaming: StreamingConfig,
610 hooks: Arc<dyn AgentHooks>,
611 hitl_engine: Option<HITLEngine>,
612 approval_handler: Arc<dyn ApprovalHandler>,
613 storage_config: StorageConfig,
614 storage: RwLock<Option<Arc<dyn AgentStorage>>>,
615 storage_init: tokio::sync::Mutex<()>,
616 reasoning_config: ReasoningConfig,
617 reflection_config: ReflectionConfig,
618 disambiguation_manager: Option<DisambiguationManager>,
619 disambiguation_epoch: AtomicU64,
621 disambiguation_admission: tokio::sync::RwLock<()>,
623 state_transition_reserved: AtomicBool,
625 persona_manager: Option<Arc<ai_agents_persona::PersonaManager>>,
627 pending_skill_id: RwLock<Option<String>>,
631 current_plan: RwLock<Option<Plan>>,
632 declared_tool_ids: Option<Vec<String>>,
634 context_initialized: AtomicBool,
636 spawner: Option<Arc<crate::spawner::AgentSpawner>>,
638 spawner_registry: Option<Arc<crate::spawner::AgentRegistry>>,
640 redispatch_depth: RwLock<u32>,
643 active_turn_context: RwLock<Option<TurnOptimizationContext>>,
645 root_user_message_committed: AtomicBool,
647 active_native_exchanges: RwLock<Vec<ActiveNativeExchange>>,
649 actor_id: RwLock<Option<String>>,
651 fact_store: RwLock<Option<Arc<ai_agents_facts::FactStore>>>,
653 fact_extractor: RwLock<Option<Arc<dyn ai_agents_facts::FactExtractor>>>,
656 actor_facts_cache: Arc<RwLock<HashMap<String, Vec<ai_agents_core::KeyFact>>>>,
658 messages_since_extraction: Arc<RwLock<usize>>,
660 actor_memory_config: Option<ai_agents_facts::ActorMemoryConfig>,
662 facts_config: Option<ai_agents_facts::FactsConfig>,
664 session_metadata: RwLock<ai_agents_core::SessionMetadata>,
666 current_session_id: RwLock<Option<String>>,
668 relationship_manager: Option<Arc<RelationshipManager>>,
670 observability_manager: Option<Arc<ObservabilityManager>>,
672 runtime_config: RuntimeConfig,
674 background_maintenance: Arc<BackgroundMaintenanceQueue>,
676 resource_locks: ToolResourceLocks,
678 runtime_control: Arc<RuntimeControlState>,
680 root_turn_gate: RootTurnGate,
682}
683
684impl std::fmt::Debug for RuntimeAgent {
685 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
686 f.debug_struct("RuntimeAgent")
687 .field("info", &self.info)
688 .field("base_system_prompt", &self.base_system_prompt)
689 .field("max_iterations", &self.max_iterations)
690 .field("skills_count", &self.skills.len())
691 .field("max_context_tokens", &self.max_context_tokens)
692 .field("has_state_machine", &self.state_machine.is_some())
693 .field("parallel_tools", &self.parallel_tools)
694 .field("streaming", &self.streaming)
695 .field("has_hooks", &true)
696 .field("has_hitl", &self.hitl_engine.is_some())
697 .field("storage_type", &self.storage_config.storage_type())
698 .field("reasoning_mode", &self.reasoning_config.mode)
699 .field("reflection_enabled", &self.reflection_config.enabled)
700 .field("declared_tool_ids", &self.declared_tool_ids)
701 .field("has_persona", &self.persona_manager.is_some())
702 .field("has_observability", &self.observability_manager.is_some())
703 .finish_non_exhaustive()
704 }
705}
706
707struct ObservabilityClarificationObserver;
708
709impl ClarificationObserver for ObservabilityClarificationObserver {
710 fn observe_question<'a>(
712 &'a self,
713 future: ClarificationQuestionFuture<'a>,
714 ) -> ClarificationQuestionFuture<'a> {
715 Box::pin(async move {
716 with_observation_purpose(ObservationPurpose::DisambiguationClarification, future).await
717 })
718 }
719
720 fn observe_parse<'a>(
722 &'a self,
723 future: ClarificationParseFuture<'a>,
724 ) -> ClarificationParseFuture<'a> {
725 Box::pin(async move {
726 with_observation_purpose(ObservationPurpose::DisambiguationClarification, future).await
727 })
728 }
729
730 fn observe_confirmation_parse<'a>(
732 &'a self,
733 future: ConfirmationParseFuture<'a>,
734 ) -> ConfirmationParseFuture<'a> {
735 Box::pin(async move {
736 with_observation_purpose(ObservationPurpose::DisambiguationClarification, future).await
737 })
738 }
739}
740
741struct ObservabilityProcessStageObserver;
742
743impl ProcessStageObserver for ObservabilityProcessStageObserver {
744 fn observe<'a>(
746 &'a self,
747 hint: ProcessPurposeHint,
748 future: ProcessStageFuture<'a>,
749 ) -> ProcessStageFuture<'a> {
750 Box::pin(async move {
751 with_observation_purpose(observation_purpose_for_process(hint), future).await
752 })
753 }
754}
755
756struct RegistryLLMGetter {
757 registry: Arc<LLMRegistry>,
758}
759
760impl LLMGetter for RegistryLLMGetter {
761 fn get_condition_llm(&self, alias: Option<&str>) -> Result<Option<Arc<dyn LLMProvider>>> {
763 Ok(
764 match self
765 .registry
766 .resolve_role_override(ai_agents_llm::LLMRole::ToolsCondition, alias)
767 .map_err(|e| AgentError::Config(e.to_string()))?
768 {
769 Some(resolved) => Some(resolved.provider),
770 None => self.get_llm(alias.unwrap_or("router")),
771 },
772 )
773 }
774 fn get_llm(&self, alias: &str) -> Option<Arc<dyn LLMProvider>> {
775 self.registry.get(alias).ok()
776 }
777}
778
779impl RuntimeAgent {
780 fn role_llm<F>(
782 &self,
783 role: ai_agents_llm::LLMRole,
784 local: Option<&str>,
785 legacy: F,
786 ) -> Result<Arc<dyn LLMProvider>>
787 where
788 F: FnOnce() -> Result<Arc<dyn LLMProvider>>,
789 {
790 match self
791 .llm_registry
792 .resolve_role_override(role, local)
793 .map_err(|e| AgentError::Config(e.to_string()))?
794 {
795 Some(resolved) => Ok(resolved.provider),
796 None => legacy(),
797 }
798 }
799
800 fn optional_role_llm<F>(
802 &self,
803 role: ai_agents_llm::LLMRole,
804 local: Option<&str>,
805 legacy: F,
806 ) -> Result<Option<Arc<dyn LLMProvider>>>
807 where
808 F: FnOnce() -> Option<Arc<dyn LLMProvider>>,
809 {
810 match self
811 .llm_registry
812 .resolve_role_override(role, local)
813 .map_err(|e| AgentError::Config(e.to_string()))?
814 {
815 Some(resolved) => Ok(Some(resolved.provider)),
816 None => Ok(legacy()),
817 }
818 }
819 #[allow(clippy::too_many_arguments)]
821 pub fn new(
822 info: AgentInfo,
823 llm_registry: Arc<LLMRegistry>,
824 memory: Arc<dyn Memory>,
825 tools: Arc<ToolRegistry>,
826 skills: Vec<SkillDefinition>,
827 system_prompt: String,
828 max_iterations: u32,
829 ) -> Self {
830 Self::try_new(
831 info,
832 llm_registry,
833 memory,
834 tools,
835 skills,
836 system_prompt,
837 max_iterations,
838 )
839 .expect("Invalid hierarchical runtime routing; use RuntimeAgent::try_new")
840 }
841
842 #[allow(clippy::too_many_arguments)]
844 pub fn try_new(
845 info: AgentInfo,
846 llm_registry: Arc<LLMRegistry>,
847 memory: Arc<dyn Memory>,
848 tools: Arc<ToolRegistry>,
849 skills: Vec<SkillDefinition>,
850 system_prompt: String,
851 max_iterations: u32,
852 ) -> Result<Self> {
853 llm_registry
854 .validate_router_roles()
855 .map_err(|error| AgentError::Config(error.to_string()))?;
856 let (skill_router, skill_executor) = if !skills.is_empty() {
857 let router_llm = match llm_registry
858 .resolve_role_override(ai_agents_llm::LLMRole::SkillsSelection, None)
859 .map_err(|e| AgentError::Config(e.to_string()))?
860 {
861 Some(resolved) => Some(resolved.provider),
862 None => llm_registry.router().ok(),
863 };
864 let router = router_llm.map(|llm| SkillRouter::new(llm, skills.clone()));
865 let executor = SkillExecutor::new(llm_registry.clone(), tools.clone());
866 (router, Some(executor))
867 } else {
868 (None, None)
869 };
870
871 let context_manager =
872 ContextManager::new(HashMap::new(), info.name.clone(), info.version.clone());
873
874 Ok(Self {
875 info,
876 llm_registry,
877 memory,
878 tools,
879 skills,
880 skill_router,
881 skill_executor,
882 base_system_prompt: system_prompt,
883 max_iterations,
884 iteration_count: RwLock::new(0),
885 max_context_tokens: 128000,
886 memory_token_budget: None,
887 recovery_manager: RecoveryManager::default(),
888 tool_security: ToolSecurityEngine::default(),
889 process_processor: None,
890 message_filters: RwLock::new(HashMap::new()),
891 state_machine: None,
892 transition_evaluator: None,
893 context_manager: Arc::new(context_manager),
894 template_renderer: TemplateRenderer::new(),
895 tool_call_history: RwLock::new(Vec::new()),
896 parallel_tools: ParallelToolsConfig::default(),
897 streaming: StreamingConfig::default(),
898 hooks: Arc::new(NoopHooks),
899 hitl_engine: None,
900 approval_handler: Arc::new(RejectAllHandler::new()),
901 storage_config: StorageConfig::default(),
902 storage: RwLock::new(None),
903 storage_init: tokio::sync::Mutex::new(()),
904 reasoning_config: ReasoningConfig::default(),
905 reflection_config: ReflectionConfig::default(),
906 disambiguation_manager: None,
907 disambiguation_epoch: AtomicU64::new(0),
908 disambiguation_admission: tokio::sync::RwLock::new(()),
909 state_transition_reserved: AtomicBool::new(false),
910 persona_manager: None,
911 pending_skill_id: RwLock::new(None),
912 current_plan: RwLock::new(None),
913 declared_tool_ids: None,
914 context_initialized: AtomicBool::new(false),
915 spawner: None,
916 spawner_registry: None,
917 redispatch_depth: RwLock::new(0),
918 active_turn_context: RwLock::new(None),
919 root_user_message_committed: AtomicBool::new(false),
920 active_native_exchanges: RwLock::new(Vec::new()),
921 actor_id: RwLock::new(None),
922 fact_store: RwLock::new(None),
923 fact_extractor: RwLock::new(None),
924 actor_facts_cache: Arc::new(RwLock::new(HashMap::new())),
925 messages_since_extraction: Arc::new(RwLock::new(0)),
926 actor_memory_config: None,
927 facts_config: None,
928 session_metadata: RwLock::new(ai_agents_core::SessionMetadata::default()),
929 current_session_id: RwLock::new(None),
930 relationship_manager: None,
931 observability_manager: None,
932 runtime_config: RuntimeConfig::default(),
933 background_maintenance: Arc::new(BackgroundMaintenanceQueue::default()),
934 resource_locks: new_tool_resource_locks(),
935 runtime_control: Arc::new(RuntimeControlState::default()),
936 root_turn_gate: Arc::new(tokio::sync::Mutex::new(())),
937 })
938 }
939
940 pub fn with_declared_tool_ids(mut self, ids: Option<Vec<String>>) -> Self {
941 self.declared_tool_ids = ids;
942 self
943 }
944
945 pub fn with_storage_config(mut self, config: StorageConfig) -> Self {
946 self.storage_config = config;
947 self
948 }
949
950 pub fn with_storage(self, storage: Arc<dyn AgentStorage>) -> Self {
951 *self.storage.write() = Some(storage);
952 self
953 }
954
955 pub(crate) fn with_shared_resource_locks(mut self, locks: ToolResourceLocks) -> Self {
956 self.resource_locks = locks;
957 self
958 }
959
960 pub fn with_reasoning(mut self, config: ReasoningConfig) -> Self {
961 self.reasoning_config = config;
962 self
963 }
964
965 pub fn with_reflection(mut self, config: ReflectionConfig) -> Self {
966 self.reflection_config = config;
967 self
968 }
969
970 pub fn with_relationships(mut self, manager: Arc<RelationshipManager>) -> Self {
972 self.relationship_manager = Some(manager);
973 self
974 }
975
976 pub fn with_observability(mut self, manager: Arc<ObservabilityManager>) -> Self {
978 self.observability_manager = Some(manager);
979 self
980 }
981
982 pub fn with_runtime_config(mut self, config: RuntimeConfig) -> Self {
984 let max_tasks = config.optimization.post_turn.max_background_tasks;
985 self.background_maintenance = Arc::new(BackgroundMaintenanceQueue::new(max_tasks));
986 self.runtime_config = config;
987 self
988 }
989
990 pub fn runtime_config(&self) -> &RuntimeConfig {
992 &self.runtime_config
993 }
994
995 pub async fn flush_background_tasks(&self) -> Result<()> {
997 self.background_maintenance.flush_all().await
998 }
999
1000 pub async fn flush_background_tasks_for_actor(&self, actor_id: &str) -> Result<()> {
1002 self.background_maintenance.flush_scope(actor_id).await
1003 }
1004
1005 pub async fn flush_background_tasks_for_purpose(
1007 &self,
1008 purpose: RuntimeTaskPurpose,
1009 ) -> Result<()> {
1010 self.background_maintenance.flush_purpose(purpose).await
1011 }
1012
1013 pub async fn flush_background_tasks_for_actor_purpose(
1015 &self,
1016 actor_id: &str,
1017 purpose: RuntimeTaskPurpose,
1018 ) -> Result<()> {
1019 self.background_maintenance
1020 .flush_scope_purpose(actor_id, purpose)
1021 .await
1022 }
1023
1024 pub async fn shutdown_background_tasks(&self) -> Result<()> {
1026 self.flush_background_tasks().await
1027 }
1028
1029 pub fn observability(&self) -> Option<Arc<ObservabilityManager>> {
1031 self.observability_manager.clone()
1032 }
1033
1034 async fn export_observability_if_configured(&self) {
1036 let Some(manager) = self.observability_manager.as_ref() else {
1037 return;
1038 };
1039 let export = &manager.config().export;
1040 if !export.write_report && !export.write_raw_events {
1041 return;
1042 }
1043 if let Err(error) = manager.export().await {
1044 warn!(error = %error, "Observability export failed");
1045 }
1046 }
1047
1048 pub fn relationship_manager(&self) -> Option<Arc<RelationshipManager>> {
1050 self.relationship_manager.clone()
1051 }
1052
1053 fn current_turn_actor_context(&self) -> Option<crate::TurnActorContext> {
1054 current_turn_actor_context()
1055 }
1056
1057 fn effective_actor_id(&self) -> Option<String> {
1058 self.current_turn_actor_context()
1059 .and_then(|ctx| ctx.effective_actor_id().map(|id| id.to_string()))
1060 .or_else(|| self.actor_id.read().clone())
1061 }
1062
1063 fn effective_origin_actor_id(&self) -> Option<String> {
1064 self.current_turn_actor_context()
1065 .and_then(|ctx| ctx.origin_actor_id.clone())
1066 .or_else(|| self.actor_id.read().clone())
1067 }
1068
1069 fn record_session_actor_if_needed(&self) {
1070 if let Some(actor_id) = self.effective_origin_actor_id() {
1071 let mut meta = self.session_metadata.write();
1072 meta.actor_id = Some(actor_id.clone());
1073 if !meta.actors.iter().any(|a| a == &actor_id) {
1074 meta.actors.push(actor_id);
1075 }
1076 }
1077 }
1078
1079 fn outbound_actor_context(&self) -> crate::TurnActorContext {
1080 let mut context = self.current_turn_actor_context().unwrap_or_default();
1081 if context.origin_actor_id.is_none() {
1082 context.origin_actor_id = self.effective_origin_actor_id();
1083 }
1084 context.sender_agent_id = Some(self.info.id.clone());
1085 context
1086 }
1087
1088 fn observation_session_id(&self) -> Option<String> {
1090 let mut current = self.current_session_id.write();
1091 if current.is_none() {
1092 *current = Some(new_observation_session_id());
1093 }
1094 current.clone()
1095 }
1096
1097 fn build_observation_context(&self, actor_id: Option<String>) -> Option<SpanContext> {
1099 let manager = self.observability_manager.as_ref()?;
1100 let context = self.build_context_with_overlays();
1101 let language = resolve_language_from_context(manager.config(), &context);
1102 let context = current_observation_context()
1103 .map(|parent| parent.child_for_agent(self.info.id.clone()).with_new_turn())
1104 .unwrap_or_else(|| SpanContext::new_root(self.info.id.clone()));
1105 Some(
1106 context
1107 .with_actor(actor_id.or_else(|| self.effective_actor_id()))
1108 .with_session(self.observation_session_id())
1109 .with_state(self.current_state())
1110 .with_language(Some(language)),
1111 )
1112 }
1113
1114 fn current_runtime_observation_context(
1116 &self,
1117 purpose: ObservationPurpose,
1118 ) -> Option<SpanContext> {
1119 let manager = self.observability_manager.as_ref()?;
1120 let context = self.build_context_with_overlays();
1121 let language = resolve_language_from_context(manager.config(), &context);
1122 let mut observation = current_observation_context()
1123 .unwrap_or_else(|| SpanContext::new_root(self.info.id.clone()));
1124 observation.agent_id = self.info.id.clone();
1125 observation.actor_id = self.effective_actor_id();
1126 observation.session_id = self.observation_session_id();
1127 observation.state = self.current_state();
1128 observation.language = Some(language);
1129 observation.purpose = purpose;
1130 Some(observation)
1131 }
1132
1133 async fn observe_purpose<F, T>(&self, purpose: ObservationPurpose, future: F) -> T
1135 where
1136 F: Future<Output = T>,
1137 {
1138 if let Some(context) = self.current_runtime_observation_context(purpose) {
1139 with_observation_context(context, future).await
1140 } else {
1141 future.await
1142 }
1143 }
1144
1145 fn chat_with_actor_context_boxed<'a>(
1149 &'a self,
1150 input: &'a str,
1151 actor_context: crate::TurnActorContext,
1152 ) -> Pin<Box<dyn Future<Output = Result<AgentResponse>> + Send + 'a>> {
1153 Box::pin(async move {
1154 let RootTurnAdmission {
1155 guard,
1156 identity_stack,
1157 } = self.acquire_root_turn().await?;
1158 let result = scope_runtime_gate_identity_stack(&identity_stack, async move {
1159 let actor_id = actor_context.effective_actor_id().map(str::to_string);
1160 let run = async move {
1161 scope_actor_context(
1162 actor_context,
1163 Box::pin(async move { self.run_loop(input).await }),
1164 )
1165 .await
1166 };
1167 let result = if let Some(context) = self.build_observation_context(actor_id) {
1168 with_observation_context(context, run).await
1169 } else {
1170 run.await
1171 };
1172 self.export_observability_if_configured().await;
1173 result
1174 })
1175 .await;
1176 drop(guard);
1177 result
1178 })
1179 }
1180
1181 async fn acquire_root_turn(&self) -> Result<RootTurnAdmission> {
1183 let gate_identity = Arc::clone(&self.root_turn_gate);
1184 let current_identity_stack = current_runtime_gate_identity_stack();
1185 if current_identity_stack
1186 .iter()
1187 .any(|owned_gate| Arc::ptr_eq(owned_gate, &gate_identity))
1188 {
1189 return Err(AgentError::Other(format!(
1190 "RuntimeAgent '{}' rejected reentrant root turn ownership",
1191 self.info.id
1192 )));
1193 }
1194 let guard = Arc::clone(&gate_identity).lock_owned().await;
1195 let mut identity_stack = Vec::with_capacity(current_identity_stack.len() + 1);
1199 identity_stack.extend(current_identity_stack.iter().cloned());
1200 identity_stack.push(gate_identity);
1201 Ok(RootTurnAdmission {
1202 guard,
1203 identity_stack: identity_stack.into(),
1204 })
1205 }
1206
1207 pub async fn chat_with_actor_context(
1211 &self,
1212 input: &str,
1213 actor_context: crate::TurnActorContext,
1214 ) -> Result<AgentResponse> {
1215 self.chat_with_actor_context_boxed(input, actor_context)
1216 .await
1217 }
1218
1219 pub async fn chat_as_actor(&self, actor_id: &str, input: &str) -> Result<AgentResponse> {
1221 let actor_context = crate::TurnActorContext::new().with_origin_actor(actor_id);
1222 self.chat_with_actor_context(input, actor_context).await
1223 }
1224
1225 pub async fn load_actor_relationship(&self) -> Result<()> {
1227 self.maybe_load_actor_relationship().await;
1228 Ok(())
1229 }
1230
1231 pub async fn update_relationship_dimension(
1233 &self,
1234 dimension: &str,
1235 delta: f64,
1236 reason: Option<&str>,
1237 ) -> Result<ai_agents_relationships::DimensionChange> {
1238 self.update_relationship_dimension_for_perspective(
1239 ai_agents_relationships::RelationshipPerspective::AgentToActor,
1240 dimension,
1241 delta,
1242 reason,
1243 )
1244 .await
1245 }
1246
1247 pub async fn update_relationship_dimension_for_perspective(
1251 &self,
1252 perspective: ai_agents_relationships::RelationshipPerspective,
1253 dimension: &str,
1254 delta: f64,
1255 reason: Option<&str>,
1256 ) -> Result<ai_agents_relationships::DimensionChange> {
1257 let manager = self
1258 .relationship_manager
1259 .as_ref()
1260 .ok_or_else(|| AgentError::Config("Relationship memory is not configured".into()))?;
1261 let actor_id = self.effective_actor_id().ok_or_else(|| {
1262 AgentError::Config("No actor ID set. Use set_actor_id() first".into())
1263 })?;
1264 let change = manager.update_dimension_for_perspective(
1265 &actor_id,
1266 perspective,
1267 dimension,
1268 delta,
1269 1.0,
1270 reason.unwrap_or("manual relationship update"),
1271 )?;
1272 self.persist_actor_relationship(&actor_id).await?;
1273 info!(
1274 actor_id = %actor_id,
1275 perspective = %change.perspective,
1276 dimension = %change.dimension,
1277 delta = change.delta,
1278 current = change.current,
1279 "relationship updated manually"
1280 );
1281 self.hooks
1282 .on_relationship_change(&actor_id, std::slice::from_ref(&change))
1283 .await;
1284 Ok(change)
1285 }
1286
1287 pub fn reasoning_config(&self) -> &ReasoningConfig {
1288 &self.reasoning_config
1289 }
1290
1291 pub fn reflection_config(&self) -> &ReflectionConfig {
1292 &self.reflection_config
1293 }
1294
1295 pub fn with_facts_config(
1298 mut self,
1299 actor_memory_config: Option<ai_agents_facts::ActorMemoryConfig>,
1300 facts_config: Option<ai_agents_facts::FactsConfig>,
1301 ) -> Self {
1302 self.actor_memory_config = actor_memory_config;
1303 self.facts_config = facts_config;
1304 self
1305 }
1306
1307 pub fn with_facts(
1310 mut self,
1311 store: Arc<ai_agents_facts::FactStore>,
1312 extractor: Option<Arc<dyn ai_agents_facts::FactExtractor>>,
1313 actor_memory_config: Option<ai_agents_facts::ActorMemoryConfig>,
1314 facts_config: Option<ai_agents_facts::FactsConfig>,
1315 ) -> Self {
1316 *self.fact_store.write() = Some(store);
1317 *self.fact_extractor.write() = extractor;
1318 self.actor_memory_config = actor_memory_config;
1319 self.facts_config = facts_config;
1320 self
1321 }
1322
1323 pub fn fact_store(&self) -> Option<Arc<ai_agents_facts::FactStore>> {
1325 self.fact_store.read().clone()
1326 }
1327
1328 pub fn actor_id(&self) -> Option<String> {
1330 self.actor_id.read().clone()
1331 }
1332
1333 pub fn set_actor_id(&self, actor_id: &str) -> ai_agents_core::Result<()> {
1335 *self.actor_id.write() = Some(actor_id.to_string());
1336 {
1337 let mut meta = self.session_metadata.write();
1338 meta.actor_id = Some(actor_id.to_string());
1339 if !meta.actors.iter().any(|a| a == actor_id) {
1340 meta.actors.push(actor_id.to_string());
1341 }
1342 }
1343 Ok(())
1344 }
1345
1346 pub fn clear_actor_id(&self) {
1348 *self.actor_id.write() = None;
1349 self.session_metadata.write().actor_id = None;
1350 }
1351
1352 pub fn set_user_id(&self, user_id: &str) -> ai_agents_core::Result<()> {
1354 self.set_actor_id(user_id)
1355 }
1356
1357 pub async fn load_actor_memory(&self) -> ai_agents_core::Result<()> {
1359 let actor_id = match self.effective_actor_id() {
1360 Some(id) => id,
1361 None => return Ok(()),
1362 };
1363
1364 let store_opt = self.fact_store.read().clone();
1365 if let Some(store) = store_opt {
1366 let facts = store.get_facts(&actor_id).await?;
1367 let count = facts.len();
1368 self.actor_facts_cache
1369 .write()
1370 .insert(actor_id.clone(), facts);
1371 self.hooks.on_actor_memory_loaded(&actor_id, count).await;
1372 tracing::debug!("loaded {} facts for actor {}", count, actor_id);
1373 }
1374
1375 Ok(())
1376 }
1377
1378 async fn maybe_load_actor_memory(&self) {
1380 let Some(actor_id) = self.effective_actor_id() else {
1381 return;
1382 };
1383 if self.actor_facts_cache.read().contains_key(&actor_id) {
1384 return;
1385 }
1386 let _ = self.load_actor_memory().await;
1387 }
1388
1389 async fn pre_turn_session_lifecycle(&self) {
1391 if *self.redispatch_depth.read() > 0 {
1392 return;
1393 }
1394 self.resolve_actor_id_from_context();
1395 self.await_background_before_next_turn().await;
1396 self.record_session_actor_if_needed();
1397 self.maybe_load_actor_memory().await;
1398 self.maybe_load_actor_relationship().await;
1399 *self.messages_since_extraction.write() += 1;
1400 }
1401
1402 async fn post_turn_session_lifecycle(&self) -> Result<()> {
1404 if *self.redispatch_depth.read() > 0 {
1405 return Ok(());
1406 }
1407 *self.messages_since_extraction.write() += 1;
1408 self.run_post_turn_maintenance().await
1409 }
1410
1411 fn begin_root_turn(&self) {
1413 if *self.redispatch_depth.read() == 0 {
1414 let mut guard = self.active_turn_context.write();
1415 if guard.is_none() {
1416 self.root_user_message_committed
1417 .store(false, Ordering::SeqCst);
1418 self.active_native_exchanges.write().clear();
1419 let max_calls = self
1420 .runtime_config
1421 .optimization
1422 .max_speculative_llm_calls_per_turn;
1423 *guard = Some(TurnOptimizationContext::new(
1424 String::new(),
1425 HashMap::new(),
1426 max_calls,
1427 ));
1428 }
1429 }
1430 }
1431
1432 fn update_active_turn_context(
1433 &self,
1434 processed_input: &str,
1435 input_context: HashMap<String, Value>,
1436 ) {
1437 if *self.redispatch_depth.read() > 0 {
1438 return;
1439 }
1440 let max_calls = self
1441 .runtime_config
1442 .optimization
1443 .max_speculative_llm_calls_per_turn;
1444 let mut guard = self.active_turn_context.write();
1445 match guard.as_mut() {
1446 Some(context) => {
1447 context.processed_input = processed_input.to_string();
1448 context.input_context = input_context;
1449 context.max_speculative_llm_calls = max_calls;
1450 }
1451 None => {
1452 *guard = Some(TurnOptimizationContext::new(
1453 processed_input,
1454 input_context,
1455 max_calls,
1456 ));
1457 }
1458 }
1459 }
1460
1461 async fn commit_root_user_message(&self, processed_input: &str) -> Result<()> {
1463 if *self.redispatch_depth.read() > 0 {
1464 return Ok(());
1465 }
1466 if !self
1467 .root_user_message_committed
1468 .swap(true, Ordering::SeqCst)
1469 {
1470 self.memory
1471 .add_message(ChatMessage::user(processed_input))
1472 .await?;
1473 if let Some(context) = self.active_turn_context.write().as_mut() {
1474 context.mark_user_message_committed();
1475 }
1476 }
1477 Ok(())
1478 }
1479
1480 fn end_root_turn(&self) {
1482 if *self.redispatch_depth.read() == 0 {
1483 self.root_user_message_committed
1484 .store(false, Ordering::SeqCst);
1485 *self.active_turn_context.write() = None;
1486 self.active_native_exchanges.write().clear();
1487 }
1488 }
1489
1490 fn reserve_active_speculative_llm_call(&self, kind: RuntimeOptimizationKind) -> bool {
1491 self.begin_root_turn();
1492 let mut guard = self.active_turn_context.write();
1493 let Some(context) = guard.as_mut() else {
1494 return false;
1495 };
1496 context.reserve_speculative_llm_call_for(kind)
1497 }
1498
1499 fn branch_context_preview(&self) -> String {
1500 let context = self.build_context_with_overlays();
1501 let mut value = serde_json::to_string_pretty(&context).unwrap_or_else(|_| "{}".to_string());
1502 const MAX_CONTEXT_PREVIEW_CHARS: usize = 2048;
1503 if value.chars().count() > MAX_CONTEXT_PREVIEW_CHARS {
1504 value = value
1505 .chars()
1506 .take(MAX_CONTEXT_PREVIEW_CHARS)
1507 .collect::<String>();
1508 value.push_str("...");
1509 }
1510 value
1511 }
1512
1513 async fn await_background_before_next_turn(&self) {
1515 let optimization = &self.runtime_config.optimization;
1516 if !optimization.enabled {
1517 return;
1518 }
1519 let actor_id = self.effective_actor_id();
1520 let post = &optimization.post_turn;
1521 self.await_background_task(
1522 post.facts.await_before_next_turn,
1523 RuntimeTaskPurpose::PostTurnFacts,
1524 actor_id.as_deref(),
1525 "facts",
1526 )
1527 .await;
1528 self.await_background_task(
1529 post.relationships.await_before_next_turn,
1530 RuntimeTaskPurpose::PostTurnRelationship,
1531 actor_id.as_deref(),
1532 "relationships",
1533 )
1534 .await;
1535 }
1536
1537 async fn await_background_task(
1538 &self,
1539 policy: AwaitBeforeNextTurn,
1540 purpose: RuntimeTaskPurpose,
1541 actor_id: Option<&str>,
1542 label: &str,
1543 ) {
1544 match policy {
1545 AwaitBeforeNextTurn::Never => {}
1546 AwaitBeforeNextTurn::Always => {
1547 if let Err(error) = self.flush_background_tasks_for_purpose(purpose).await {
1548 warn!(label = label, error = %error, "background maintenance flush failed");
1549 }
1550 }
1551 AwaitBeforeNextTurn::SameActor => {
1552 if let Some(actor_id) = actor_id
1553 && let Err(error) = self
1554 .flush_background_tasks_for_actor_purpose(actor_id, purpose)
1555 .await
1556 {
1557 warn!(label = label, actor_id = %actor_id, error = %error, "actor background maintenance flush failed");
1558 }
1559 }
1560 }
1561 }
1562
1563 async fn run_post_turn_maintenance(&self) -> Result<()> {
1565 let optimization = &self.runtime_config.optimization;
1566 if !optimization.enabled {
1567 self.auto_extract_facts().await;
1568 self.auto_update_relationship().await;
1569 return Ok(());
1570 }
1571
1572 let facts_mode = effective_maintenance_mode(
1573 optimization.post_turn.facts.mode,
1574 optimization.parallel_post_turn_memory,
1575 );
1576 let relationships_mode = effective_maintenance_mode(
1577 optimization.post_turn.relationships.mode,
1578 optimization.parallel_post_turn_memory,
1579 );
1580
1581 match (facts_mode, relationships_mode) {
1582 (MaintenanceMode::InlineSerial, MaintenanceMode::InlineSerial) => {
1583 self.auto_extract_facts().await;
1584 self.auto_update_relationship().await;
1585 }
1586 (MaintenanceMode::InlineParallel, MaintenanceMode::InlineParallel) => {
1587 let facts = self.auto_extract_facts();
1588 let relationships = self.auto_update_relationship();
1589 tokio::join!(facts, relationships);
1590 }
1591 (MaintenanceMode::Background, MaintenanceMode::Background) => {
1592 self.schedule_facts_background().await?;
1593 self.schedule_relationship_background().await?;
1594 }
1595 (MaintenanceMode::Background, MaintenanceMode::InlineParallel)
1596 | (MaintenanceMode::Background, MaintenanceMode::InlineSerial) => {
1597 self.schedule_facts_background().await?;
1598 self.auto_update_relationship().await;
1599 }
1600 (MaintenanceMode::InlineParallel, MaintenanceMode::Background)
1601 | (MaintenanceMode::InlineSerial, MaintenanceMode::Background) => {
1602 self.auto_extract_facts().await;
1603 self.schedule_relationship_background().await?;
1604 }
1605 _ => {
1606 self.auto_extract_facts().await;
1607 self.auto_update_relationship().await;
1608 }
1609 }
1610 Ok(())
1611 }
1612
1613 async fn schedule_facts_background(&self) -> Result<()> {
1614 let policy = self.runtime_config.optimization.post_turn.facts.clone();
1615 let should_extract = self
1616 .facts_config
1617 .as_ref()
1618 .map(|c| c.enabled && c.auto_extract)
1619 .unwrap_or(false);
1620 if !should_extract {
1621 return Ok(());
1622 }
1623 let msgs_since = *self.messages_since_extraction.read();
1624 if msgs_since < 2 {
1625 return Ok(());
1626 }
1627 let Some(actor_id) = self.effective_actor_id() else {
1628 self.record_skipped_maintenance(
1629 "facts",
1630 ObservationPurpose::FactsExtraction,
1631 "missing_actor",
1632 Some(&policy),
1633 );
1634 return Ok(());
1635 };
1636 let Some(extractor) = self.fact_extractor.read().clone() else {
1637 return Ok(());
1638 };
1639 let messages = match self.memory.get_messages(None).await {
1640 Ok(messages) => messages,
1641 Err(error) => {
1642 warn!(error = %error, "failed to snapshot messages for fact extraction");
1643 return Ok(());
1644 }
1645 };
1646 let messages = Self::readable_native_messages(messages)?;
1647 let recent: Vec<_> = messages
1648 .iter()
1649 .rev()
1650 .take(msgs_since)
1651 .rev()
1652 .cloned()
1653 .collect();
1654 if recent.is_empty() {
1655 return Ok(());
1656 }
1657 let existing = self
1658 .actor_facts_cache
1659 .read()
1660 .get(&actor_id)
1661 .cloned()
1662 .unwrap_or_default();
1663 let categories = self
1664 .facts_config
1665 .as_ref()
1666 .map(|c| c.custom_categories.clone())
1667 .unwrap_or_default();
1668 let store = self.fact_store.read().clone();
1669 let cache = Arc::clone(&self.actor_facts_cache);
1670 let counter = Arc::clone(&self.messages_since_extraction);
1671 let hooks = Arc::clone(&self.hooks);
1672 let agent_id = self.info.id.clone();
1673 let observation = current_observation_context();
1674 let key = MaintenanceSequenceKey::actor(
1675 agent_id,
1676 actor_id.clone(),
1677 RuntimeTaskPurpose::PostTurnFacts,
1678 );
1679 let actor_for_task = actor_id.clone();
1680 let task = async move {
1681 let run = async move {
1682 let facts = extractor
1683 .extract(&recent, &existing, Some(&actor_for_task), &categories)
1684 .await?;
1685 if !facts.is_empty() {
1686 if let Some(store) = store {
1687 let authoritative = store.add_facts(&actor_for_task, facts.clone()).await?;
1688 cache.write().insert(actor_for_task.clone(), authoritative);
1689 } else {
1690 cache
1691 .write()
1692 .entry(actor_for_task.clone())
1693 .or_default()
1694 .extend(facts.clone());
1695 }
1696 {
1697 let mut count = counter.write();
1698 if *count <= msgs_since {
1699 *count = 0;
1700 } else {
1701 *count -= msgs_since;
1702 }
1703 }
1704 hooks.on_facts_extracted(&actor_for_task, &facts).await;
1705 }
1706 Ok(())
1707 };
1708 if let Some(context) = observation {
1709 with_observation_context(
1710 context.with_purpose(ObservationPurpose::FactsExtraction),
1711 run,
1712 )
1713 .await
1714 } else {
1715 run.await
1716 }
1717 };
1718 self.spawn_or_handle_background(Some(key), task, "facts", &policy)
1719 .await
1720 }
1721
1722 async fn schedule_relationship_background(&self) -> Result<()> {
1723 let policy = self
1724 .runtime_config
1725 .optimization
1726 .post_turn
1727 .relationships
1728 .clone();
1729 let Some(manager) = self.relationship_manager.as_ref().cloned() else {
1730 return Ok(());
1731 };
1732 let Some(actor_id) = self.effective_actor_id() else {
1733 self.record_skipped_maintenance(
1734 "relationships",
1735 ObservationPurpose::RelationshipUpdate,
1736 "missing_actor",
1737 Some(&policy),
1738 );
1739 return Ok(());
1740 };
1741 let recent_messages = manager.config().auto_update.recent_messages;
1742 let messages = match self.memory.get_messages(Some(recent_messages)).await {
1743 Ok(messages) => messages,
1744 Err(error) => {
1745 warn!(actor = %actor_id, error = %error, "failed to snapshot messages for relationship update");
1746 return Ok(());
1747 }
1748 };
1749 let messages = Self::readable_native_messages(messages)?;
1750 let storage = self.storage.read().clone();
1751 let hooks = Arc::clone(&self.hooks);
1752 let agent_id = self.info.id.clone();
1753 let observation = current_observation_context();
1754 let key = MaintenanceSequenceKey::actor(
1755 agent_id.clone(),
1756 actor_id.clone(),
1757 RuntimeTaskPurpose::PostTurnRelationship,
1758 );
1759 let actor_for_task = actor_id.clone();
1760 let task = async move {
1761 let run = async move {
1762 if manager.config().auto_update.enabled {
1763 let update = manager.auto_update(&actor_for_task, &messages).await?;
1764 if !update.changes.is_empty() {
1765 hooks
1766 .on_relationship_change(&actor_for_task, &update.changes)
1767 .await;
1768 }
1769 if let Some(ref event) = update.event {
1770 hooks.on_notable_event(&actor_for_task, event).await;
1771 }
1772 }
1773 if manager.config().persistence.enabled
1774 && let (Some(storage), Some(value)) =
1775 (storage, manager.relationship_as_value(&actor_for_task)?)
1776 {
1777 storage
1778 .save_relationship(&agent_id, &actor_for_task, &value)
1779 .await?;
1780 }
1781 Ok(())
1782 };
1783 if let Some(context) = observation {
1784 with_observation_context(
1785 context.with_purpose(ObservationPurpose::RelationshipUpdate),
1786 run,
1787 )
1788 .await
1789 } else {
1790 run.await
1791 }
1792 };
1793 self.spawn_or_handle_background(Some(key), task, "relationships", &policy)
1794 .await
1795 }
1796
1797 async fn spawn_or_handle_background<F>(
1799 &self,
1800 key: Option<MaintenanceSequenceKey>,
1801 task: F,
1802 label: &'static str,
1803 policy: &crate::optimization::config::MaintenanceTaskPolicy,
1804 ) -> Result<()>
1805 where
1806 F: Future<Output = Result<()>> + Send + 'static,
1807 {
1808 if self.background_maintenance.is_full() {
1809 match self
1810 .runtime_config
1811 .optimization
1812 .post_turn
1813 .on_background_overflow
1814 {
1815 BackgroundOverflowPolicy::RunInline => {
1816 record_background_maintenance_event(
1817 self.observability_manager.as_ref(),
1818 label,
1819 EventStatus::Success,
1820 0,
1821 "inline_overflow",
1822 None,
1823 Some(policy),
1824 );
1825 let start = Instant::now();
1826 match task.await {
1827 Ok(()) => record_background_maintenance_event(
1828 self.observability_manager.as_ref(),
1829 label,
1830 EventStatus::Success,
1831 start.elapsed().as_millis() as u64,
1832 "inline_completed",
1833 None,
1834 Some(policy),
1835 ),
1836 Err(error) => {
1837 warn!(label = label, error = %error, "inline maintenance fallback failed");
1838 record_background_maintenance_event(
1839 self.observability_manager.as_ref(),
1840 label,
1841 EventStatus::Error,
1842 start.elapsed().as_millis() as u64,
1843 "inline_failed",
1844 Some(error.to_string()),
1845 Some(policy),
1846 );
1847 return Err(error);
1848 }
1849 }
1850 }
1851 BackgroundOverflowPolicy::Drop => {
1852 self.record_skipped_maintenance(
1853 label,
1854 ObservationPurpose::Other(label.to_string()),
1855 "queue_full",
1856 Some(policy),
1857 );
1858 }
1859 BackgroundOverflowPolicy::Error => {
1860 record_background_maintenance_event(
1861 self.observability_manager.as_ref(),
1862 label,
1863 EventStatus::Error,
1864 0,
1865 "queue_full",
1866 None,
1867 Some(policy),
1868 );
1869 warn!(label = label, "background maintenance queue full");
1870 return Err(AgentError::Other(format!(
1871 "background maintenance queue is full for {}",
1872 label
1873 )));
1874 }
1875 }
1876 return Ok(());
1877 }
1878
1879 record_background_maintenance_event(
1880 self.observability_manager.as_ref(),
1881 label,
1882 EventStatus::Success,
1883 0,
1884 "scheduled",
1885 None,
1886 Some(policy),
1887 );
1888 let manager = self.observability_manager.clone();
1889 let policy_for_task = policy.clone();
1890 let observed_task = async move {
1891 let start = Instant::now();
1892 let result = task.await;
1893 match &result {
1894 Ok(()) => record_background_maintenance_event(
1895 manager.as_ref(),
1896 label,
1897 EventStatus::Success,
1898 start.elapsed().as_millis() as u64,
1899 "completed",
1900 None,
1901 Some(&policy_for_task),
1902 ),
1903 Err(error) => record_background_maintenance_event(
1904 manager.as_ref(),
1905 label,
1906 EventStatus::Error,
1907 start.elapsed().as_millis() as u64,
1908 "failed",
1909 Some(error.to_string()),
1910 Some(&policy_for_task),
1911 ),
1912 }
1913 result
1914 };
1915
1916 if let Err(error) = self.background_maintenance.spawn(key, observed_task) {
1917 record_background_maintenance_event(
1918 self.observability_manager.as_ref(),
1919 label,
1920 EventStatus::Error,
1921 0,
1922 "spawn_failed",
1923 Some(error.to_string()),
1924 Some(policy),
1925 );
1926 warn!(label = label, error = %error, "background maintenance spawn failed");
1927 return Err(error);
1928 }
1929 Ok(())
1930 }
1931
1932 fn record_skipped_maintenance(
1934 &self,
1935 label: &str,
1936 purpose: ObservationPurpose,
1937 reason: &str,
1938 policy: Option<&crate::optimization::config::MaintenanceTaskPolicy>,
1939 ) {
1940 if let Some(manager) = self.observability_manager.as_ref() {
1941 let mut tags = background_maintenance_tags(label, "skipped", Some(reason), policy);
1942 tags.insert("runtime.skip_reason".to_string(), reason.to_string());
1943 manager.record_lifecycle_event(
1944 EventType::MemoryOperation {
1945 operation: format!("{}_maintenance", label),
1946 },
1947 purpose,
1948 EventStatus::Skipped,
1949 0,
1950 tags,
1951 None,
1952 );
1953 }
1954 }
1955
1956 pub fn actor_facts(&self) -> Vec<ai_agents_core::KeyFact> {
1958 let Some(actor_id) = self.effective_actor_id() else {
1959 return Vec::new();
1960 };
1961 self.actor_facts_cache
1962 .read()
1963 .get(&actor_id)
1964 .cloned()
1965 .unwrap_or_default()
1966 }
1967
1968 pub fn relationship_memory_text(&self) -> Option<String> {
1970 self.format_relationship_for_context().map(|(_, text)| text)
1971 }
1972
1973 pub async fn extract_facts(
1975 &self,
1976 last_n: usize,
1977 ) -> ai_agents_core::Result<Vec<ai_agents_core::KeyFact>> {
1978 self.extract_facts_with_source(last_n, "manual").await
1979 }
1980
1981 async fn extract_facts_with_source(
1982 &self,
1983 last_n: usize,
1984 source: &'static str,
1985 ) -> ai_agents_core::Result<Vec<ai_agents_core::KeyFact>> {
1986 let extractor = match self.fact_extractor.read().clone() {
1987 Some(e) => e,
1988 None => return Ok(vec![]),
1989 };
1990
1991 let messages = Self::readable_native_messages(self.memory.get_messages(None).await?)?;
1992 let recent: Vec<_> = messages.iter().rev().take(last_n).rev().cloned().collect();
1993
1994 if recent.is_empty() {
1995 return Ok(vec![]);
1996 }
1997
1998 let actor_id = self.effective_actor_id();
1999 let existing = actor_id
2000 .as_ref()
2001 .and_then(|aid| self.actor_facts_cache.read().get(aid).cloned())
2002 .unwrap_or_default();
2003
2004 let categories = self
2005 .facts_config
2006 .as_ref()
2007 .map(|c| c.custom_categories.clone())
2008 .unwrap_or_default();
2009
2010 let facts = self
2011 .observe_purpose(
2012 ObservationPurpose::FactsExtraction,
2013 extractor.extract(&recent, &existing, actor_id.as_deref(), &categories),
2014 )
2015 .await?;
2016
2017 if !facts.is_empty() {
2019 let fact_store_opt = self.fact_store.read().clone();
2020 let mut stored_total = 0usize;
2021 let mut cache_updated = false;
2022 if let (Some(store), Some(aid)) = (fact_store_opt, &actor_id) {
2023 let authoritative = store.add_facts(aid, facts.clone()).await?;
2025 stored_total = authoritative.len();
2026 self.actor_facts_cache
2027 .write()
2028 .insert(aid.clone(), authoritative);
2029 cache_updated = true;
2030 } else if let Some(aid) = &actor_id {
2031 let mut cache = self.actor_facts_cache.write();
2032 let entry = cache.entry(aid.clone()).or_default();
2033 entry.extend(facts.clone());
2034 stored_total = entry.len();
2035 cache_updated = true;
2036 }
2037
2038 info!(
2039 actor_id = %actor_id.as_deref().unwrap_or("<none>"),
2040 source = source,
2041 requested_messages = last_n,
2042 message_count = recent.len(),
2043 extracted_count = facts.len(),
2044 cache_updated = cache_updated,
2045 stored_total = stored_total,
2046 "facts extracted"
2047 );
2048
2049 if let Some(ref aid) = actor_id {
2050 self.hooks.on_facts_extracted(aid, &facts).await;
2051 }
2052 }
2053
2054 Ok(facts)
2055 }
2056
2057 fn resolve_actor_id_from_context(&self) {
2060 if self
2061 .current_turn_actor_context()
2062 .and_then(|ctx| ctx.effective_actor_id().map(str::to_string))
2063 .is_some()
2064 {
2065 return;
2066 }
2067
2068 if let Some(ref am_config) = self.actor_memory_config
2069 && am_config.identification.method == ai_agents_facts::IdentificationMethod::FromContext
2070 && let Some(ref path) = am_config.identification.context_path
2071 {
2072 let val = self
2074 .context_manager
2075 .get_path(path)
2076 .or_else(|| self.context_manager.get(path));
2077 if let Some(val) = val
2078 && let Some(id_str) = val.as_str()
2079 {
2080 let current = self.actor_id.read().clone();
2081 if current.as_deref() != Some(id_str) {
2082 *self.actor_id.write() = Some(id_str.to_string());
2083 let mut meta = self.session_metadata.write();
2084 meta.actor_id = Some(id_str.to_string());
2085 if !meta.actors.iter().any(|a| a == id_str) {
2086 meta.actors.push(id_str.to_string());
2087 }
2088 }
2089 }
2090 }
2091 }
2092
2093 fn format_actor_facts_for_context(&self) -> String {
2095 let should_inject = self
2097 .facts_config
2098 .as_ref()
2099 .map(|c| c.inject_in_context)
2100 .unwrap_or(true);
2101 if !should_inject {
2102 return String::new();
2103 }
2104
2105 let Some(actor_id) = self.effective_actor_id() else {
2106 return String::new();
2107 };
2108
2109 let facts = self
2110 .actor_facts_cache
2111 .read()
2112 .get(&actor_id)
2113 .cloned()
2114 .unwrap_or_default();
2115 if facts.is_empty() {
2116 return String::new();
2117 }
2118
2119 let am_config = self.actor_memory_config.as_ref();
2120 let facts_budget = self
2123 .memory_token_budget
2124 .as_ref()
2125 .map(|b| b.allocation.facts as usize)
2126 .filter(|n| *n > 0);
2127 let default_max = am_config.map(|c| c.injection.max_tokens).unwrap_or(800);
2128 let max_tokens = facts_budget.unwrap_or(default_max);
2129
2130 let filtered: Vec<ai_agents_core::KeyFact> = if let Some(cfg) = am_config {
2132 if cfg.injection.mode == ai_agents_facts::InjectionMode::OnDemand {
2133 return String::new();
2134 }
2135 if cfg.injection.mode == ai_agents_facts::InjectionMode::Category
2136 && !cfg.injection.categories.is_empty()
2137 {
2138 facts
2139 .iter()
2140 .filter(|f| {
2141 cfg.injection
2142 .categories
2143 .iter()
2144 .any(|c| f.category.to_string() == *c)
2145 })
2146 .cloned()
2147 .collect()
2148 } else {
2149 facts.clone()
2150 }
2151 } else {
2152 facts.clone()
2153 };
2154
2155 if filtered.is_empty() {
2156 return String::new();
2157 }
2158
2159 if let Some(store) = self.fact_store.read().clone() {
2160 store.format_for_context(&filtered, max_tokens)
2161 } else {
2162 String::new()
2163 }
2164 }
2165
2166 fn build_context_with_staged(&self, staged: &HashMap<String, Value>) -> HashMap<String, Value> {
2167 let context = self.build_context_with_overlays();
2168 let mut root = Value::Object(context.into_iter().collect());
2169 for (path, value) in staged {
2170 if let Ok(updated) = ai_agents_core::set_dot_path(root.clone(), path, value.clone()) {
2171 root = updated;
2172 }
2173 }
2174 match root {
2175 Value::Object(obj) => obj.into_iter().collect(),
2176 _ => HashMap::new(),
2177 }
2178 }
2179
2180 fn build_context_with_overlays(&self) -> HashMap<String, Value> {
2181 let mut context = self.context_manager.get_all();
2182 let mut root = Value::Object(context.clone().into_iter().collect());
2183
2184 if let Some(turn_ctx) = self.current_turn_actor_context() {
2185 if let Some(ref origin_actor_id) = turn_ctx.origin_actor_id
2186 && let Ok(updated) = ai_agents_core::set_dot_path(
2187 root.clone(),
2188 "interaction.origin_actor_id",
2189 serde_json::json!(origin_actor_id),
2190 )
2191 {
2192 root = updated;
2193 }
2194 if let Some(ref sender_agent_id) = turn_ctx.sender_agent_id
2195 && let Ok(updated) = ai_agents_core::set_dot_path(
2196 root.clone(),
2197 "interaction.sender_agent_id",
2198 serde_json::json!(sender_agent_id),
2199 )
2200 {
2201 root = updated;
2202 }
2203 }
2204
2205 if let Some(ref actor_id) = self.effective_actor_id()
2206 && let Ok(updated) = ai_agents_core::set_dot_path(
2207 root.clone(),
2208 "interaction.actor_id",
2209 serde_json::json!(actor_id),
2210 )
2211 {
2212 root = updated;
2213 }
2214
2215 if let Some(manager) = self.relationship_manager.as_ref()
2216 && let Some(actor_id) = self.effective_actor_id()
2217 && let Some(value) = manager.to_context_value(&actor_id)
2218 && let Ok(updated) = ai_agents_core::set_dot_path(
2219 root.clone(),
2220 &manager.config().injection.context_path,
2221 value,
2222 )
2223 {
2224 root = updated;
2225 }
2226
2227 if let Value::Object(obj) = root {
2228 context = obj.into_iter().collect();
2229 }
2230
2231 context
2232 }
2233
2234 fn resolve_actor_name_from_context(&self) -> Option<String> {
2235 for path in ["actor.name", "user.name", "player.name", "customer.name"] {
2236 if let Some(value) = self.context_manager.get_path(path)
2237 && let Some(name) = value.as_str()
2238 {
2239 return Some(name.to_string());
2240 }
2241 }
2242 None
2243 }
2244
2245 async fn maybe_load_actor_relationship(&self) {
2246 let Some(manager) = self.relationship_manager.as_ref() else {
2247 return;
2248 };
2249 let Some(actor_id) = self.effective_actor_id() else {
2250 return;
2251 };
2252
2253 let mut should_fire_loaded = false;
2254 if manager.get(&actor_id).is_none() {
2255 let mut loaded = false;
2256 if manager.config().persistence.enabled {
2257 let storage = self.storage.read().clone();
2258 if let Some(storage) = storage {
2259 match storage.load_relationship(&self.info.id, &actor_id).await {
2260 Ok(Some(value)) => match manager.insert_from_value(value) {
2261 Ok(_) => loaded = true,
2262 Err(e) => {
2263 warn!(actor = %actor_id, error = %e, "failed to restore relationship")
2264 }
2265 },
2266 Ok(None) => {}
2267 Err(e) => {
2268 warn!(actor = %actor_id, error = %e, "failed to load relationship")
2269 }
2270 }
2271 }
2272 }
2273
2274 if !loaded {
2275 manager.get_or_create(&actor_id, self.resolve_actor_name_from_context().as_deref());
2276 }
2277 should_fire_loaded = true;
2278 }
2279
2280 let actor_name = self.resolve_actor_name_from_context();
2281 let relationship = manager.touch_interaction(&actor_id, actor_name.as_deref());
2282 if should_fire_loaded {
2283 self.hooks
2284 .on_relationship_loaded(&actor_id, &relationship)
2285 .await;
2286 }
2287 }
2288
2289 fn format_relationship_for_context(&self) -> Option<(String, String)> {
2290 let manager = self.relationship_manager.as_ref()?;
2291 if !manager.config().injection.enabled {
2292 return None;
2293 }
2294 let actor_id = self.effective_actor_id()?;
2295 let relationship = manager.get(&actor_id)?;
2296 let local_cap = manager.config().injection.max_tokens;
2297 let global_cap = self
2298 .memory_token_budget
2299 .as_ref()
2300 .map(|b| b.allocation.relationships as usize)
2301 .filter(|n| *n > 0);
2302 let max_tokens = global_cap.map(|g| g.min(local_cap)).unwrap_or(local_cap);
2303 let text = ai_agents_relationships::format_relationship(
2304 &relationship,
2305 &manager.config().injection.format,
2306 max_tokens,
2307 );
2308 if text.is_empty() {
2309 None
2310 } else {
2311 Some((manager.config().injection.prompt_variable.clone(), text))
2312 }
2313 }
2314
2315 async fn persist_actor_relationship(&self, actor_id: &str) -> Result<()> {
2316 let Some(manager) = self.relationship_manager.as_ref() else {
2317 return Ok(());
2318 };
2319 if !manager.config().persistence.enabled {
2320 return Ok(());
2321 }
2322 let storage = self.storage.read().clone();
2323 let Some(storage) = storage else {
2324 return Ok(());
2325 };
2326 if let Some(value) = manager.relationship_as_value(actor_id)? {
2327 storage
2328 .save_relationship(&self.info.id, actor_id, &value)
2329 .await?;
2330 }
2331 Ok(())
2332 }
2333
2334 async fn auto_update_relationship(&self) {
2335 let Some(manager) = self.relationship_manager.as_ref() else {
2336 return;
2337 };
2338 let Some(actor_id) = self.effective_actor_id() else {
2339 return;
2340 };
2341 if !manager.config().auto_update.enabled {
2342 let _ = self.persist_actor_relationship(&actor_id).await;
2343 return;
2344 }
2345
2346 let recent_messages = manager.config().auto_update.recent_messages;
2347 let messages = match self.memory.get_messages(Some(recent_messages)).await {
2348 Ok(messages) => messages,
2349 Err(e) => {
2350 warn!(actor = %actor_id, error = %e, "failed to read messages for relationship update");
2351 return;
2352 }
2353 };
2354 let messages = match Self::readable_native_messages(messages) {
2355 Ok(messages) => messages,
2356 Err(error) => {
2357 warn!(actor = %actor_id, error = %error, "failed to project native history for relationship update");
2358 return;
2359 }
2360 };
2361
2362 match self
2363 .observe_purpose(
2364 ObservationPurpose::RelationshipUpdate,
2365 manager.auto_update(&actor_id, &messages),
2366 )
2367 .await
2368 {
2369 Ok(update) => {
2370 if !update.changes.is_empty() {
2371 self.hooks
2372 .on_relationship_change(&actor_id, &update.changes)
2373 .await;
2374 }
2375 if let Some(ref event) = update.event {
2376 self.hooks.on_notable_event(&actor_id, event).await;
2377 }
2378 let persisted = match self.persist_actor_relationship(&actor_id).await {
2379 Ok(()) => true,
2380 Err(e) => {
2381 warn!(actor = %actor_id, error = %e, "failed to persist relationship");
2382 false
2383 }
2384 };
2385 if !update.changes.is_empty() || update.event.is_some() {
2386 let changed_dimensions: Vec<String> = update
2387 .changes
2388 .iter()
2389 .map(|change| format!("{}:{}", change.perspective, change.dimension))
2390 .collect();
2391 info!(
2392 actor_id = %actor_id,
2393 change_count = update.changes.len(),
2394 changed_dimensions = ?changed_dimensions,
2395 event_present = update.event.is_some(),
2396 persisted = persisted,
2397 "relationship updated"
2398 );
2399 } else {
2400 debug!(actor_id = %actor_id, persisted = persisted, "relationship evaluation ran but found no changes");
2401 }
2402 }
2403 Err(e) => warn!(actor = %actor_id, error = %e, "relationship update failed"),
2404 }
2405 }
2406
2407 async fn auto_extract_facts(&self) {
2409 let should_extract = self
2410 .facts_config
2411 .as_ref()
2412 .map(|c| c.enabled && c.auto_extract)
2413 .unwrap_or(false);
2414
2415 if !should_extract {
2416 debug!("fact extraction skipped because auto extraction is disabled");
2417 return;
2418 }
2419
2420 let msgs_since = *self.messages_since_extraction.read();
2421 if msgs_since < 2 {
2422 debug!(
2423 messages_since_extraction = msgs_since,
2424 "fact extraction skipped until threshold is reached"
2425 );
2426 return;
2427 }
2428
2429 match self.extract_facts_with_source(msgs_since, "auto").await {
2430 Ok(facts) => {
2431 if !facts.is_empty() {
2432 *self.messages_since_extraction.write() = 0;
2433 } else {
2434 debug!("fact extraction ran but found no new facts");
2435 }
2436 }
2437 Err(e) => {
2438 warn!("fact extraction failed: {}", e);
2439 }
2440 }
2441 }
2442
2443 pub fn with_persona(mut self, manager: Arc<ai_agents_persona::PersonaManager>) -> Self {
2444 self.persona_manager = Some(manager);
2445 self
2446 }
2447
2448 pub fn persona_manager(&self) -> Option<&Arc<ai_agents_persona::PersonaManager>> {
2449 self.persona_manager.as_ref()
2450 }
2451
2452 pub fn with_disambiguation(mut self, config: DisambiguationConfig) -> Self {
2453 if config.is_enabled() {
2454 let manager = DisambiguationManager::new(config, Arc::clone(&self.llm_registry))
2455 .with_clarification_observer(Arc::new(ObservabilityClarificationObserver));
2456 self.disambiguation_manager = Some(manager);
2457 }
2458 self
2459 }
2460
2461 pub fn disambiguation_manager(&self) -> Option<&DisambiguationManager> {
2462 self.disambiguation_manager.as_ref()
2463 }
2464
2465 pub fn has_disambiguation(&self) -> bool {
2466 self.disambiguation_manager
2467 .as_ref()
2468 .is_some_and(|m| m.is_enabled())
2469 }
2470
2471 pub async fn init_storage(&self) -> Result<()> {
2472 let _guard = self.storage_init.lock().await;
2476 let mut storage = self.storage.read().clone();
2477 if storage.is_none() && !self.storage_config.is_none() {
2478 let storage_config = self.convert_storage_config();
2479 storage = create_storage(&storage_config).await?;
2480 *self.storage.write() = storage.clone();
2481 }
2482
2483 self.validate_storage_requirements(storage.as_deref())?;
2484 self.complete_facts_init().await?;
2485 Ok(())
2486 }
2487
2488 fn validate_storage_requirements(&self, storage: Option<&dyn AgentStorage>) -> Result<()> {
2489 let facts_required = self
2490 .facts_config
2491 .as_ref()
2492 .is_some_and(|config| config.enabled)
2493 || self
2494 .actor_memory_config
2495 .as_ref()
2496 .is_some_and(|config| config.enabled);
2497 let relationships_required = self
2498 .relationship_manager
2499 .as_ref()
2500 .is_some_and(|manager| manager.config().persistence.enabled);
2501
2502 let Some(storage) = storage else {
2503 let mut requirements = Vec::new();
2504 if facts_required {
2505 requirements.push("actor facts or actor memory");
2506 }
2507 if relationships_required {
2508 requirements.push("persistent relationships");
2509 }
2510 if requirements.is_empty() {
2511 return Ok(());
2512 }
2513 return Err(AgentError::Config(format!(
2514 "Storage is required for enabled {} but none is configured or injected",
2515 requirements.join(" and ")
2516 )));
2517 };
2518
2519 if facts_required && !storage.supports(StorageCapability::ActorFacts) {
2523 return Err(AgentError::UnsupportedStorageCapability(
2524 StorageCapability::ActorFacts,
2525 ));
2526 }
2527 if relationships_required && !storage.supports(StorageCapability::ActorRelationships) {
2528 return Err(AgentError::UnsupportedStorageCapability(
2529 StorageCapability::ActorRelationships,
2530 ));
2531 }
2532 Ok(())
2533 }
2534
2535 async fn complete_facts_init(&self) -> Result<()> {
2539 if self.fact_store.read().is_some() {
2540 return Ok(());
2541 }
2542 let storage = match self.storage.read().clone() {
2543 Some(s) => s,
2544 None => return Ok(()),
2545 };
2546
2547 let facts_enabled = self
2548 .facts_config
2549 .as_ref()
2550 .map(|f| f.enabled)
2551 .unwrap_or(false);
2552 let actor_memory_enabled = self
2553 .actor_memory_config
2554 .as_ref()
2555 .map(|a| a.enabled)
2556 .unwrap_or(false);
2557
2558 if !facts_enabled && !actor_memory_enabled {
2559 return Ok(());
2560 }
2561
2562 let fc = self.facts_config.clone().unwrap_or_default();
2563 let store = Arc::new(ai_agents_facts::FactStore::new(
2564 storage,
2565 self.info.id.clone(),
2566 fc.clone(),
2567 ));
2568
2569 let extractor: Option<Arc<dyn ai_agents_facts::FactExtractor>> = if facts_enabled {
2570 let extractor_llm = self.optional_role_llm(
2571 ai_agents_llm::LLMRole::MemoryFacts,
2572 fc.extractor_llm.as_deref(),
2573 || {
2574 fc.extractor_llm
2575 .as_ref()
2576 .and_then(|alias| self.llm_registry.get(alias).ok())
2577 .or_else(|| self.llm_registry.router().ok())
2578 .or_else(|| self.llm_registry.default().ok())
2579 },
2580 )?;
2581 extractor_llm.map(|llm| {
2582 Arc::new(ai_agents_facts::LLMFactExtractor::new(llm, fc.clone()))
2583 as Arc<dyn ai_agents_facts::FactExtractor>
2584 })
2585 } else {
2586 None
2587 };
2588
2589 *self.fact_store.write() = Some(store);
2590 *self.fact_extractor.write() = extractor;
2591 debug!(
2592 agent = %self.info.id,
2593 facts_enabled,
2594 actor_memory_enabled,
2595 "facts storage initialized"
2596 );
2597 Ok(())
2598 }
2599
2600 fn convert_storage_config(&self) -> StorageStorageConfig {
2601 crate::spec::storage::to_storage_config(&self.storage_config)
2602 }
2603
2604 pub fn storage(&self) -> Option<Arc<dyn AgentStorage>> {
2605 self.storage.read().clone()
2606 }
2607
2608 pub fn storage_config(&self) -> &StorageConfig {
2609 &self.storage_config
2610 }
2611
2612 pub fn spawner(&self) -> Option<&Arc<crate::spawner::AgentSpawner>> {
2614 self.spawner.as_ref()
2615 }
2616
2617 pub fn spawner_registry(&self) -> Option<&Arc<crate::spawner::AgentRegistry>> {
2619 self.spawner_registry.as_ref()
2620 }
2621
2622 pub fn has_spawner(&self) -> bool {
2623 self.spawner_registry.is_some()
2624 }
2625
2626 pub fn with_spawner_handles(
2627 mut self,
2628 spawner: Arc<crate::spawner::AgentSpawner>,
2629 registry: Arc<crate::spawner::AgentRegistry>,
2630 ) -> Self {
2631 self.spawner = Some(spawner);
2632 self.spawner_registry = Some(registry);
2633 self
2634 }
2635
2636 pub fn with_hooks(mut self, hooks: Arc<dyn AgentHooks>) -> Self {
2637 self.hooks = hooks;
2638 self
2639 }
2640
2641 pub fn with_parallel_tools(mut self, config: ParallelToolsConfig) -> Self {
2642 self.parallel_tools = config;
2643 self
2644 }
2645
2646 pub fn with_streaming(mut self, config: StreamingConfig) -> Self {
2647 self.streaming = config;
2648 self
2649 }
2650
2651 pub fn with_hitl(mut self, engine: HITLEngine, handler: Arc<dyn ApprovalHandler>) -> Self {
2652 self.hitl_engine = Some(engine);
2653 self.approval_handler = handler;
2654 self
2655 }
2656
2657 pub fn with_max_context_tokens(mut self, tokens: u32) -> Self {
2658 self.max_context_tokens = tokens;
2659 self
2660 }
2661
2662 pub fn with_memory_token_budget(mut self, budget: MemoryTokenBudget) -> Self {
2663 self.memory_token_budget = Some(budget);
2664 self
2665 }
2666
2667 pub fn with_recovery_manager(mut self, manager: RecoveryManager) -> Self {
2668 self.recovery_manager = manager;
2669 self
2670 }
2671
2672 pub fn with_tool_security(mut self, engine: ToolSecurityEngine) -> Self {
2673 self.tool_security = engine;
2674 self
2675 }
2676
2677 pub fn runtime_control(&self) -> RuntimeControlHandle {
2679 RuntimeControlHandle {
2680 state: Arc::clone(&self.runtime_control),
2681 }
2682 }
2683
2684 pub fn set_question_handler(&self, handler: Option<Arc<dyn QuestionHandler>>) {
2686 self.tools.set_question_handler(handler);
2687 }
2688
2689 pub fn set_diagnostics_provider(&self, provider: Arc<dyn DiagnosticsProvider>) {
2691 self.tools.set_diagnostics_provider(provider);
2692 }
2693
2694 pub fn set_command_runner(&self, runner: Arc<dyn CommandRunner>) {
2696 self.tools.set_command_runner(runner);
2697 }
2698
2699 pub fn set_web_search_provider(&self, provider: Arc<dyn ai_agents_tools::WebSearchProvider>) {
2701 self.tools.set_web_search_provider(provider);
2702 }
2703
2704 pub fn todos(&self) -> Vec<TodoItem> {
2706 self.tools.todos()
2707 }
2708
2709 fn active_tool_security(&self) -> ToolSecurityEngine {
2711 self.runtime_control
2712 .tool_security_override
2713 .read()
2714 .clone()
2715 .unwrap_or_else(|| self.tool_security.clone())
2716 }
2717
2718 fn runtime_safety_snapshot(&self) -> RuntimeSafetySnapshot {
2720 let _guard = self.runtime_control.snapshot_guard.read();
2721 RuntimeSafetySnapshot {
2722 version: self.runtime_control.version.load(Ordering::SeqCst),
2723 emergency_deny: self.runtime_control.emergency_deny.load(Ordering::SeqCst),
2724 tool_security: self
2725 .runtime_control
2726 .tool_security_override
2727 .read()
2728 .clone()
2729 .unwrap_or_else(|| self.tool_security.clone()),
2730 tool_scope_override: self.runtime_control.tool_scope_override.read().clone(),
2731 }
2732 }
2733
2734 fn admit_tool_execution(
2736 &self,
2737 expected_runtime_version: u64,
2738 expected_policy_version: u64,
2739 expected_state_generation: Option<u64>,
2740 canonical_id: &str,
2741 ) -> SecurityCheckResult {
2742 let _guard = self.runtime_control.snapshot_guard.read();
2743 if self.runtime_control.emergency_deny.load(Ordering::SeqCst) {
2744 return SecurityCheckResult::Block {
2745 reason: "runtime emergency deny is enabled".to_string(),
2746 };
2747 }
2748 let runtime_version = self.runtime_control.version.load(Ordering::SeqCst);
2749 let security_engine = self
2750 .runtime_control
2751 .tool_security_override
2752 .read()
2753 .clone()
2754 .unwrap_or_else(|| self.tool_security.clone());
2755 if runtime_version != expected_runtime_version
2756 || security_engine.policy_version() != expected_policy_version
2757 {
2758 return SecurityCheckResult::Block {
2759 reason: "runtime safety controls changed before admission".to_string(),
2760 };
2761 }
2762 let current_state_generation = self
2763 .state_machine
2764 .as_ref()
2765 .map(|state_machine| state_machine.generation());
2766 if current_state_generation != expected_state_generation {
2767 return SecurityCheckResult::Block {
2768 reason: "state scope changed before admission".to_string(),
2769 };
2770 }
2771 security_engine.admit_tool_execution(canonical_id)
2772 }
2773
2774 pub fn with_process_processor(mut self, processor: ProcessProcessor) -> Self {
2775 let processor = processor.with_stage_observer(Arc::new(ObservabilityProcessStageObserver));
2776 self.process_processor = Some(processor);
2777 self
2778 }
2779
2780 pub fn with_state_machine(
2781 mut self,
2782 state_machine: Arc<StateMachine>,
2783 evaluator: Arc<dyn TransitionEvaluator>,
2784 ) -> Self {
2785 self.state_machine = Some(state_machine);
2786 self.transition_evaluator = Some(evaluator);
2787 self
2788 }
2789
2790 pub fn with_context_manager(mut self, manager: Arc<ContextManager>) -> Self {
2791 self.context_manager = manager;
2792 self
2793 }
2794
2795 pub fn register_message_filter(&self, name: impl Into<String>, filter: Arc<dyn MessageFilter>) {
2796 self.message_filters.write().insert(name.into(), filter);
2797 }
2798
2799 pub fn set_context(&self, key: &str, value: Value) -> Result<()> {
2800 self.context_manager.update(key, value)
2801 }
2802
2803 pub fn update_context(&self, path: &str, value: Value) -> Result<()> {
2804 self.context_manager.update(path, value)
2805 }
2806
2807 pub fn get_context(&self) -> HashMap<String, Value> {
2808 self.build_context_with_overlays()
2809 }
2810
2811 pub fn remove_context(&self, key: &str) -> Option<Value> {
2812 self.context_manager.remove(key)
2813 }
2814
2815 pub async fn refresh_context(&self, key: &str) -> Result<()> {
2816 self.context_manager.refresh(key).await
2817 }
2818
2819 pub fn register_context_provider(&self, name: &str, provider: Arc<dyn ContextProvider>) {
2820 self.context_manager.register_provider(name, provider);
2821 }
2822
2823 pub fn current_state(&self) -> Option<String> {
2824 self.state_machine.as_ref().map(|sm| sm.current())
2825 }
2826
2827 async fn invalidate_pending_confirmation(&self, reason: &'static str) {
2829 self.disambiguation_epoch.fetch_add(1, Ordering::SeqCst);
2830 let Some(disambiguator) = self.disambiguation_manager.as_ref() else {
2831 return;
2832 };
2833 if disambiguator.has_pending_confirmation().await {
2834 disambiguator.clear_pending().await;
2835 *self.pending_skill_id.write() = None;
2836 info!(
2837 confirmation_event = "invalidated",
2838 invalidation_reason = reason,
2839 "Runtime invalidated pending confirmation"
2840 );
2841 }
2842 }
2843
2844 async fn admit_disambiguation_redispatch(
2846 &self,
2847 expected_epoch: u64,
2848 expected_state_generation: Option<u64>,
2849 ) -> Result<tokio::sync::RwLockReadGuard<'_, ()>> {
2850 let admission = self.disambiguation_admission.read().await;
2851 let state_generation = self
2852 .state_machine
2853 .as_ref()
2854 .map(|state_machine| state_machine.generation());
2855 if self.disambiguation_epoch.load(Ordering::SeqCst) != expected_epoch
2856 || state_generation != expected_state_generation
2857 {
2858 return Err(AgentError::Other(
2859 "Disambiguation ownership changed before redispatch admission".to_string(),
2860 ));
2861 }
2862 Ok(admission)
2863 }
2864
2865 fn reserve_state_transition(&self) -> Option<StateTransitionReservation<'_>> {
2867 self.state_transition_reserved
2868 .compare_exchange(false, true, Ordering::SeqCst, Ordering::SeqCst)
2869 .ok()
2870 .map(|_| StateTransitionReservation {
2871 reserved: &self.state_transition_reserved,
2872 })
2873 }
2874
2875 async fn admit_optional_disambiguation_ownership(
2877 &self,
2878 ownership: Option<DisambiguationOwnership>,
2879 ) -> Result<Option<tokio::sync::RwLockReadGuard<'_, ()>>> {
2880 match ownership {
2881 Some(ownership) => self
2882 .admit_disambiguation_redispatch(ownership.epoch, ownership.state_generation)
2883 .await
2884 .map(Some),
2885 None => Ok(None),
2886 }
2887 }
2888
2889 pub async fn transition_to(&self, state: &str) -> Result<()> {
2891 let Some(ref sm) = self.state_machine else {
2892 return Ok(());
2893 };
2894 let claim_admission = self.disambiguation_admission.write().await;
2895 let reservation = self.reserve_state_transition().ok_or_else(|| {
2896 AgentError::Other("Another state transition is already in progress".to_string())
2897 })?;
2898 let from_state = sm.current();
2899 let expected_state_generation = sm.generation();
2900 let expected_disambiguation_epoch = self.disambiguation_epoch.load(Ordering::SeqCst);
2901 let history_before = sm.history();
2902 drop(claim_admission);
2903
2904 self.execute_state_exit_actions(&from_state).await;
2905
2906 let admission = self.disambiguation_admission.write().await;
2907 if sm.current() != from_state
2908 || sm.generation() != expected_state_generation
2909 || self.disambiguation_epoch.load(Ordering::SeqCst) != expected_disambiguation_epoch
2910 {
2911 return Err(AgentError::Other(
2912 "State ownership changed during manual transition preparation".to_string(),
2913 ));
2914 }
2915 sm.transition_to(state, "manual transition")?;
2916 self.invalidate_pending_confirmation("state_transition")
2917 .await;
2918 let entered = sm.current();
2919 let is_reentry = Self::state_was_previously_entered(&entered, &from_state, &history_before);
2920 drop(admission);
2921
2922 self.execute_state_enter_actions(&entered, is_reentry).await;
2923 drop(reservation);
2924 info!(to = %entered, "Manual state transition");
2925 Ok(())
2926 }
2927
2928 pub fn state_history(&self) -> Vec<StateTransitionEvent> {
2929 self.state_machine
2930 .as_ref()
2931 .map(|sm| sm.history())
2932 .unwrap_or_default()
2933 }
2934
2935 pub fn session_metadata(&self) -> ai_agents_core::SessionMetadata {
2937 self.session_metadata.read().clone()
2938 }
2939
2940 pub async fn delete_actor_data(&self, actor_id: &str) -> Result<()> {
2943 let allowed = self
2944 .actor_memory_config
2945 .as_ref()
2946 .map(|c| c.privacy.allow_deletion)
2947 .unwrap_or(true);
2948 if !allowed {
2949 return Err(AgentError::Config(
2950 "privacy.allow_deletion is false; actor data deletion is not permitted".into(),
2951 ));
2952 }
2953 let storage = self.storage.read().clone();
2954 if let Some(storage) = storage {
2955 if !storage.supports(StorageCapability::ActorDataDeletion) {
2959 return Err(AgentError::UnsupportedStorageCapability(
2960 StorageCapability::ActorDataDeletion,
2961 ));
2962 }
2963 storage.delete_actor_data(&self.info.id, actor_id).await?;
2964 } else {
2965 let store = { self.fact_store.read().clone() };
2969 if let Some(store) = store {
2970 store.delete_actor_data(actor_id).await?;
2971 }
2972 }
2973 if let Some(manager) = self.relationship_manager.as_ref() {
2974 manager.remove(actor_id);
2975 }
2976 self.actor_facts_cache.write().remove(actor_id);
2977 Ok(())
2978 }
2979
2980 pub fn set_session_metadata(&self, meta: ai_agents_core::SessionMetadata) {
2982 *self.session_metadata.write() = meta;
2983 }
2984
2985 pub async fn cleanup_expired_sessions(&self) -> Result<usize> {
2987 let storage = self.storage.read().clone();
2988 match storage {
2989 Some(s) => {
2990 let count = s.cleanup_expired().await?;
2991 if count > 0 {
2992 self.hooks.on_sessions_expired(count).await;
2993 }
2994 Ok(count)
2995 }
2996 None => Err(AgentError::Config(
2997 "No storage configured. Use with_storage_config() or with_storage() first".into(),
2998 )),
2999 }
3000 }
3001
3002 pub async fn list_sessions_filtered(
3004 &self,
3005 filter: &ai_agents_core::SessionFilter,
3006 ) -> Result<Vec<ai_agents_core::SessionSummary>> {
3007 let storage = self.storage.read().clone();
3008 match storage {
3009 Some(s) => s.list_sessions_filtered(filter).await,
3010 None => Err(AgentError::Config(
3011 "No storage configured. Use with_storage_config() or with_storage() first".into(),
3012 )),
3013 }
3014 }
3015
3016 pub async fn save_state(&self) -> Result<AgentSnapshot> {
3017 let memory_snapshot = self.memory.snapshot().await?;
3018 let state_machine_snapshot = self.state_machine.as_ref().map(|sm| sm.snapshot());
3019 let context_snapshot = self.context_manager.snapshot();
3020
3021 let mut snapshot = AgentSnapshot::new(self.info.id.clone())
3022 .with_memory(memory_snapshot)
3023 .with_context(context_snapshot)
3024 .with_state_machine(
3025 state_machine_snapshot.unwrap_or_else(|| StateMachineSnapshot {
3026 current_state: String::new(),
3027 previous_state: None,
3028 turn_count: 0,
3029 no_transition_count: 0,
3030 history: vec![],
3031 }),
3032 );
3033
3034 if let Some(ref persona) = self.persona_manager {
3035 snapshot.persona = Some(persona.snapshot_as_value()?);
3036 }
3037
3038 if let Some(ref relationships) = self.relationship_manager {
3039 snapshot.relationships = Some(relationships.snapshot_as_value()?);
3040 }
3041
3042 Ok(snapshot)
3043 }
3044
3045 pub(crate) fn prepared_persistence_spec(
3047 &self,
3048 spec: &crate::spec::AgentSpec,
3049 ) -> Result<crate::spec::AgentSpec> {
3050 if spec.llm.router_roles().is_none() {
3051 if self.llm_registry.router_roles().is_some() {
3052 return Err(AgentError::Config(
3053 "Hierarchy runtime has no matching declared child spec".into(),
3054 ));
3055 }
3056 return Ok(spec.clone());
3057 }
3058 if spec.llm.router_roles() != self.llm_registry.router_roles()
3059 || spec.llm.get_default_alias() != self.llm_registry.default_alias()
3060 {
3061 return Err(AgentError::Config(
3062 "Hierarchy child spec does not match its live LLM routing".into(),
3063 ));
3064 }
3065 let mut prepared = spec.clone();
3066 prepared.skills = self
3067 .skills
3068 .iter()
3069 .cloned()
3070 .map(ai_agents_skills::SkillRef::Inline)
3071 .collect();
3072 prepared.reasoning = self.reasoning_config.clone();
3073 prepared.error_recovery = self.recovery_manager.config().clone();
3074 if let Some(manager) = &self.disambiguation_manager {
3075 prepared.disambiguation = manager.config().clone();
3076 }
3077 if let Some(config) = &self.facts_config {
3078 prepared.memory.facts = Some(config.clone());
3079 }
3080 prepared.reflection = self.reflection_config.clone();
3081 if let Some(processor) = &self.process_processor {
3082 prepared.process = processor.config().clone();
3083 }
3084 if let Some(engine) = &self.hitl_engine {
3085 prepared.hitl = Some(engine.config().clone());
3086 }
3087 if let Some(machine) = &self.state_machine {
3088 prepared.states = Some(machine.config().clone());
3089 }
3090 Ok(prepared)
3091 }
3092
3093 pub async fn save_state_full(&self) -> Result<AgentSnapshot> {
3095 let mut snapshot = self.save_state().await?;
3096 if let Some(ref registry) = self.spawner_registry {
3097 let entries = registry.persistence_entries()?;
3098 if !entries.is_empty() {
3099 snapshot = snapshot.with_spawned_agents(entries);
3100 }
3101 }
3102 Ok(snapshot)
3103 }
3104
3105 pub async fn restore_state(&self, snapshot: AgentSnapshot) -> Result<()> {
3107 let _admission = self.disambiguation_admission.write().await;
3108 if self.state_transition_reserved.load(Ordering::SeqCst) {
3109 return Err(AgentError::Other(
3110 "Cannot restore state while a state transition is in progress".to_string(),
3111 ));
3112 }
3113 self.invalidate_pending_confirmation("state_restore").await;
3114 *self.pending_skill_id.write() = None;
3115 if let Some(disambiguator) = self.disambiguation_manager.as_ref() {
3116 disambiguator.clear_pending().await;
3117 }
3118 self.memory.restore(snapshot.memory).await?;
3119 self.active_native_exchanges.write().clear();
3120
3121 if let (Some(sm), Some(sm_snapshot)) = (&self.state_machine, snapshot.state_machine)
3122 && !sm_snapshot.current_state.is_empty()
3123 {
3124 sm.restore(sm_snapshot)?;
3125 }
3126
3127 self.context_manager.restore(snapshot.context);
3128
3129 if let (Some(persona_value), Some(persona_manager)) =
3130 (snapshot.persona, &self.persona_manager)
3131 {
3132 persona_manager.restore_from_value(persona_value)?;
3133 }
3134
3135 if let (Some(relationship_value), Some(relationship_manager)) =
3136 (snapshot.relationships, &self.relationship_manager)
3137 {
3138 relationship_manager.restore_from_value(relationship_value)?;
3139 }
3140
3141 info!(agent_id = %snapshot.agent_id, "State restored");
3142 Ok(())
3143 }
3144
3145 pub async fn save_to(&self, storage: &dyn AgentStorage, session_id: &str) -> Result<()> {
3146 let snapshot = self.save_state().await?;
3147 storage.save(session_id, &snapshot).await
3148 }
3149
3150 async fn load_session_restore(
3151 storage: &dyn AgentStorage,
3152 session_id: &str,
3153 ) -> Result<Option<StoredSessionRestore>> {
3154 let Some(snapshot) = storage.load(session_id).await? else {
3155 return Ok(None);
3156 };
3157 let metadata = if storage.supports(StorageCapability::SessionMetadata) {
3161 storage.load_metadata(session_id).await?
3162 } else {
3163 None
3164 };
3165 Ok(Some(StoredSessionRestore { snapshot, metadata }))
3166 }
3167
3168 async fn capture_session_restore_point(&self) -> Result<RuntimeSessionRestorePoint> {
3169 Ok(RuntimeSessionRestorePoint {
3170 snapshot: self.save_state().await?,
3171 metadata: self.session_metadata(),
3172 actor_id: self.actor_id(),
3173 session_id: self.current_session_id.read().clone(),
3174 })
3175 }
3176
3177 async fn apply_session_restore_unchecked(
3178 &self,
3179 session_id: &str,
3180 stored: StoredSessionRestore,
3181 ) -> Result<()> {
3182 self.restore_state(stored.snapshot).await?;
3183 let metadata = stored.metadata.unwrap_or_default();
3184 if let Some(actor_id) = metadata.actor_id.as_deref() {
3185 self.set_actor_id(actor_id)?;
3186 } else {
3187 self.clear_actor_id();
3188 }
3189 self.set_session_metadata(metadata);
3190 *self.current_session_id.write() = Some(session_id.to_string());
3191 Ok(())
3192 }
3193
3194 async fn restore_session_restore_point(
3195 &self,
3196 restore_point: &RuntimeSessionRestorePoint,
3197 ) -> Result<()> {
3198 self.restore_state(restore_point.snapshot.clone()).await?;
3199 if let Some(actor_id) = restore_point.actor_id.as_deref() {
3200 self.set_actor_id(actor_id)?;
3201 } else {
3202 self.clear_actor_id();
3203 }
3204 self.set_session_metadata(restore_point.metadata.clone());
3205 *self.current_session_id.write() = restore_point.session_id.clone();
3206 Ok(())
3207 }
3208
3209 async fn apply_session_restore(
3210 &self,
3211 session_id: &str,
3212 stored: StoredSessionRestore,
3213 ) -> Result<()> {
3214 let before = self.capture_session_restore_point().await?;
3215 if let Err(error) = self
3216 .apply_session_restore_unchecked(session_id, stored)
3217 .await
3218 {
3219 return match self.restore_session_restore_point(&before).await {
3220 Ok(()) => Err(error),
3221 Err(rollback_error) => Err(AgentError::Other(format!(
3222 "Session restore failed: {error}; rollback failed: {rollback_error}"
3223 ))),
3224 };
3225 }
3226 Ok(())
3227 }
3228
3229 async fn rollback_session_restore_set(
3230 parent: Option<(&RuntimeAgent, &RuntimeSessionRestorePoint)>,
3231 children: &[(String, Arc<RuntimeAgent>, RuntimeSessionRestorePoint)],
3232 ) -> Vec<String> {
3233 let mut errors = Vec::new();
3234 if let Some((agent, restore_point)) = parent
3235 && let Err(error) = agent.restore_session_restore_point(restore_point).await
3236 {
3237 errors.push(format!("parent: {error}"));
3238 }
3239 for (id, agent, restore_point) in children {
3240 if let Err(error) = agent.restore_session_restore_point(restore_point).await {
3241 errors.push(format!("child '{id}': {error}"));
3242 }
3243 }
3244 errors
3245 }
3246
3247 fn restore_failure(error: impl std::fmt::Display, rollback_errors: Vec<String>) -> AgentError {
3248 if rollback_errors.is_empty() {
3249 AgentError::Other(format!(
3250 "Session restore failed: {error}; runtime state was rolled back"
3251 ))
3252 } else {
3253 AgentError::Other(format!(
3254 "Session restore failed: {error}; rollback also failed for {}",
3255 rollback_errors.join(", ")
3256 ))
3257 }
3258 }
3259
3260 pub async fn load_from(&self, storage: &dyn AgentStorage, session_id: &str) -> Result<bool> {
3261 let Some(stored) = Self::load_session_restore(storage, session_id).await? else {
3262 return Ok(false);
3263 };
3264 self.apply_session_restore(session_id, stored).await?;
3265 Ok(true)
3266 }
3267
3268 pub async fn save_session(&self, session_id: &str) -> Result<()> {
3269 let storage = self.storage.read().clone();
3270 match storage {
3271 Some(s) => {
3272 let is_new = {
3274 let cur = self.current_session_id.read().clone();
3275 cur.as_deref() != Some(session_id)
3276 };
3277 if is_new {
3278 *self.current_session_id.write() = Some(session_id.to_string());
3279 self.hooks.on_session_created(session_id).await;
3280 }
3281
3282 {
3284 let now = chrono::Utc::now();
3285 let msg_count = self
3286 .memory
3287 .get_messages(None)
3288 .await
3289 .map(|v| v.len())
3290 .unwrap_or(0);
3291 let mut meta = self.session_metadata.write();
3292 meta.last_active = now;
3293 meta.message_count = msg_count;
3294 if meta.actor_id.is_none() {
3295 meta.actor_id = self.actor_id.read().clone();
3296 }
3297 }
3298
3299 let snapshot = self.save_state().await?;
3300 if s.supports(StorageCapability::SessionMetadata) {
3304 let metadata = self.session_metadata.read().clone();
3305 s.save_snapshot_with_metadata(session_id, &snapshot, &metadata)
3306 .await
3307 } else {
3308 s.save(session_id, &snapshot).await
3309 }
3310 }
3311 None => Err(AgentError::Config(
3312 "No storage configured. Use with_storage_config() or with_storage() first".into(),
3313 )),
3314 }
3315 }
3316
3317 pub async fn load_session(&self, session_id: &str) -> Result<bool> {
3318 let storage = self.storage.read().clone();
3319 match storage {
3320 Some(storage) => self.load_from(storage.as_ref(), session_id).await,
3321 None => Err(AgentError::Config(
3322 "No storage configured. Use with_storage_config() or with_storage() first".into(),
3323 )),
3324 }
3325 }
3326
3327 pub async fn restore_session_full(&self, session_id: &str) -> Result<usize> {
3329 self.init_storage().await?;
3330 let storage = self.storage.read().clone().ok_or_else(|| {
3331 AgentError::Config(
3332 "No storage configured. Use with_storage_config() or with_storage() first".into(),
3333 )
3334 })?;
3335 let target_parent = Self::load_session_restore(storage.as_ref(), session_id)
3336 .await?
3337 .ok_or_else(|| AgentError::Persistence(format!("Session not found: {session_id}")))?;
3338 let manifest = target_parent
3339 .snapshot
3340 .spawned_agents
3341 .clone()
3342 .unwrap_or_default();
3343
3344 let registry = self.spawner_registry.as_ref().cloned();
3345 let spawner = if manifest.is_empty() {
3346 self.spawner.as_ref().cloned()
3347 } else {
3348 Some(self.spawner.as_ref().cloned().ok_or_else(|| {
3349 AgentError::Config(
3350 "Saved session contains child agents but this runtime has no spawner".into(),
3351 )
3352 })?)
3353 };
3354 let registry = if manifest.is_empty() {
3355 registry
3356 } else {
3357 Some(registry.ok_or_else(|| {
3358 AgentError::Config(
3359 "Saved session contains child agents but this runtime has no registry".into(),
3360 )
3361 })?)
3362 };
3363
3364 let mut target_ids = HashSet::with_capacity(manifest.len());
3365 let mut prepared = Vec::with_capacity(manifest.len());
3366 for entry in manifest {
3367 if !target_ids.insert(entry.id.clone()) {
3368 return Err(AgentError::InvalidSpec(format!(
3369 "Saved child manifest contains duplicate ID: {}",
3370 entry.id
3371 )));
3372 }
3373 let spec = crate::spec::AgentSpec::from_yaml_strict(&entry.spec_yaml)?;
3374 if spec.llm.router_roles().is_some()
3375 && spec
3376 .skills
3377 .iter()
3378 .any(|skill| !matches!(skill, ai_agents_skills::SkillRef::Inline(_)))
3379 {
3380 return Err(AgentError::Config(format!(
3381 "Cannot restore child '{}': hierarchy snapshot requires prepared inline skills",
3382 entry.id
3383 )));
3384 }
3385 spawner
3386 .as_ref()
3387 .expect("non-empty manifests require a spawner")
3388 .validate_explicit_child(&entry.id, &spec)?;
3389 if let Some(live) = registry
3390 .as_ref()
3391 .and_then(|registry| registry.get_spawned(&entry.id))
3392 && (spec.llm.router_roles().is_some()
3393 || live.agent.llm_registry.router_roles().is_some())
3394 {
3395 let live_spec = live.agent.prepared_persistence_spec(&live.spec)?;
3396 if spec.routing_projection()? != live_spec.routing_projection()? {
3397 return Err(AgentError::Config(format!(
3398 "Cannot restore child '{}': hierarchical LLM routing differs from the live child",
3399 entry.id
3400 )));
3401 }
3402 }
3403 prepared.push((entry.id, spec));
3404 }
3405
3406 let current_ids = registry
3407 .as_ref()
3408 .map(|registry| {
3409 registry
3410 .list()
3411 .into_iter()
3412 .map(|info| info.id)
3413 .collect::<HashSet<_>>()
3414 })
3415 .unwrap_or_default();
3416 let removal_count = current_ids.difference(&target_ids).count();
3417 let additions = prepared
3418 .iter()
3419 .filter(|(id, _)| !current_ids.contains(id))
3420 .cloned()
3421 .collect::<Vec<_>>();
3422
3423 let mut existing = Vec::new();
3424 if let Some(registry) = registry.as_ref() {
3425 for (id, _) in prepared.iter().filter(|(id, _)| current_ids.contains(id)) {
3426 let agent = registry.get(id).ok_or_else(|| {
3427 AgentError::Config(format!("Retained child disappeared during restore: {id}"))
3428 })?;
3429 let child_storage = agent.storage().ok_or_else(|| {
3430 AgentError::Config(format!("Child '{id}' has no storage for session restore"))
3431 })?;
3432 let stored = Self::load_session_restore(child_storage.as_ref(), session_id)
3433 .await?
3434 .ok_or_else(|| {
3435 AgentError::Persistence(format!(
3436 "Child '{id}' has no saved session '{session_id}'"
3437 ))
3438 })?;
3439 existing.push((id.clone(), agent, stored));
3440 }
3441 }
3442
3443 let mut staged = Vec::with_capacity(additions.len());
3444 if !additions.is_empty() {
3445 let spawner = spawner
3446 .as_ref()
3447 .expect("restored additions require a spawner");
3448 let reservations = spawner.reserve_restore_capacity(additions.len(), removal_count)?;
3449 for ((id, spec), reservation) in additions.into_iter().zip(reservations) {
3450 let spawned = spawner
3451 .spawn_with_reserved_capacity(id.clone(), spec, reservation)
3452 .await?;
3453 let child_storage = spawned.agent.storage().ok_or_else(|| {
3454 AgentError::Config(format!("Child '{id}' has no storage for session restore"))
3455 })?;
3456 let stored = Self::load_session_restore(child_storage.as_ref(), session_id)
3457 .await?
3458 .ok_or_else(|| {
3459 AgentError::Persistence(format!(
3460 "Child '{id}' has no saved session '{session_id}'"
3461 ))
3462 })?;
3463 staged.push((spawned, stored));
3464 }
3465 } else if let Some(spawner) = spawner.as_ref() {
3466 spawner.reserve_restore_capacity(0, removal_count)?;
3467 }
3468
3469 let parent_before = self.capture_session_restore_point().await?;
3470 let mut existing_before = Vec::with_capacity(existing.len());
3471 for (id, agent, _) in &existing {
3472 existing_before.push((
3473 id.clone(),
3474 Arc::clone(agent),
3475 agent.capture_session_restore_point().await?,
3476 ));
3477 }
3478
3479 for (_, agent, stored) in &existing {
3483 if let Err(error) = agent
3484 .apply_session_restore_unchecked(session_id, stored.clone())
3485 .await
3486 {
3487 drop(staged);
3488 let rollback_errors =
3489 Self::rollback_session_restore_set(None, &existing_before).await;
3490 return Err(Self::restore_failure(error, rollback_errors));
3491 }
3492 }
3493 for (spawned, stored) in &staged {
3494 if let Err(error) = spawned
3495 .agent
3496 .apply_session_restore_unchecked(session_id, stored.clone())
3497 .await
3498 {
3499 drop(staged);
3500 let rollback_errors =
3501 Self::rollback_session_restore_set(None, &existing_before).await;
3502 return Err(Self::restore_failure(error, rollback_errors));
3503 }
3504 }
3505 if let Err(error) = self
3506 .apply_session_restore_unchecked(session_id, target_parent)
3507 .await
3508 {
3509 drop(staged);
3510 let rollback_errors =
3511 Self::rollback_session_restore_set(Some((self, &parent_before)), &existing_before)
3512 .await;
3513 return Err(Self::restore_failure(error, rollback_errors));
3514 }
3515
3516 if let Some(registry) = registry.as_ref()
3517 && let Err(error) = registry
3518 .reconcile(
3519 &target_ids,
3520 staged.into_iter().map(|(spawned, _)| spawned).collect(),
3521 )
3522 .await
3523 {
3524 let rollback_errors =
3525 Self::rollback_session_restore_set(Some((self, &parent_before)), &existing_before)
3526 .await;
3527 return Err(Self::restore_failure(error, rollback_errors));
3528 }
3529
3530 Ok(target_ids.len())
3531 }
3532
3533 pub async fn delete_session(&self, session_id: &str) -> Result<()> {
3534 let storage = self.storage.read().clone();
3535 match storage {
3536 Some(s) => s.delete(session_id).await,
3537 None => Err(AgentError::Config(
3538 "No storage configured. Use with_storage_config() or with_storage() first".into(),
3539 )),
3540 }
3541 }
3542
3543 pub async fn list_sessions(&self) -> Result<Vec<String>> {
3544 let storage = self.storage.read().clone();
3545 match storage {
3546 Some(s) => s.list_sessions().await,
3547 None => Err(AgentError::Config(
3548 "No storage configured. Use with_storage_config() or with_storage() first".into(),
3549 )),
3550 }
3551 }
3552
3553 fn estimate_tokens(&self, text: &str) -> u32 {
3554 (text.len() as f32 / 4.0).ceil() as u32
3555 }
3556
3557 fn estimate_total_tokens(&self, messages: &[ChatMessage]) -> u32 {
3558 messages
3559 .iter()
3560 .map(|m| self.estimate_tokens(&m.content))
3561 .sum()
3562 }
3563
3564 fn native_safe_prefix_at_least(messages: &[ChatMessage], required: usize) -> Result<usize> {
3566 let inspection =
3567 inspect_native_history(messages).map_err(|error| AgentError::LLM(error.to_string()))?;
3568 let has_signed_history = !inspection.exchanges().is_empty();
3569 for count in required.min(messages.len())..=messages.len() {
3570 if has_signed_history
3571 && count < messages.len()
3572 && messages[count].role != ai_agents_core::Role::User
3573 {
3574 continue;
3575 }
3576 if inspection.is_safe_prefix_len(count)
3577 && inspect_native_history(&messages[count..]).is_ok()
3578 {
3579 return Ok(count);
3580 }
3581 }
3582 Err(AgentError::LLM(
3583 "Context limits cannot remove a complete native history prefix".to_string(),
3584 ))
3585 }
3586
3587 fn truncate_context(&self, messages: &mut Vec<ChatMessage>, keep_recent: usize) -> Result<()> {
3589 if messages.len() <= keep_recent + 1 {
3590 return Ok(());
3591 }
3592 let system_msg = messages.remove(0);
3593 let required = messages.len().saturating_sub(keep_recent);
3594 let to_remove = Self::native_safe_prefix_at_least(messages, required)?;
3595 messages.drain(..to_remove);
3596 messages.insert(0, system_msg);
3597 Ok(())
3598 }
3599
3600 fn get_filter(&self, config: &FilterConfig) -> Arc<dyn MessageFilter> {
3601 match config {
3602 FilterConfig::KeepRecent(n) => Arc::new(KeepRecentFilter::new(*n)),
3603 FilterConfig::ByRole { keep_roles } => Arc::new(ByRoleFilter::new(keep_roles.clone())),
3604 FilterConfig::SkipPattern { skip_if_contains } => {
3605 Arc::new(SkipPatternFilter::new(skip_if_contains.clone()))
3606 }
3607 FilterConfig::Custom { name } => {
3608 let filters = self.message_filters.read();
3609 filters
3610 .get(name)
3611 .cloned()
3612 .unwrap_or_else(|| Arc::new(KeepRecentFilter::new(10)))
3613 }
3614 }
3615 }
3616
3617 async fn summarize_context(
3619 &self,
3620 messages: &mut Vec<ChatMessage>,
3621 summarizer_llm: Option<&str>,
3622 max_summary_tokens: u32,
3623 custom_prompt: Option<&str>,
3624 keep_recent: usize,
3625 filter: Option<&FilterConfig>,
3626 ) -> Result<()> {
3627 let system_msg = messages.remove(0);
3628
3629 let required = messages.len().saturating_sub(keep_recent);
3630 if required == 0 {
3631 messages.insert(0, system_msg);
3632 return Ok(());
3633 }
3634 let to_summarize_count = Self::native_safe_prefix_at_least(messages, required)?;
3635
3636 let recent_msgs: Vec<ChatMessage> = messages.drain(to_summarize_count..).collect();
3637 let mut to_summarize = std::mem::take(messages);
3638
3639 if let Some(filter_config) = filter {
3640 let filter = self.get_filter(filter_config);
3641 to_summarize = filter.filter(to_summarize);
3642 }
3643
3644 if to_summarize.is_empty() {
3645 *messages = recent_msgs;
3646 messages.insert(0, system_msg);
3647 return Ok(());
3648 }
3649
3650 let to_summarize = Self::readable_native_messages(to_summarize)?;
3651 let conversation_text = to_summarize
3652 .iter()
3653 .map(|m| format!("{:?}: {}", m.role, m.content))
3654 .collect::<Vec<_>>()
3655 .join("\n");
3656
3657 let default_prompt = format!(
3658 "Summarize the following conversation in under {} tokens, preserving key information:\n\n{}",
3659 max_summary_tokens, conversation_text
3660 );
3661
3662 let summary_prompt = custom_prompt
3663 .map(|p| format!("{}\n\n{}", p, conversation_text))
3664 .unwrap_or(default_prompt);
3665
3666 let summarizer = self.role_llm(
3667 ai_agents_llm::LLMRole::ContextSummarize,
3668 summarizer_llm,
3669 || {
3670 Ok(if let Some(alias) = summarizer_llm {
3671 self.llm_registry
3672 .get(alias)
3673 .map_err(|e| AgentError::Config(e.to_string()))?
3674 } else {
3675 self.llm_registry
3676 .router()
3677 .or_else(|_| self.llm_registry.default())
3678 .map_err(|e| AgentError::Config(e.to_string()))?
3679 })
3680 },
3681 )?;
3682
3683 let summary_msgs = vec![ChatMessage::user(&summary_prompt)];
3684 let response = self
3685 .observe_purpose(
3686 ObservationPurpose::Summarization,
3687 summarizer.complete(&summary_msgs, None),
3688 )
3689 .await?;
3690
3691 let summary_message = ChatMessage::system(format!(
3692 "[Previous conversation summary]\n{}",
3693 response.content
3694 ));
3695
3696 *messages = vec![system_msg, summary_message];
3697 messages.extend(recent_msgs);
3698
3699 debug!(
3700 summarized_count = to_summarize_count,
3701 kept_recent = keep_recent,
3702 "Context summarized"
3703 );
3704
3705 Ok(())
3706 }
3707
3708 fn render_system_prompt(&self) -> Result<String> {
3709 let mut context = self.build_context_with_overlays();
3710
3711 let facts_text = self.format_actor_facts_for_context();
3713 if !facts_text.is_empty() {
3714 context.insert(
3715 "actor_facts".to_string(),
3716 serde_json::Value::String(facts_text),
3717 );
3718 }
3719
3720 if let Some((key, text)) = self.format_relationship_for_context() {
3721 context.insert(key, serde_json::Value::String(text));
3722 }
3723
3724 self.template_renderer
3725 .render(&self.base_system_prompt, &context)
3726 }
3727
3728 fn canonical_unique_tool_ids(&self, ids: &[String]) -> Vec<String> {
3730 let mut seen = HashSet::new();
3731 ids.iter()
3732 .filter_map(|id| self.tools.canonical_id(id))
3733 .filter(|canonical_id| seen.insert(canonical_id.clone()))
3734 .collect()
3735 }
3736
3737 fn get_top_level_tool_ids_for_scope(&self, scope_override: Option<&[String]>) -> Vec<String> {
3739 let Some(declared) = self.declared_tool_ids.as_deref() else {
3740 return Vec::new();
3741 };
3742 let mut effective = self.canonical_unique_tool_ids(declared);
3743 if let Some(scope) = scope_override {
3744 let scope: HashSet<String> =
3745 self.canonical_unique_tool_ids(scope).into_iter().collect();
3746 effective.retain(|canonical_id| scope.contains(canonical_id));
3747 }
3748 effective
3749 }
3750
3751 async fn get_available_tool_ids(&self) -> Result<Vec<String>> {
3754 Ok(self.get_available_tool_ids_snapshot().await?.tool_ids)
3755 }
3756
3757 async fn get_available_tool_ids_snapshot(&self) -> Result<AvailableToolIdsSnapshot> {
3759 let scope_override = self.runtime_control.tool_scope_override.read().clone();
3760 self.get_available_tool_ids_snapshot_for_scope(scope_override.as_deref())
3761 .await
3762 }
3763
3764 async fn get_available_tool_ids_snapshot_for_scope(
3766 &self,
3767 scope_override: Option<&[String]>,
3768 ) -> Result<AvailableToolIdsSnapshot> {
3769 let mut available = self.get_top_level_tool_ids_for_scope(scope_override);
3770 let (state_generation, state_scopes) = self
3771 .state_machine
3772 .as_ref()
3773 .map(|state_machine| {
3774 let (generation, scopes) = state_machine.current_tool_scope_snapshot();
3775 (Some(generation), scopes)
3776 })
3777 .unwrap_or((None, Vec::new()));
3778
3779 if available.is_empty() || state_scopes.is_empty() {
3780 return Ok(AvailableToolIdsSnapshot {
3781 tool_ids: available,
3782 state_generation,
3783 });
3784 }
3785
3786 let eval_ctx = self.build_evaluation_context().await?;
3787 let llm_getter = RegistryLLMGetter {
3788 registry: self.llm_registry.clone(),
3789 };
3790 let evaluator = ConditionEvaluator::new(llm_getter);
3791
3792 for state_scope in state_scopes {
3793 if state_scope.is_empty() {
3794 available.clear();
3795 break;
3796 }
3797
3798 let mut allowed = HashSet::new();
3799 for tool_ref in &state_scope {
3800 let tool_id = tool_ref.id();
3801 let Some(canonical_id) = self.tools.canonical_id(tool_id) else {
3802 continue;
3803 };
3804 let condition_matches = if let Some(condition) = tool_ref.condition() {
3805 match evaluator.evaluate(condition, &eval_ctx).await {
3806 Ok(matches) => matches,
3807 Err(error @ AgentError::Config(_))
3808 if self.llm_registry.router_roles().is_some() =>
3809 {
3810 return Err(error);
3811 }
3812 Err(error) => {
3813 warn!(tool = tool_id, error = %error, "Error evaluating tool condition");
3814 false
3815 }
3816 }
3817 } else {
3818 true
3819 };
3820 if condition_matches {
3821 allowed.insert(canonical_id);
3822 } else {
3823 debug!(tool = tool_id, "Tool condition not met, skipping");
3824 }
3825 }
3826 available.retain(|canonical_id| allowed.contains(canonical_id));
3827 if available.is_empty() {
3828 break;
3829 }
3830 }
3831
3832 Ok(AvailableToolIdsSnapshot {
3833 tool_ids: available,
3834 state_generation,
3835 })
3836 }
3837
3838 async fn build_evaluation_context(&self) -> Result<EvaluationContext> {
3839 let context = self.build_context_with_overlays();
3840 let messages = Self::readable_native_messages(self.memory.get_messages(Some(10)).await?)?;
3841 let tool_history = self.tool_call_history.read().clone();
3842
3843 let (state_name, turn_count, previous_state) = if let Some(ref sm) = self.state_machine {
3844 (Some(sm.current()), sm.turn_count(), sm.previous())
3845 } else {
3846 (None, 0, None)
3847 };
3848
3849 Ok(EvaluationContext::default()
3850 .with_context(context)
3851 .with_state(state_name, turn_count, previous_state)
3852 .with_called_tools(tool_history)
3853 .with_messages(messages))
3854 }
3855
3856 fn record_tool_call(&self, tool_id: &str, result: Value) {
3857 self.tool_call_history.write().push(ToolCallRecord {
3858 tool_id: tool_id.to_string(),
3859 result,
3860 timestamp: chrono::Utc::now(),
3861 });
3862 }
3863
3864 async fn get_effective_system_prompt_with_persona_hooks(
3865 &self,
3866 fire_persona_hooks: bool,
3867 include_tool_prompt: bool,
3868 ) -> Result<String> {
3869 let rendered_base = self.render_system_prompt()?;
3870
3871 let persona_prefix = if let Some(ref persona) = self.persona_manager {
3872 let context = self.build_context_with_overlays();
3873 if fire_persona_hooks {
3874 let render_result = persona.render_prompt(&context)?;
3875 for content in &render_result.newly_revealed {
3876 self.hooks.on_secret_revealed(content).await;
3877 }
3878 render_result.prompt
3879 } else {
3880 persona.render_prompt_preview(&context)?
3881 }
3882 } else {
3883 String::new()
3884 };
3885
3886 if let Some(ref sm) = self.state_machine
3887 && let Some(state_def) = sm.current_definition()
3888 {
3889 let state_prompt = if let Some(ref prompt) = state_def.prompt {
3890 let context = self.build_context_with_overlays();
3891 self.template_renderer.render_with_state(
3892 prompt,
3893 &context,
3894 &sm.current(),
3895 sm.previous().as_deref(),
3896 sm.turn_count(),
3897 state_def.max_turns,
3898 )?
3899 } else {
3900 String::new()
3901 };
3902
3903 let combined = match state_def.prompt_mode {
3904 PromptMode::Append => {
3905 if state_prompt.is_empty() {
3906 rendered_base
3907 } else {
3908 format!(
3909 "{}\n\n[Current State: {}]\n{}",
3910 rendered_base,
3911 sm.current(),
3912 state_prompt
3913 )
3914 }
3915 }
3916 PromptMode::Replace => {
3917 if state_prompt.is_empty() {
3918 rendered_base
3919 } else {
3920 state_prompt
3921 }
3922 }
3923 PromptMode::Prepend => {
3924 if state_prompt.is_empty() {
3925 rendered_base
3926 } else {
3927 format!("{}\n\n{}", state_prompt, rendered_base)
3928 }
3929 }
3930 };
3931
3932 let with_persona = if persona_prefix.is_empty() {
3934 combined
3935 } else {
3936 format!("{}\n\n{}", persona_prefix, combined)
3937 };
3938
3939 if include_tool_prompt {
3940 let available_tool_ids = self.get_available_tool_ids().await?;
3941 if !available_tool_ids.is_empty() {
3942 let tools_prompt = self.tools.generate_scoped_prompt_with_mode(
3943 &available_tool_ids,
3944 None,
3945 self.parallel_tools.enabled,
3946 self.runtime_config.tool_schema_prompt_mode,
3947 );
3948 if !tools_prompt.is_empty() {
3949 return Ok(format!("{}\n\n{}", with_persona, tools_prompt));
3950 }
3951 }
3952 }
3953 return Ok(with_persona);
3954 }
3955
3956 let with_persona = if persona_prefix.is_empty() {
3958 rendered_base
3959 } else {
3960 format!("{}\n\n{}", persona_prefix, rendered_base)
3961 };
3962
3963 if include_tool_prompt {
3964 let available_tool_ids = self.get_available_tool_ids().await?;
3965 let tools_prompt = self.tools.generate_scoped_prompt_with_mode(
3966 &available_tool_ids,
3967 None,
3968 self.parallel_tools.enabled,
3969 self.runtime_config.tool_schema_prompt_mode,
3970 );
3971 if !tools_prompt.is_empty() {
3972 return Ok(format!("{}\n\n{}", with_persona, tools_prompt));
3973 }
3974 }
3975 Ok(with_persona)
3976 }
3977
3978 fn get_state_llm(&self) -> Result<Arc<dyn LLMProvider>> {
3979 if let Some(ref sm) = self.state_machine
3980 && let Some(state_def) = sm.current_definition()
3981 && let Some(ref llm_alias) = state_def.llm
3982 {
3983 return self
3984 .llm_registry
3985 .get(llm_alias)
3986 .map_err(|e| AgentError::Config(e.to_string()));
3987 }
3988 self.llm_registry
3989 .default()
3990 .map_err(|e| AgentError::Config(e.to_string()))
3991 }
3992
3993 fn get_effective_reasoning_config(&self) -> ReasoningConfig {
3994 if let Some(ref sm) = self.state_machine
3995 && let Some(state_def) = sm.current_definition()
3996 && let Some(ref state_reasoning) = state_def.reasoning
3997 {
3998 return state_reasoning.clone();
3999 }
4000 self.reasoning_config.clone()
4001 }
4002
4003 fn get_effective_reflection_config(&self) -> ReflectionConfig {
4004 if let Some(ref sm) = self.state_machine
4005 && let Some(state_def) = sm.current_definition()
4006 && let Some(ref state_reflection) = state_def.reflection
4007 {
4008 return state_reflection.clone();
4009 }
4010 self.reflection_config.clone()
4011 }
4012
4013 fn routing_reflection_mode(&self) -> ReflectionMode {
4015 self.get_effective_reflection_config().enabled
4016 }
4017
4018 fn get_skill_reasoning_config(&self, skill: &SkillDefinition) -> ReasoningConfig {
4019 skill
4020 .reasoning
4021 .clone()
4022 .unwrap_or_else(|| self.get_effective_reasoning_config())
4023 }
4024
4025 fn get_skill_reflection_config(&self, skill: &SkillDefinition) -> ReflectionConfig {
4026 skill
4027 .reflection
4028 .clone()
4029 .unwrap_or_else(|| self.get_effective_reflection_config())
4030 }
4031
4032 async fn build_disambiguation_context(&self) -> Result<DisambiguationContext> {
4034 let context_config = self
4035 .disambiguation_manager
4036 .as_ref()
4037 .map(|manager| manager.config().context.clone())
4038 .unwrap_or_default();
4039 let recent_messages = if context_config.recent_messages == 0 {
4040 Vec::new()
4041 } else {
4042 Self::readable_native_messages(
4043 self.memory
4044 .get_messages(Some(context_config.recent_messages))
4045 .await?,
4046 )?
4047 .iter()
4048 .map(|message| format!("{:?}: {}", message.role, message.content))
4049 .collect()
4050 };
4051
4052 let current_state = self.current_state().map(|s| s.to_string());
4053
4054 let state_prompt: Option<String> = self
4057 .state_machine
4058 .as_ref()
4059 .and_then(|sm| sm.current_definition())
4060 .and_then(|def| def.prompt.clone());
4061
4062 let available_tools = if context_config.include_available_tools {
4063 self.get_available_tool_ids().await?
4064 } else {
4065 Vec::new()
4066 };
4067
4068 let available_skills: Vec<String> = self.skills.iter().map(|s| s.id.clone()).collect();
4069
4070 let mut user_context = self.build_context_with_overlays();
4071 user_context.remove(DISAMBIGUATION_STATE_GENERATION_KEY);
4072 if let Some(state_generation) = self
4073 .state_machine
4074 .as_ref()
4075 .map(|state_machine| state_machine.generation())
4076 {
4077 user_context.insert(
4078 DISAMBIGUATION_STATE_GENERATION_KEY.to_string(),
4079 serde_json::json!(state_generation),
4080 );
4081 }
4082
4083 let available_intents: Vec<String> = if let Some(ref sm) = self.state_machine {
4085 sm.current_definition()
4086 .map(|def| {
4087 def.transitions
4088 .iter()
4089 .filter_map(|t| t.intent.clone())
4090 .collect()
4091 })
4092 .unwrap_or_default()
4093 } else {
4094 Vec::new()
4095 };
4096
4097 Ok(DisambiguationContext::from_agent_state(
4098 recent_messages,
4099 current_state,
4100 state_prompt,
4101 available_tools,
4102 available_skills,
4103 available_intents,
4104 user_context,
4105 ))
4106 }
4107
4108 fn get_available_skills(&self) -> Vec<&SkillDefinition> {
4109 if let Some(ref sm) = self.state_machine
4110 && let Some(state_def) = sm.current_definition()
4111 {
4112 let parent_def = sm.get_parent_definition();
4113 let effective_skills = state_def.get_effective_skills(parent_def.as_ref());
4114 if !effective_skills.is_empty() {
4115 return self
4116 .skills
4117 .iter()
4118 .filter(|s| effective_skills.contains(&&s.id))
4119 .collect();
4120 }
4121 }
4122 self.skills.iter().collect()
4123 }
4124
4125 async fn build_messages(&self) -> Result<Vec<ChatMessage>> {
4126 self.build_messages_internal(true, None, true).await
4127 }
4128
4129 async fn build_messages_for_draft(&self, user_message: &str) -> Result<Vec<ChatMessage>> {
4130 self.build_messages_internal(false, Some(user_message), true)
4131 .await
4132 }
4133
4134 async fn build_messages_internal(
4135 &self,
4136 fire_persona_hooks: bool,
4137 ephemeral_user_message: Option<&str>,
4138 include_tool_prompt: bool,
4139 ) -> Result<Vec<ChatMessage>> {
4140 let system_prompt = self
4141 .get_effective_system_prompt_with_persona_hooks(fire_persona_hooks, include_tool_prompt)
4142 .await?;
4143 let mut messages = vec![ChatMessage::system(&system_prompt)];
4144
4145 let context = self.memory.get_context().await?;
4146 let history = if let Some(ref budget) = self.memory_token_budget {
4147 context.to_llm_messages_with_allocation(&budget.allocation)
4148 } else {
4149 context.to_llm_messages()
4150 };
4151 messages.extend(history);
4152 if let Some(user_message) = ephemeral_user_message {
4153 messages.push(ChatMessage::user(user_message));
4154 }
4155
4156 let total_tokens = self.estimate_total_tokens(&messages);
4157
4158 if total_tokens > self.max_context_tokens {
4159 debug!(
4160 total = total_tokens,
4161 limit = self.max_context_tokens,
4162 "Context overflow"
4163 );
4164
4165 match &self.recovery_manager.config().llm.on_context_overflow {
4166 ContextOverflowAction::Error => {
4167 return Err(AgentError::LLM(format!(
4168 "Context overflow: {} tokens > {} limit",
4169 total_tokens, self.max_context_tokens
4170 )));
4171 }
4172 ContextOverflowAction::Truncate { keep_recent } => {
4173 self.truncate_context(&mut messages, *keep_recent)?;
4174 }
4175 ContextOverflowAction::Summarize {
4176 summarizer_llm,
4177 max_summary_tokens,
4178 custom_prompt,
4179 keep_recent,
4180 filter,
4181 } => {
4182 self.summarize_context(
4183 &mut messages,
4184 summarizer_llm.as_deref(),
4185 *max_summary_tokens,
4186 custom_prompt.as_deref(),
4187 *keep_recent,
4188 filter.as_ref(),
4189 )
4190 .await?;
4191 }
4192 }
4193 }
4194
4195 self.validate_active_native_history(&messages, true)?;
4196 Ok(messages)
4197 }
4198
4199 async fn main_tool_protocol(
4200 &self,
4201 llm: &dyn LLMProvider,
4202 ephemeral_new_turn: bool,
4203 ) -> Result<MainToolProtocol> {
4204 let mut choice = llm.configured_tool_choice();
4205 if matches!(choice.as_ref(), Some(ToolChoice::None)) {
4206 return Ok(MainToolProtocol {
4207 choice,
4208 tool_ids: Vec::new(),
4209 definitions: Vec::new(),
4210 });
4211 }
4212
4213 let mut tool_ids = self.get_available_tool_ids().await?;
4214 tool_ids.sort();
4215 tool_ids.dedup();
4216 if let Some(ToolChoice::Specific(expected)) = choice.as_ref() {
4217 let canonical = self.tools.canonical_id(expected).ok_or_else(|| {
4218 AgentError::Config(format!(
4219 "specific tool choice '{expected}' is not registered"
4220 ))
4221 })?;
4222 if canonical != *expected {
4223 return Err(AgentError::Config(format!(
4224 "specific tool choice must use canonical ID '{canonical}', not '{expected}'"
4225 )));
4226 }
4227 if !tool_ids.iter().any(|tool_id| tool_id == expected) {
4228 return Err(AgentError::Config(format!(
4229 "specific tool choice '{expected}' is outside the effective tool grant"
4230 )));
4231 }
4232 }
4233 if matches!(
4234 choice.as_ref(),
4235 Some(ToolChoice::Required | ToolChoice::Specific(_))
4236 ) && tool_ids.is_empty()
4237 {
4238 return Err(AgentError::Config(
4239 "required tool choice has no tool inside the effective grant".to_string(),
4240 ));
4241 }
4242 if !ephemeral_new_turn
4243 && let Some(configured_choice) = choice.as_ref()
4244 && matches!(
4245 configured_choice,
4246 ToolChoice::Required | ToolChoice::Specific(_)
4247 )
4248 && self
4249 .tool_choice_satisfied_in_current_turn(configured_choice, &tool_ids)
4250 .await?
4251 {
4252 choice = Some(ToolChoice::Auto);
4253 }
4254 if let Some(ToolChoice::Specific(expected)) = choice.as_ref() {
4255 tool_ids.retain(|tool_id| tool_id == expected);
4256 }
4257
4258 let definitions = tool_ids
4259 .iter()
4260 .map(|tool_id| {
4261 let tool = self.tools.get(tool_id).ok_or_else(|| {
4262 AgentError::Config(format!(
4263 "effective tool '{tool_id}' disappeared before provider exposure"
4264 ))
4265 })?;
4266 Ok(LLMToolDefinition {
4267 name: tool_id.clone(),
4268 description: tool.description().to_string(),
4269 input_schema: tool.input_schema(),
4270 })
4271 })
4272 .collect::<Result<Vec<_>>>()?;
4273
4274 Ok(MainToolProtocol {
4278 choice,
4279 tool_ids,
4280 definitions,
4281 })
4282 }
4283
4284 async fn tool_choice_satisfied_in_current_turn(
4285 &self,
4286 choice: &ToolChoice,
4287 effective_tool_ids: &[String],
4288 ) -> Result<bool> {
4289 let messages = self.memory.get_messages(None).await?;
4290 let mut saw_tool_result = false;
4291 for message in messages.iter().rev() {
4292 match message.role {
4293 ai_agents_core::Role::Tool | ai_agents_core::Role::Function => {
4294 saw_tool_result = true;
4295 }
4296 ai_agents_core::Role::Assistant if saw_tool_result => {
4297 let Some(calls) = self.parse_tool_calls(&message.content)? else {
4298 continue;
4299 };
4300 let calls_are_effective = !calls.is_empty()
4301 && calls.iter().all(|call| {
4302 self.tools
4303 .canonical_id(&call.name)
4304 .is_some_and(|canonical| effective_tool_ids.contains(&canonical))
4305 });
4306 return Ok(calls_are_effective
4307 && match choice {
4308 ToolChoice::Required => true,
4309 ToolChoice::Specific(expected) => calls.iter().all(|call| {
4310 self.tools.canonical_id(&call.name).as_deref()
4311 == Some(expected.as_str())
4312 }),
4313 _ => false,
4314 });
4315 }
4316 ai_agents_core::Role::User => return Ok(false),
4317 _ => {}
4318 }
4319 }
4320 Ok(false)
4321 }
4322
4323 fn provider_can_use_native_tools(
4324 &self,
4325 llm: &dyn LLMProvider,
4326 protocol: &MainToolProtocol,
4327 ) -> bool {
4328 let Some(choice) = protocol.choice.as_ref() else {
4329 return false;
4330 };
4331 if matches!(choice, ToolChoice::None) || protocol.definitions.is_empty() {
4332 return false;
4333 }
4334 llm.supports_tool_choice(choice)
4335 && protocol.definitions.iter().all(|definition| {
4336 !definition.name.is_empty()
4337 && definition.name.len() <= 64
4338 && definition
4339 .name
4340 .bytes()
4341 .all(|byte| byte.is_ascii_alphanumeric() || matches!(byte, b'_' | b'-'))
4342 })
4343 }
4344
4345 fn prompt_messages_for_tool_protocol(
4346 &self,
4347 messages: &[ChatMessage],
4348 protocol: &MainToolProtocol,
4349 corrective: bool,
4350 ) -> Vec<ChatMessage> {
4351 let mut messages = messages.to_vec();
4352 let Some(choice) = protocol.choice.as_ref() else {
4353 return messages;
4354 };
4355 if matches!(choice, ToolChoice::None) || protocol.tool_ids.is_empty() {
4356 return messages;
4357 }
4358
4359 let mut tool_prompt = self.tools.generate_scoped_prompt_with_mode(
4360 &protocol.tool_ids,
4361 None,
4362 self.parallel_tools.enabled,
4363 self.runtime_config.tool_schema_prompt_mode,
4364 );
4365 match choice {
4366 ToolChoice::Required => tool_prompt.push_str(
4367 "\n\nYou must call at least one listed tool before giving a final answer.",
4368 ),
4369 ToolChoice::Specific(tool_id) => tool_prompt.push_str(&format!(
4370 "\n\nYou must call the '{tool_id}' tool before giving a final answer."
4371 )),
4372 ToolChoice::Auto => {}
4373 ToolChoice::None => return messages,
4374 _ => return messages,
4375 }
4376 if let Some(system) = messages
4377 .iter_mut()
4378 .find(|message| message.role == ai_agents_core::Role::System)
4379 {
4380 system.content.push_str("\n\n");
4381 system.content.push_str(&tool_prompt);
4382 } else {
4383 messages.insert(0, ChatMessage::system(tool_prompt));
4384 }
4385 if corrective {
4386 let instruction = match choice {
4387 ToolChoice::Required => {
4388 "Your previous response did not call a required tool. Call at least one listed tool now and return only the JSON tool call."
4389 }
4390 ToolChoice::Specific(tool_id) => {
4391 messages.push(ChatMessage::user(format!(
4392 "Your previous response did not call the required '{tool_id}' tool. Call it now and return only the JSON tool call."
4393 )));
4394 return messages;
4395 }
4396 _ => return messages,
4397 };
4398 messages.push(ChatMessage::user(instruction));
4399 }
4400 messages
4401 }
4402
4403 async fn invoke_main_provider(
4404 &self,
4405 llm: Arc<dyn LLMProvider>,
4406 messages: &[ChatMessage],
4407 protocol: &MainToolProtocol,
4408 corrective: bool,
4409 ) -> std::result::Result<MainProviderResponse, LLMError> {
4410 let use_native = self.provider_can_use_native_tools(llm.as_ref(), protocol);
4411 let response = if use_native {
4412 let request = LLMToolRequest {
4413 tools: protocol.definitions.clone(),
4414 choice: protocol
4415 .choice
4416 .clone()
4417 .expect("native tool requests require an explicit choice"),
4418 };
4419 self.observe_purpose(
4420 ObservationPurpose::MainResponse,
4421 llm.complete_with_tools(messages, None, &request),
4422 )
4423 .await?
4424 } else {
4425 let prompt_messages =
4426 self.prompt_messages_for_tool_protocol(messages, protocol, corrective);
4427 self.observe_purpose(
4428 ObservationPurpose::MainResponse,
4429 llm.complete(&prompt_messages, None),
4430 )
4431 .await?
4432 };
4433 Ok(MainProviderResponse {
4434 response,
4435 used_native_tools: use_native,
4436 })
4437 }
4438
4439 async fn complete_main_attempt_with_recovery(
4440 &self,
4441 llm: Arc<dyn LLMProvider>,
4442 messages: &[ChatMessage],
4443 protocol: &MainToolProtocol,
4444 corrective: bool,
4445 ) -> Result<MainProviderResponse> {
4446 let primary_result = self
4448 .recovery_manager
4449 .with_llm_retry(
4450 "llm_call",
4451 None,
4452 || {
4453 let llm = Arc::clone(&llm);
4454 async move {
4455 self.invoke_main_provider(llm, messages, protocol, corrective)
4456 .await
4457 }
4458 },
4459 |error| llm.is_terminal_error(error),
4460 )
4461 .await;
4462
4463 match primary_result {
4464 Ok(response) => Ok(response),
4465 Err(ai_agents_recovery::RetryFailure::Terminal { error, .. }) => {
4466 Err(AgentError::LLM(error.to_string()))
4467 }
4468 Err(failure) => {
4469 let primary_error = AgentError::LLM(failure.into_error().to_string());
4470 match &self.recovery_manager.config().llm.on_failure {
4471 LLMFailureAction::FallbackLlm { fallback_llm } => {
4472 let fallback = self.llm_registry.get(fallback_llm).map_err(|error| {
4473 AgentError::Config(format!(
4474 "Fallback LLM '{fallback_llm}' not found: {error}"
4475 ))
4476 })?;
4477 self.invoke_main_provider(fallback, messages, protocol, corrective)
4478 .await
4479 .map_err(|error| AgentError::LLM(error.to_string()))
4480 }
4481 LLMFailureAction::FallbackResponse { message } => {
4482 if matches!(
4483 protocol.choice.as_ref(),
4484 Some(ToolChoice::Required | ToolChoice::Specific(_))
4485 ) {
4486 Err(AgentError::LLM(format!(
4487 "Required tool selection failed and cannot be satisfied by a static fallback response: {primary_error}"
4488 )))
4489 } else {
4490 Ok(MainProviderResponse {
4491 response: LLMResponse::new(message.clone(), FinishReason::Stop),
4492 used_native_tools: false,
4493 })
4494 }
4495 }
4496 LLMFailureAction::Error => Err(primary_error),
4497 }
4498 }
4499 }
4500 }
4501
4502 fn normalize_main_provider_response(
4503 &self,
4504 mut response: LLMResponse,
4505 protocol: &MainToolProtocol,
4506 ) -> Result<(LLMResponse, bool)> {
4507 let provider_state = response
4508 .take_provider_state()
4509 .map_err(|error| AgentError::LLM(error.to_string()))?;
4510 let native_calls = response
4511 .tool_calls()
4512 .map_err(|error| AgentError::LLM(error.to_string()))?;
4513 let calls = match native_calls {
4514 Some(calls) => {
4515 response.content = encode_native_tool_call_markers(&calls, provider_state.as_ref())
4516 .map_err(|error| AgentError::LLM(error.to_string()))?;
4517 Some(calls)
4518 }
4519 None if provider_state.is_some() => {
4520 return Err(AgentError::LLM(
4521 "Provider returned replay state without native tool calls".to_string(),
4522 ));
4523 }
4524 None if !matches!(protocol.choice.as_ref(), Some(ToolChoice::None)) => {
4525 self.parse_tool_calls(response.content.trim())?
4526 }
4527 None => None,
4528 };
4529
4530 if protocol.choice.is_some()
4531 && let Some(calls) = calls.as_ref()
4532 && calls.iter().any(|call| {
4533 self.tools
4534 .canonical_id(&call.name)
4535 .is_none_or(|canonical| !protocol.tool_ids.contains(&canonical))
4536 })
4537 {
4538 return Err(AgentError::LLM(
4539 "Provider returned a tool call outside the effective grant".to_string(),
4540 ));
4541 }
4542
4543 let compliant = match protocol.choice.as_ref() {
4544 Some(ToolChoice::Required) => calls.as_ref().is_some_and(|calls| !calls.is_empty()),
4545 Some(ToolChoice::Specific(expected)) => calls.as_ref().is_some_and(|calls| {
4546 !calls.is_empty()
4547 && calls.iter().all(|call| {
4548 self.tools.canonical_id(&call.name).as_deref() == Some(expected.as_str())
4549 })
4550 }),
4551 _ => true,
4552 };
4553 Ok((response, compliant))
4554 }
4555
4556 async fn complete_main_llm_with_recovery(
4557 &self,
4558 llm: Arc<dyn LLMProvider>,
4559 messages: &[ChatMessage],
4560 protocol: &MainToolProtocol,
4561 ) -> Result<LLMResponse> {
4562 let first = self
4563 .complete_main_attempt_with_recovery(Arc::clone(&llm), messages, protocol, false)
4564 .await?;
4565 let (response, compliant) =
4566 self.normalize_main_provider_response(first.response, protocol)?;
4567 if compliant {
4568 return Ok(response);
4569 }
4570 if first.used_native_tools {
4571 return Err(AgentError::LLM(
4572 "Provider returned no compliant native call for required tool choice".to_string(),
4573 ));
4574 }
4575
4576 let corrected = self
4577 .complete_main_attempt_with_recovery(llm, messages, protocol, true)
4578 .await?;
4579 let (response, compliant) =
4580 self.normalize_main_provider_response(corrected.response, protocol)?;
4581 if compliant {
4582 return Ok(response);
4583 }
4584 Err(AgentError::LLM(
4585 "Provider returned no compliant tool call after one corrective retry".to_string(),
4586 ))
4587 }
4588
4589 async fn open_main_stream_with_recovery(
4596 &self,
4597 llm: Arc<dyn LLMProvider>,
4598 messages: &[ChatMessage],
4599 protocol: &MainToolProtocol,
4600 ) -> Result<MainStreamSource> {
4601 debug_assert!(
4602 protocol.choice.is_none(),
4603 "streaming raw path must not run with explicit tool choice"
4604 );
4605 let primary = self
4606 .recovery_manager
4607 .with_llm_retry(
4608 "llm_stream_open",
4609 None,
4610 || {
4611 let llm = Arc::clone(&llm);
4612 async move {
4613 self.observe_purpose(
4614 ObservationPurpose::MainResponse,
4615 llm.complete_stream(messages, None),
4616 )
4617 .await
4618 }
4619 },
4620 |error| llm.is_terminal_error(error),
4621 )
4622 .await;
4623
4624 match primary {
4625 Ok(stream) => Ok(MainStreamSource::Stream(stream)),
4626 Err(ai_agents_recovery::RetryFailure::Terminal { error, .. }) => {
4627 Err(AgentError::LLM(error.to_string()))
4628 }
4629 Err(failure) => {
4630 let primary_error = AgentError::LLM(failure.into_error().to_string());
4631 match &self.recovery_manager.config().llm.on_failure {
4632 LLMFailureAction::FallbackLlm { fallback_llm } => {
4633 let fallback = self.llm_registry.get(fallback_llm).map_err(|error| {
4634 AgentError::Config(format!(
4635 "Fallback LLM '{fallback_llm}' not found: {error}"
4636 ))
4637 })?;
4638 if fallback.supports(LLMFeature::Streaming) {
4639 let stream = self
4640 .observe_purpose(
4641 ObservationPurpose::MainResponse,
4642 fallback.complete_stream(messages, None),
4643 )
4644 .await
4645 .map_err(|error| AgentError::LLM(error.to_string()))?;
4646 Ok(MainStreamSource::Stream(stream))
4647 } else {
4648 let response = self
4649 .observe_purpose(
4650 ObservationPurpose::MainResponse,
4651 fallback.complete(messages, None),
4652 )
4653 .await
4654 .map_err(|error| AgentError::LLM(error.to_string()))?;
4655 Ok(MainStreamSource::StaticResponse(response.content))
4656 }
4657 }
4658 LLMFailureAction::FallbackResponse { message } => {
4659 Ok(MainStreamSource::StaticResponse(message.clone()))
4660 }
4661 LLMFailureAction::Error => Err(primary_error),
4662 }
4663 }
4664 }
4665 }
4666
4667 fn main_stream_must_buffer(
4676 &self,
4677 reasoning_mode: &ReasoningMode,
4678 protocol: &MainToolProtocol,
4679 ) -> bool {
4680 protocol.choice.is_some()
4681 || self.get_effective_reflection_config().requires_evaluation()
4682 || matches!(reasoning_mode, ReasoningMode::CoT | ReasoningMode::React)
4683 }
4684
4685 fn is_native_tool_call_content(content: &str) -> Result<bool> {
4687 decode_native_tool_call_markers(content)
4688 .map(|batch| batch.is_some())
4689 .map_err(|error| AgentError::LLM(error.to_string()))
4690 }
4691
4692 fn tool_result_message(
4694 tool_call: &ToolCall,
4695 output: &str,
4696 native_tool_call: bool,
4697 ) -> Result<ChatMessage> {
4698 if !native_tool_call {
4699 return Ok(ChatMessage::function(&tool_call.name, output));
4700 }
4701 let output = serde_json::from_str::<serde_json::Value>(output)
4702 .unwrap_or_else(|_| serde_json::Value::String(output.to_string()));
4703 let content = encode_native_tool_result_marker(tool_call, output)
4704 .map_err(|error| AgentError::LLM(error.to_string()))?;
4705 Ok(ChatMessage::function(&tool_call.name, content))
4706 }
4707
4708 fn remember_active_native_exchange(&self, content: &str) -> Result<()> {
4710 let Some(batch) = decode_native_tool_call_markers(content)
4711 .map_err(|error| AgentError::LLM(error.to_string()))?
4712 else {
4713 return Ok(());
4714 };
4715 let Some(state) = batch.provider_state() else {
4716 return Ok(());
4717 };
4718 let expected = ActiveNativeExchange {
4719 exchange_id: state.exchange_id().to_string(),
4720 call_ids: batch.calls().iter().map(|call| call.id.clone()).collect(),
4721 };
4722 let mut active = self.active_native_exchanges.write();
4723 if let Some(existing) = active
4724 .iter()
4725 .find(|existing| existing.exchange_id == expected.exchange_id)
4726 {
4727 if existing.call_ids != expected.call_ids {
4728 return Err(AgentError::LLM(format!(
4729 "Active native exchange '{}' changed its call identities",
4730 expected.exchange_id
4731 )));
4732 }
4733 } else {
4734 active.push(expected);
4735 }
4736 Ok(())
4737 }
4738
4739 fn validate_active_native_history(
4741 &self,
4742 messages: &[ChatMessage],
4743 require_complete: bool,
4744 ) -> Result<()> {
4745 let expected = self.active_native_exchanges.read().clone();
4746 if expected.is_empty() {
4747 return Ok(());
4748 }
4749 let inspection =
4750 inspect_native_history(messages).map_err(|error| AgentError::LLM(error.to_string()))?;
4751 let expected_count = expected.len();
4752 for (index, expected) in expected.iter().enumerate() {
4753 let Some(exchange) = inspection
4754 .exchanges()
4755 .iter()
4756 .find(|exchange| exchange.state().exchange_id() == expected.exchange_id)
4757 else {
4758 return Err(AgentError::LLM(format!(
4759 "Active native exchange '{}' was removed before provider continuation",
4760 expected.exchange_id
4761 )));
4762 };
4763 let must_be_complete = require_complete || index + 1 < expected_count;
4764 if exchange.call_ids() != expected.call_ids
4765 || (must_be_complete && !exchange.is_complete())
4766 {
4767 return Err(AgentError::LLM(format!(
4768 "Active native exchange '{}' is incomplete before provider continuation",
4769 expected.exchange_id
4770 )));
4771 }
4772 }
4773 Ok(())
4774 }
4775
4776 async fn remember_committed_native_exchange(&self, content: &str) -> Result<()> {
4778 self.remember_active_native_exchange(content)?;
4779 if !self.active_native_exchanges.read().is_empty() {
4780 let messages = self.memory.get_messages(None).await?;
4781 self.validate_active_native_history(&messages, false)?;
4782 }
4783 Ok(())
4784 }
4785
4786 fn readable_native_messages(mut messages: Vec<ChatMessage>) -> Result<Vec<ChatMessage>> {
4788 for message in &mut messages {
4789 if matches!(
4790 message.role,
4791 ai_agents_core::Role::Assistant
4792 | ai_agents_core::Role::Tool
4793 | ai_agents_core::Role::Function
4794 ) {
4795 message.content = native_readable_projection(&message.content)
4796 .map_err(|error| AgentError::LLM(error.to_string()))?;
4797 }
4798 }
4799 Ok(messages)
4800 }
4801
4802 fn parse_main_tool_calls(
4804 &self,
4805 content: &str,
4806 protocol: &MainToolProtocol,
4807 ) -> Result<Option<Vec<ToolCall>>> {
4808 if matches!(protocol.choice.as_ref(), Some(ToolChoice::None)) {
4809 Ok(None)
4810 } else {
4811 self.parse_tool_calls(content)
4812 }
4813 }
4814
4815 fn parse_tool_calls(&self, content: &str) -> Result<Option<Vec<ToolCall>>> {
4817 if let Some(batch) = decode_native_tool_call_markers(content)
4818 .map_err(|error| AgentError::LLM(error.to_string()))?
4819 {
4820 return Ok(Some(batch.into_parts().0));
4821 }
4822 if let Ok(parsed) = serde_json::from_str::<serde_json::Value>(content) {
4824 if let Some(arr) = parsed.as_array() {
4826 let calls: Vec<ToolCall> = arr
4827 .iter()
4828 .filter_map(|v| self.extract_tool_call_from_value(v))
4829 .collect();
4830 if !calls.is_empty() {
4831 return Ok(Some(calls));
4832 }
4833 }
4834 if let Some(tool_call) = self.extract_tool_call_from_value(&parsed) {
4836 return Ok(Some(vec![tool_call]));
4837 }
4838 }
4839
4840 if let Some(json_str) = self.extract_json_from_content(content)
4842 && let Ok(parsed) = serde_json::from_str::<serde_json::Value>(&json_str)
4843 {
4844 if let Some(arr) = parsed.as_array() {
4846 let calls: Vec<ToolCall> = arr
4847 .iter()
4848 .filter_map(|v| self.extract_tool_call_from_value(v))
4849 .collect();
4850 if !calls.is_empty() {
4851 return Ok(Some(calls));
4852 }
4853 }
4854 if let Some(tool_call) = self.extract_tool_call_from_value(&parsed) {
4856 return Ok(Some(vec![tool_call]));
4857 }
4858 }
4859
4860 Ok(None)
4861 }
4862
4863 fn extract_tool_call_from_value(&self, parsed: &serde_json::Value) -> Option<ToolCall> {
4864 if let Some(tool_name) = parsed.get("tool").and_then(|v| v.as_str()) {
4865 let arguments = parsed
4866 .get("arguments")
4867 .cloned()
4868 .unwrap_or(serde_json::json!({}));
4869 return Some(ToolCall {
4870 id: parsed
4871 .get("id")
4872 .and_then(|value| value.as_str())
4873 .filter(|id| !id.is_empty())
4874 .map(str::to_string)
4875 .unwrap_or_else(|| uuid::Uuid::new_v4().to_string()),
4876 name: tool_name.to_string(),
4877 arguments,
4878 });
4879 }
4880 None
4881 }
4882
4883 fn extract_json_from_content(&self, content: &str) -> Option<String> {
4885 if let Some(result) = self.extract_json_array_from_content(content) {
4887 return Some(result);
4888 }
4889 self.extract_json_object_from_content(content)
4890 }
4891
4892 fn extract_json_array_from_content(&self, content: &str) -> Option<String> {
4894 let start = content.find('[')?;
4895 let content_from_start = &content[start..];
4896
4897 let mut depth = 0;
4898 let mut end = 0;
4899 for (i, ch) in content_from_start.char_indices() {
4900 match ch {
4901 '[' => depth += 1,
4902 ']' => {
4903 depth -= 1;
4904 if depth == 0 {
4905 end = i + 1;
4906 break;
4907 }
4908 }
4909 _ => {}
4910 }
4911 }
4912
4913 if end > 0 {
4914 let json_str = &content_from_start[..end];
4915 if json_str.contains("\"tool\"") {
4917 return Some(json_str.to_string());
4918 }
4919 }
4920
4921 None
4922 }
4923
4924 fn extract_json_object_from_content(&self, content: &str) -> Option<String> {
4926 let start = content.find('{')?;
4927 let content_from_start = &content[start..];
4928
4929 let mut depth = 0;
4931 let mut end = 0;
4932 for (i, ch) in content_from_start.char_indices() {
4933 match ch {
4934 '{' => depth += 1,
4935 '}' => {
4936 depth -= 1;
4937 if depth == 0 {
4938 end = i + 1;
4939 break;
4940 }
4941 }
4942 _ => {}
4943 }
4944 }
4945
4946 if end > 0 {
4947 let json_str = &content_from_start[..end];
4948 if json_str.contains("\"tool\"") {
4950 return Some(json_str.to_string());
4951 }
4952 }
4953
4954 None
4955 }
4956
4957 #[allow(clippy::too_many_arguments)]
4961 fn record_from_parts(
4962 &self,
4963 request: &ToolExecutionRequest,
4964 canonical_id: String,
4965 executed_arguments: Value,
4966 started_at: chrono::DateTime<chrono::Utc>,
4967 start: Instant,
4968 executed: bool,
4969 success: bool,
4970 output: String,
4971 metadata: HashMap<String, Value>,
4972 policy: ToolPolicyDecisionRecord,
4973 approval: Option<ToolApprovalRecord>,
4974 timed_out: bool,
4975 output_truncated: bool,
4976 ) -> ToolExecutionRecord {
4977 let versions = ToolDecisionVersions {
4978 policy: self.active_tool_security().policy_version(),
4979 registry: self.tools.version(),
4980 runtime_control: self.runtime_control.version.load(Ordering::SeqCst),
4981 state: self
4982 .state_machine
4983 .as_ref()
4984 .map(|state_machine| state_machine.generation()),
4985 };
4986 self.record_from_parts_at(
4987 request,
4988 canonical_id,
4989 executed_arguments,
4990 started_at,
4991 start,
4992 executed,
4993 success,
4994 output,
4995 metadata,
4996 policy,
4997 approval,
4998 timed_out,
4999 output_truncated,
5000 versions,
5001 )
5002 }
5003
5004 #[allow(clippy::too_many_arguments)]
5006 fn record_from_parts_at(
5007 &self,
5008 request: &ToolExecutionRequest,
5009 canonical_id: String,
5010 executed_arguments: Value,
5011 started_at: chrono::DateTime<chrono::Utc>,
5012 start: Instant,
5013 executed: bool,
5014 success: bool,
5015 output: String,
5016 metadata: HashMap<String, Value>,
5017 policy: ToolPolicyDecisionRecord,
5018 approval: Option<ToolApprovalRecord>,
5019 timed_out: bool,
5020 output_truncated: bool,
5021 versions: ToolDecisionVersions,
5022 ) -> ToolExecutionRecord {
5023 ToolExecutionRecord {
5024 call_id: request.call_id.clone(),
5025 requested_name: request.requested_name.clone(),
5026 canonical_id,
5027 source: request.source.clone(),
5028 arguments: request.arguments.clone(),
5029 executed_arguments,
5030 policy_version: versions.policy,
5031 registry_version: versions.registry,
5032 runtime_config_version: versions.runtime_control,
5033 executed,
5034 success,
5035 output,
5036 metadata,
5037 policy,
5038 approval,
5039 started_at,
5040 duration_ms: start.elapsed().as_millis() as u64,
5041 timed_out,
5042 cancelled: false,
5043 cancellation_reason: None,
5044 output_truncated,
5045 }
5046 }
5047
5048 async fn finish_tool_record(&self, record: &ToolExecutionRecord) {
5050 let result = ToolResult {
5051 success: record.success,
5052 output: record.model_output_string(),
5053 metadata: if record.metadata.is_empty() {
5054 None
5055 } else {
5056 Some(record.metadata.clone())
5057 },
5058 };
5059 self.hooks
5060 .on_tool_complete(&record.canonical_id, &result, record.duration_ms)
5061 .await;
5062 self.hooks.on_tool_execution_record(record).await;
5063 self.record_tool_call(&record.canonical_id, record.model_output_value());
5064 if !record.success {
5065 self.hooks
5066 .on_error(&AgentError::Tool(record.output.clone()))
5067 .await;
5068 }
5069 }
5070
5071 async fn finish_tool_record_after_resource_guards(
5073 &self,
5074 resource_guards: ToolResourceGuards,
5075 record: &ToolExecutionRecord,
5076 ) {
5077 drop(resource_guards);
5078 self.finish_tool_record(record).await;
5079 }
5080
5081 fn validated_tool_timeout(timeout_ms: u64) -> Result<ValidatedToolTimeout> {
5085 if timeout_ms > MAX_TOOL_TIMEOUT_MS {
5086 return Err(AgentError::Config(format!(
5087 "effective tool timeout_ms must be no greater than {MAX_TOOL_TIMEOUT_MS} milliseconds"
5088 )));
5089 }
5090 let timer = Duration::from_millis(timeout_ms);
5091 let deadline_delta = chrono::Duration::from_std(timer).map_err(|_| {
5092 AgentError::Config(format!(
5093 "effective tool timeout_ms cannot be represented as a UTC deadline: {timeout_ms}"
5094 ))
5095 })?;
5096 Ok(ValidatedToolTimeout {
5097 timer,
5098 deadline_delta,
5099 })
5100 }
5101
5102 fn effective_tool_limits(
5106 security_engine: &ToolSecurityEngine,
5107 canonical_id: &str,
5108 safety: &ToolSafetyMetadata,
5109 classification: &ToolCallClassification,
5110 recovery_timeout_ms: Option<u64>,
5111 ) -> Result<(ToolExecutionLimits, ValidatedToolTimeout)> {
5112 if let Some(timeout_ms) = classification.timeout_ms {
5113 Self::validated_tool_timeout(timeout_ms)?;
5114 }
5115 if let Some(timeout_ms) = recovery_timeout_ms {
5116 Self::validated_tool_timeout(timeout_ms)?;
5117 }
5118
5119 let mut limits = security_engine.effective_limits(canonical_id, safety, classification);
5120 if let Some(recovery_timeout_ms) = recovery_timeout_ms {
5121 limits.timeout_ms = Some(limits.timeout_ms.map_or(recovery_timeout_ms, |timeout_ms| {
5122 timeout_ms.min(recovery_timeout_ms)
5123 }));
5124 }
5125 let timeout_ms = limits
5126 .timeout_ms
5127 .unwrap_or_else(|| security_engine.get_tool_timeout(canonical_id));
5128 let timeout = Self::validated_tool_timeout(timeout_ms)?;
5129 Ok((limits, timeout))
5130 }
5131
5132 async fn execute_resolved_tool_once(
5134 &self,
5135 tool: Arc<dyn ai_agents_core::Tool>,
5136 args: Value,
5137 mut ctx: ToolExecutionContext,
5138 timeout: ValidatedToolTimeout,
5139 ) -> Result<(ToolResult, bool, bool, bool)> {
5140 if self.runtime_control.emergency_deny.load(Ordering::SeqCst) {
5141 return Ok((
5142 ToolResult::error("Tool execution cancelled by runtime control"),
5143 false,
5144 true,
5145 false,
5146 ));
5147 }
5148 ctx.deadline = Some(
5153 chrono::Utc::now()
5154 .checked_add_signed(timeout.deadline_delta)
5155 .ok_or_else(|| {
5156 AgentError::Config(
5157 "effective tool timeout_ms exceeds the current UTC deadline range"
5158 .to_string(),
5159 )
5160 })?,
5161 );
5162 let invoked = Arc::new(AtomicBool::new(false));
5166 let invoked_by_future = Arc::clone(&invoked);
5167 let actor_context = current_turn_actor_context();
5168 let future = async move {
5169 invoked_by_future.store(true, Ordering::SeqCst);
5170 if let Some(actor_context) = actor_context {
5171 scope_actor_context(actor_context, tool.execute(args, ctx)).await
5172 } else {
5173 tool.execute(args, ctx).await
5174 }
5175 };
5176 tokio::pin!(future);
5177 let timer = tokio::time::sleep(timeout.timer);
5178 tokio::pin!(timer);
5179 let mut cancel_tick = tokio::time::interval(std::time::Duration::from_millis(50));
5180
5181 loop {
5182 tokio::select! {
5183 result = &mut future => return Ok((result, false, false, true)),
5184 _ = &mut timer => {
5185 return Ok((
5186 ToolResult::error("Tool execution timed out"),
5187 true,
5188 false,
5189 invoked.load(Ordering::SeqCst),
5190 ));
5191 }
5192 _ = cancel_tick.tick() => {
5193 if self.runtime_control.emergency_deny.load(Ordering::SeqCst) {
5194 return Ok((
5195 ToolResult::error("Tool execution cancelled by runtime control"),
5196 false,
5197 true,
5198 invoked.load(Ordering::SeqCst),
5199 ));
5200 }
5201 }
5202 }
5203 }
5204 }
5205
5206 fn truncate_tool_output(output: String, max_chars: Option<usize>) -> (String, bool) {
5208 let Some(max_chars) = max_chars else {
5209 return (output, false);
5210 };
5211 let mut chars = output.chars();
5212 let truncated: String = chars.by_ref().take(max_chars).collect();
5213 if chars.next().is_some() {
5214 (truncated, true)
5215 } else {
5216 (output, false)
5217 }
5218 }
5219
5220 async fn acquire_tool_resource_locks(&self, keys: &[String]) -> Option<ToolResourceGuards> {
5222 let locks = {
5223 let mut table = self.resource_locks.write();
5224 table.retain(|_, lock| lock.strong_count() > 0);
5225 keys.iter()
5226 .map(|key| {
5227 if let Some(lock) = table.get(key).and_then(Weak::upgrade) {
5228 lock
5229 } else {
5230 let lock = Arc::new(tokio::sync::Mutex::new(()));
5231 table.insert(key.clone(), Arc::downgrade(&lock));
5232 lock
5233 }
5234 })
5235 .collect::<Vec<_>>()
5236 };
5237 let mut resource_guards = ToolResourceGuards {
5238 guards: Vec::with_capacity(locks.len()),
5239 locks: Arc::clone(&self.resource_locks),
5240 };
5241 let mut locks = locks.into_iter();
5242 while let Some(lock) = locks.next() {
5243 let mut lock = Box::pin(lock.lock_owned());
5244 loop {
5245 tokio::select! {
5246 guard = &mut lock => {
5247 resource_guards.guards.push(guard);
5248 break;
5249 }
5250 _ = tokio::time::sleep(std::time::Duration::from_millis(10)) => {
5251 if self.runtime_control.emergency_deny.load(Ordering::SeqCst) {
5252 drop(lock);
5253 drop(locks);
5254 drop(resource_guards);
5255 return None;
5256 }
5257 }
5258 }
5259 }
5260 }
5261 Some(resource_guards)
5262 }
5263
5264 async fn run_tool_with_retries(
5268 &self,
5269 canonical_id: &str,
5270 tool: Arc<dyn ai_agents_core::Tool>,
5271 args: Value,
5272 ctx: ToolExecutionContext,
5273 timeout: ValidatedToolTimeout,
5274 max_retries: u32,
5275 ) -> Result<(ToolResult, bool, bool, bool)> {
5276 let max_retries = if ctx.classification.safely_retryable {
5277 max_retries
5278 } else {
5279 0
5280 };
5281 let mut attempts = 0;
5282 let mut invoked = false;
5283 loop {
5284 let (result, timed_out, cancelled, attempt_invoked) = self
5285 .execute_resolved_tool_once(tool.clone(), args.clone(), ctx.clone(), timeout)
5286 .await?;
5287 invoked |= attempt_invoked;
5288 if result.success || timed_out || cancelled || attempts >= max_retries {
5289 return Ok((result, timed_out, cancelled, invoked));
5290 }
5291 attempts += 1;
5292 warn!(tool = %canonical_id, attempt = attempts, error = %result.output, "Retrying failed tool call");
5293 }
5294 }
5295
5296 fn host_tool_unavailability(&self, canonical_id: &str) -> Option<(&'static str, &'static str)> {
5298 match canonical_id {
5299 "command" if !self.tools.command_runner_available() => Some((
5300 "Command runner is unavailable",
5301 "command runner is unavailable",
5302 )),
5303 "diagnostics" if !self.tools.diagnostics_available() => Some((
5304 "Diagnostics provider is unavailable",
5305 "diagnostics provider is unavailable",
5306 )),
5307 "web_search" if !self.tools.web_search_available() => Some((
5308 "Web search provider is unavailable",
5309 "web search provider is unavailable",
5310 )),
5311 _ => None,
5312 }
5313 }
5314
5315 fn execute_tool_record(
5317 &self,
5318 request: ToolExecutionRequest,
5319 ) -> Pin<Box<dyn Future<Output = Result<ToolExecutionRecord>> + Send + '_>> {
5320 Box::pin(self.execute_tool_record_inner(request, ToolFallbackState::default()))
5321 }
5322
5323 async fn execute_tool_record_inner(
5327 &self,
5328 request: ToolExecutionRequest,
5329 fallback_state: ToolFallbackState,
5330 ) -> Result<ToolExecutionRecord> {
5331 let started_at = chrono::Utc::now();
5332 let start = Instant::now();
5333 info!(tool = %request.requested_name, args = %request.arguments, "Executing tool");
5334
5335 if self.runtime_control.emergency_deny.load(Ordering::SeqCst) {
5336 let record = self.record_from_parts(
5337 &request,
5338 request.requested_name.clone(),
5339 request.arguments.clone(),
5340 started_at,
5341 start,
5342 false,
5343 false,
5344 "Tool execution is disabled by runtime control".to_string(),
5345 HashMap::new(),
5346 ToolPolicyDecisionRecord::deny("runtime emergency deny is enabled"),
5347 None,
5348 false,
5349 false,
5350 );
5351 self.finish_tool_record(&record).await;
5352 return Ok(record);
5353 }
5354
5355 let Some(resolved) = self.tools.resolve(&request.requested_name) else {
5356 let record = self.record_from_parts(
5357 &request,
5358 request.requested_name.clone(),
5359 request.arguments.clone(),
5360 started_at,
5361 start,
5362 false,
5363 false,
5364 format!("Tool '{}' is unavailable", request.requested_name),
5365 HashMap::new(),
5366 ToolPolicyDecisionRecord::unavailable(format!(
5367 "Tool '{}' is not registered",
5368 request.requested_name
5369 )),
5370 None,
5371 false,
5372 false,
5373 );
5374 self.finish_tool_record(&record).await;
5375 return Ok(record);
5376 };
5377
5378 let canonical_id = resolved.identity.canonical_id.clone();
5379
5380 let initial_scope_snapshot = self.get_available_tool_ids_snapshot().await?;
5381 if !initial_scope_snapshot
5382 .tool_ids
5383 .iter()
5384 .any(|id| id == &canonical_id)
5385 {
5386 let record = self.record_from_parts(
5387 &request,
5388 canonical_id.clone(),
5389 request.arguments.clone(),
5390 started_at,
5391 start,
5392 false,
5393 false,
5394 format!(
5395 "Tool '{}' is not available in the current scope",
5396 canonical_id
5397 ),
5398 HashMap::new(),
5399 ToolPolicyDecisionRecord::deny(format!(
5400 "Tool '{}' is not granted by the current top-level and state tool scope",
5401 canonical_id
5402 )),
5403 None,
5404 false,
5405 false,
5406 );
5407 self.finish_tool_record(&record).await;
5408 return Ok(record);
5409 }
5410
5411 let approval_control_snapshot = self.runtime_safety_snapshot();
5412 let security_engine = approval_control_snapshot.tool_security.clone();
5413 if let Some(reason) = fallback_state.rejection_reason(&canonical_id) {
5414 let mut metadata = HashMap::new();
5415 metadata.insert(
5416 "fallback_chain".to_string(),
5417 serde_json::to_value(&fallback_state.visited_canonical_ids).unwrap_or(Value::Null),
5418 );
5419 let record = self.record_from_parts(
5420 &request,
5421 canonical_id,
5422 request.arguments.clone(),
5423 started_at,
5424 start,
5425 false,
5426 false,
5427 format!("Denied: {reason}"),
5428 metadata,
5429 ToolPolicyDecisionRecord::deny(reason),
5430 None,
5431 false,
5432 false,
5433 );
5434 self.finish_tool_record(&record).await;
5435 return Ok(record);
5436 }
5437 let admitted_canonical_id = canonical_id.clone();
5438 let fallback_state = fallback_state.with_current(canonical_id.clone());
5439 let bindings = resolved.tool.policy_bindings();
5440 let mut executed_arguments = security_engine.prepare_tool_arguments_with_bindings(
5441 &canonical_id,
5442 &request.arguments,
5443 &bindings,
5444 );
5445 let mut metadata = HashMap::new();
5446 let safety = resolved.tool.safety_metadata();
5447 let classification = resolved.tool.classify_call(&executed_arguments);
5448 let initial_recovery_timeout_ms = self.recovery_manager.get_tool_timeout(&canonical_id);
5449 let (limits, _) = Self::effective_tool_limits(
5450 &security_engine,
5451 &canonical_id,
5452 &safety,
5453 &classification,
5454 initial_recovery_timeout_ms,
5455 )?;
5456 self.hooks
5457 .on_tool_start(&canonical_id, &executed_arguments)
5458 .await;
5459 metadata.insert(
5460 "classification".to_string(),
5461 serde_json::to_value(&classification).unwrap_or(Value::Null),
5462 );
5463 metadata.insert(
5464 "effective_limits".to_string(),
5465 serde_json::to_value(&limits).unwrap_or(Value::Null),
5466 );
5467 let policy_snapshot = security_engine.policy_snapshot(&canonical_id);
5468 if !policy_snapshot.is_null() {
5469 metadata.insert("policy_snapshot".to_string(), policy_snapshot.clone());
5470 }
5471
5472 let mut approval_record = Some(ToolApprovalRecord {
5473 status: ToolApprovalStatus::NotRequired,
5474 reason: None,
5475 modified_arguments: None,
5476 });
5477
5478 let mut security_result = security_engine
5479 .validate_tool_execution_with_bindings(&canonical_id, &executed_arguments, &bindings)
5480 .await?;
5481 if (security_result.is_allowed()
5486 || matches!(
5487 &security_result,
5488 SecurityCheckResult::RequireConfirmation { .. }
5489 ))
5490 && let Some((output, reason)) = self.host_tool_unavailability(&canonical_id)
5491 {
5492 let record = self.record_from_parts(
5493 &request,
5494 canonical_id,
5495 executed_arguments,
5496 started_at,
5497 start,
5498 false,
5499 false,
5500 output.to_string(),
5501 metadata,
5502 ToolPolicyDecisionRecord::unavailable(reason),
5503 Some(ToolApprovalRecord {
5504 status: ToolApprovalStatus::Unavailable,
5505 reason: Some(reason.to_string()),
5506 modified_arguments: None,
5507 }),
5508 false,
5509 false,
5510 );
5511 self.finish_tool_record(&record).await;
5512 return Ok(record);
5513 }
5514 match &security_result {
5515 SecurityCheckResult::Allow => {}
5516 SecurityCheckResult::Warn { message } => {
5517 warn!(tool = %canonical_id, message = %message, "Tool security warning");
5518 }
5519 SecurityCheckResult::Block { reason } => {
5520 let record = self.record_from_parts(
5521 &request,
5522 canonical_id,
5523 executed_arguments,
5524 started_at,
5525 start,
5526 false,
5527 false,
5528 format!("Denied: {}", reason),
5529 metadata,
5530 ToolPolicyDecisionRecord::deny(reason.clone()),
5531 approval_record,
5532 false,
5533 false,
5534 );
5535 self.finish_tool_record(&record).await;
5536 return Ok(record);
5537 }
5538 SecurityCheckResult::Unavailable { reason } => {
5539 let record = self.record_from_parts(
5540 &request,
5541 canonical_id,
5542 executed_arguments,
5543 started_at,
5544 start,
5545 false,
5546 false,
5547 format!("Unavailable: {}", reason),
5548 metadata,
5549 ToolPolicyDecisionRecord::unavailable(reason.clone()),
5550 approval_record,
5551 false,
5552 false,
5553 );
5554 self.finish_tool_record(&record).await;
5555 return Ok(record);
5556 }
5557 SecurityCheckResult::RequireConfirmation { message } => {
5558 if self.hitl_engine.is_none() {
5559 approval_record = Some(ToolApprovalRecord {
5560 status: ToolApprovalStatus::Unavailable,
5561 reason: Some("No HITL engine configured".to_string()),
5562 modified_arguments: None,
5563 });
5564 let record = self.record_from_parts(
5565 &request,
5566 canonical_id,
5567 executed_arguments,
5568 started_at,
5569 start,
5570 false,
5571 false,
5572 format!("Approval unavailable: {}", message),
5573 metadata,
5574 ToolPolicyDecisionRecord::approval(message.clone()),
5575 approval_record,
5576 false,
5577 false,
5578 );
5579 self.finish_tool_record(&record).await;
5580 return Ok(record);
5581 }
5582
5583 let check_result = HITLCheckResult::required(
5584 ApprovalTrigger::tool(&canonical_id, executed_arguments.clone()),
5585 HashMap::new(),
5586 message.clone(),
5587 None,
5588 );
5589 match self.request_hitl_approval(check_result).await? {
5590 ApprovalResult::Approved => {
5591 merge_approved_record(&mut approval_record);
5592 }
5593 ApprovalResult::Modified { changes } => {
5594 if let Some(obj) = executed_arguments.as_object_mut() {
5595 for (key, value) in changes {
5596 obj.insert(key, value);
5597 }
5598 }
5599 security_result = security_engine
5600 .validate_tool_execution_with_bindings(
5601 &canonical_id,
5602 &executed_arguments,
5603 &bindings,
5604 )
5605 .await?;
5606 if !matches!(
5607 security_result,
5608 SecurityCheckResult::Allow
5609 | SecurityCheckResult::Warn { .. }
5610 | SecurityCheckResult::RequireConfirmation { .. }
5611 ) {
5612 let reason = security_result
5613 .reason()
5614 .unwrap_or("modified arguments failed policy")
5615 .to_string();
5616 let record = self.record_from_parts(
5617 &request,
5618 canonical_id,
5619 executed_arguments.clone(),
5620 started_at,
5621 start,
5622 false,
5623 false,
5624 reason.clone(),
5625 metadata,
5626 ToolPolicyDecisionRecord::deny(reason),
5627 Some(ToolApprovalRecord {
5628 status: ToolApprovalStatus::Modified,
5629 reason: None,
5630 modified_arguments: Some(executed_arguments),
5631 }),
5632 false,
5633 false,
5634 );
5635 self.finish_tool_record(&record).await;
5636 return Ok(record);
5637 }
5638 approval_record = Some(ToolApprovalRecord {
5639 status: ToolApprovalStatus::Modified,
5640 reason: None,
5641 modified_arguments: Some(executed_arguments.clone()),
5642 });
5643 }
5644 ApprovalResult::Rejected { reason } => {
5645 let reason = reason.unwrap_or_else(|| "rejected".to_string());
5646 approval_record = Some(ToolApprovalRecord {
5647 status: ToolApprovalStatus::Rejected,
5648 reason: Some(reason.clone()),
5649 modified_arguments: None,
5650 });
5651 let record = self.record_from_parts(
5652 &request,
5653 canonical_id,
5654 executed_arguments,
5655 started_at,
5656 start,
5657 false,
5658 false,
5659 format!("Approval rejected: {}", reason),
5660 metadata,
5661 ToolPolicyDecisionRecord::approval(reason),
5662 approval_record,
5663 false,
5664 false,
5665 );
5666 self.finish_tool_record(&record).await;
5667 return Ok(record);
5668 }
5669 ApprovalResult::Timeout => {
5670 approval_record = Some(ToolApprovalRecord {
5671 status: ToolApprovalStatus::Timeout,
5672 reason: Some("approval timeout".to_string()),
5673 modified_arguments: None,
5674 });
5675 let record = self.record_from_parts(
5676 &request,
5677 canonical_id,
5678 executed_arguments,
5679 started_at,
5680 start,
5681 false,
5682 false,
5683 "Approval timed out".to_string(),
5684 metadata,
5685 ToolPolicyDecisionRecord::approval("approval timeout"),
5686 approval_record,
5687 false,
5688 false,
5689 );
5690 self.finish_tool_record(&record).await;
5691 return Ok(record);
5692 }
5693 }
5694 }
5695 }
5696
5697 if approval_record
5698 .as_ref()
5699 .is_some_and(|record| matches!(record.status, ToolApprovalStatus::NotRequired))
5700 && let Some(message) =
5701 security_engine.classification_approval_message(&canonical_id, &classification)
5702 {
5703 if self.hitl_engine.is_none() {
5704 approval_record = Some(ToolApprovalRecord {
5705 status: ToolApprovalStatus::Unavailable,
5706 reason: Some("No HITL engine configured".to_string()),
5707 modified_arguments: None,
5708 });
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 format!("Approval unavailable: {}", message),
5718 metadata,
5719 ToolPolicyDecisionRecord::approval(message),
5720 approval_record,
5721 false,
5722 false,
5723 );
5724 self.finish_tool_record(&record).await;
5725 return Ok(record);
5726 }
5727 let check_result = HITLCheckResult::required(
5728 ApprovalTrigger::tool(&canonical_id, executed_arguments.clone()),
5729 HashMap::new(),
5730 message.clone(),
5731 None,
5732 );
5733 match self.request_hitl_approval(check_result).await? {
5734 ApprovalResult::Approved => {
5735 merge_approved_record(&mut approval_record);
5736 }
5737 ApprovalResult::Modified { changes } => {
5738 if let Some(obj) = executed_arguments.as_object_mut() {
5739 for (key, value) in changes {
5740 obj.insert(key, value);
5741 }
5742 }
5743 let modified_security = security_engine
5744 .validate_tool_execution_with_bindings(
5745 &canonical_id,
5746 &executed_arguments,
5747 &bindings,
5748 )
5749 .await?;
5750 if !matches!(
5751 modified_security,
5752 SecurityCheckResult::Allow | SecurityCheckResult::Warn { .. }
5753 ) {
5754 let reason = modified_security
5755 .reason()
5756 .unwrap_or("modified arguments failed policy")
5757 .to_string();
5758 let record = self.record_from_parts(
5759 &request,
5760 canonical_id,
5761 executed_arguments.clone(),
5762 started_at,
5763 start,
5764 false,
5765 false,
5766 reason.clone(),
5767 metadata,
5768 ToolPolicyDecisionRecord::deny(reason),
5769 Some(ToolApprovalRecord {
5770 status: ToolApprovalStatus::Modified,
5771 reason: None,
5772 modified_arguments: Some(executed_arguments),
5773 }),
5774 false,
5775 false,
5776 );
5777 self.finish_tool_record(&record).await;
5778 return Ok(record);
5779 }
5780 approval_record = Some(ToolApprovalRecord {
5781 status: ToolApprovalStatus::Modified,
5782 reason: None,
5783 modified_arguments: Some(executed_arguments.clone()),
5784 });
5785 }
5786 ApprovalResult::Rejected { reason } => {
5787 let reason = reason.unwrap_or_else(|| "rejected".to_string());
5788 let record = self.record_from_parts(
5789 &request,
5790 canonical_id,
5791 executed_arguments,
5792 started_at,
5793 start,
5794 false,
5795 false,
5796 format!("Approval rejected: {}", reason),
5797 metadata,
5798 ToolPolicyDecisionRecord::approval(reason.clone()),
5799 Some(ToolApprovalRecord {
5800 status: ToolApprovalStatus::Rejected,
5801 reason: Some(reason),
5802 modified_arguments: None,
5803 }),
5804 false,
5805 false,
5806 );
5807 self.finish_tool_record(&record).await;
5808 return Ok(record);
5809 }
5810 ApprovalResult::Timeout => {
5811 let record = self.record_from_parts(
5812 &request,
5813 canonical_id,
5814 executed_arguments,
5815 started_at,
5816 start,
5817 false,
5818 false,
5819 "Approval timed out".to_string(),
5820 metadata,
5821 ToolPolicyDecisionRecord::approval("approval timeout"),
5822 Some(ToolApprovalRecord {
5823 status: ToolApprovalStatus::Timeout,
5824 reason: Some("approval timeout".to_string()),
5825 modified_arguments: None,
5826 }),
5827 false,
5828 false,
5829 );
5830 self.finish_tool_record(&record).await;
5831 return Ok(record);
5832 }
5833 }
5834 }
5835
5836 let hitl_lang_ctx = self.build_hitl_language_context();
5837 if let Some(ref hitl_engine) = self.hitl_engine {
5838 let check_result = self
5839 .observe_purpose(
5840 ObservationPurpose::HitlLocalization,
5841 hitl_engine.check_tool_with_localization(
5842 &canonical_id,
5843 &executed_arguments,
5844 &hitl_lang_ctx,
5845 self.approval_handler.as_ref(),
5846 Some(&self.llm_registry),
5847 ),
5848 )
5849 .await?;
5850 if check_result.is_required() {
5851 match self.request_hitl_approval(check_result).await? {
5852 ApprovalResult::Approved => {
5853 merge_approved_record(&mut approval_record);
5854 }
5855 ApprovalResult::Modified { changes } => {
5856 if let Some(obj) = executed_arguments.as_object_mut() {
5857 for (key, value) in changes {
5858 obj.insert(key, value);
5859 }
5860 }
5861 let modified_security = security_engine
5862 .validate_tool_execution_with_bindings(
5863 &canonical_id,
5864 &executed_arguments,
5865 &bindings,
5866 )
5867 .await?;
5868 if !matches!(
5869 modified_security,
5870 SecurityCheckResult::Allow | SecurityCheckResult::Warn { .. }
5871 ) {
5872 let reason = modified_security
5873 .reason()
5874 .unwrap_or("modified arguments failed policy")
5875 .to_string();
5876 let record = self.record_from_parts(
5877 &request,
5878 canonical_id,
5879 executed_arguments.clone(),
5880 started_at,
5881 start,
5882 false,
5883 false,
5884 reason.clone(),
5885 metadata,
5886 ToolPolicyDecisionRecord::deny(reason),
5887 Some(ToolApprovalRecord {
5888 status: ToolApprovalStatus::Modified,
5889 reason: None,
5890 modified_arguments: Some(executed_arguments),
5891 }),
5892 false,
5893 false,
5894 );
5895 self.finish_tool_record(&record).await;
5896 return Ok(record);
5897 }
5898 approval_record = Some(ToolApprovalRecord {
5899 status: ToolApprovalStatus::Modified,
5900 reason: None,
5901 modified_arguments: Some(executed_arguments.clone()),
5902 });
5903 }
5904 ApprovalResult::Rejected { reason } => {
5905 let reason = reason.unwrap_or_else(|| "rejected".to_string());
5906 let record = self.record_from_parts(
5907 &request,
5908 canonical_id,
5909 executed_arguments,
5910 started_at,
5911 start,
5912 false,
5913 false,
5914 format!("Approval rejected: {}", reason),
5915 metadata,
5916 ToolPolicyDecisionRecord::approval(reason.clone()),
5917 Some(ToolApprovalRecord {
5918 status: ToolApprovalStatus::Rejected,
5919 reason: Some(reason),
5920 modified_arguments: None,
5921 }),
5922 false,
5923 false,
5924 );
5925 self.finish_tool_record(&record).await;
5926 return Ok(record);
5927 }
5928 ApprovalResult::Timeout => {
5929 let record = self.record_from_parts(
5930 &request,
5931 canonical_id,
5932 executed_arguments,
5933 started_at,
5934 start,
5935 false,
5936 false,
5937 "Approval timed out".to_string(),
5938 metadata,
5939 ToolPolicyDecisionRecord::approval("approval timeout"),
5940 Some(ToolApprovalRecord {
5941 status: ToolApprovalStatus::Timeout,
5942 reason: Some("approval timeout".to_string()),
5943 modified_arguments: None,
5944 }),
5945 false,
5946 false,
5947 );
5948 self.finish_tool_record(&record).await;
5949 return Ok(record);
5950 }
5951 }
5952 }
5953
5954 let condition_check = self
5955 .observe_purpose(
5956 ObservationPurpose::HitlLocalization,
5957 hitl_engine.check_conditions_with_localization(
5958 &executed_arguments,
5959 &hitl_lang_ctx,
5960 self.approval_handler.as_ref(),
5961 Some(&self.llm_registry),
5962 ),
5963 )
5964 .await?;
5965 if condition_check.is_required() {
5966 match self.request_hitl_approval(condition_check).await? {
5967 ApprovalResult::Approved => {
5968 merge_approved_record(&mut approval_record);
5969 }
5970 ApprovalResult::Modified { changes } => {
5971 if let Some(obj) = executed_arguments.as_object_mut() {
5972 for (key, value) in changes {
5973 obj.insert(key, value);
5974 }
5975 }
5976 let modified_security = security_engine
5977 .validate_tool_execution_with_bindings(
5978 &canonical_id,
5979 &executed_arguments,
5980 &bindings,
5981 )
5982 .await?;
5983 if !matches!(
5984 modified_security,
5985 SecurityCheckResult::Allow | SecurityCheckResult::Warn { .. }
5986 ) {
5987 let reason = modified_security
5988 .reason()
5989 .unwrap_or("modified arguments failed policy")
5990 .to_string();
5991 let record = self.record_from_parts(
5992 &request,
5993 canonical_id,
5994 executed_arguments,
5995 started_at,
5996 start,
5997 false,
5998 false,
5999 reason.clone(),
6000 metadata,
6001 ToolPolicyDecisionRecord::deny(reason),
6002 approval_record,
6003 false,
6004 false,
6005 );
6006 self.finish_tool_record(&record).await;
6007 return Ok(record);
6008 }
6009 approval_record = Some(ToolApprovalRecord {
6010 status: ToolApprovalStatus::Modified,
6011 reason: None,
6012 modified_arguments: Some(executed_arguments.clone()),
6013 });
6014 }
6015 ApprovalResult::Rejected { reason } => {
6016 let reason = reason.unwrap_or_else(|| "rejected".to_string());
6017 let record = self.record_from_parts(
6018 &request,
6019 canonical_id,
6020 executed_arguments,
6021 started_at,
6022 start,
6023 false,
6024 false,
6025 format!("Approval rejected: {}", reason),
6026 metadata,
6027 ToolPolicyDecisionRecord::approval(reason.clone()),
6028 Some(ToolApprovalRecord {
6029 status: ToolApprovalStatus::Rejected,
6030 reason: Some(reason),
6031 modified_arguments: None,
6032 }),
6033 false,
6034 false,
6035 );
6036 self.finish_tool_record(&record).await;
6037 return Ok(record);
6038 }
6039 ApprovalResult::Timeout => {
6040 let record = self.record_from_parts(
6041 &request,
6042 canonical_id,
6043 executed_arguments,
6044 started_at,
6045 start,
6046 false,
6047 false,
6048 "Approval timed out".to_string(),
6049 metadata,
6050 ToolPolicyDecisionRecord::approval("approval timeout"),
6051 Some(ToolApprovalRecord {
6052 status: ToolApprovalStatus::Timeout,
6053 reason: Some("approval timeout".to_string()),
6054 modified_arguments: None,
6055 }),
6056 false,
6057 false,
6058 );
6059 self.finish_tool_record(&record).await;
6060 return Ok(record);
6061 }
6062 }
6063 }
6064 }
6065
6066 executed_arguments = security_engine.prepare_tool_arguments_with_bindings(
6071 &canonical_id,
6072 &executed_arguments,
6073 &bindings,
6074 );
6075 if let Some(record) = approval_record.as_mut()
6076 && matches!(record.status, ToolApprovalStatus::Modified)
6077 {
6078 record.modified_arguments = Some(executed_arguments.clone());
6079 }
6080 let binding_security_result = security_engine
6081 .validate_tool_execution_with_bindings(&canonical_id, &executed_arguments, &bindings)
6082 .await?;
6083 let approval_confirmation_required = matches!(
6084 binding_security_result,
6085 SecurityCheckResult::RequireConfirmation { .. }
6086 ) || security_engine
6087 .classification_approval_message(
6088 &canonical_id,
6089 &resolved.tool.classify_call(&executed_arguments),
6090 )
6091 .is_some();
6092 let approval_binding = approval_record.as_ref().and_then(|record| {
6093 matches!(
6094 record.status,
6095 ToolApprovalStatus::Approved | ToolApprovalStatus::Modified
6096 )
6097 .then(|| ToolApprovalBinding {
6098 canonical_id: canonical_id.clone(),
6099 arguments: executed_arguments.clone(),
6100 confirmation_required: approval_confirmation_required,
6101 policy_version: security_engine.policy_version(),
6102 runtime_control_version: approval_control_snapshot.version,
6103 state_generation: initial_scope_snapshot.state_generation,
6104 reviewed_tool: Arc::clone(&resolved.tool),
6105 })
6106 });
6107
6108 let control_snapshot = self.runtime_safety_snapshot();
6113 let resolved = self.tools.resolve(&request.requested_name);
6114 let registry_version = self.tools.version();
6115 let mut versions = ToolDecisionVersions {
6116 policy: control_snapshot.tool_security.policy_version(),
6117 registry: registry_version,
6118 runtime_control: control_snapshot.version,
6119 state: None,
6120 };
6121 metadata.insert(
6122 "runtime_scope_snapshot".to_string(),
6123 serde_json::to_value(&control_snapshot.tool_scope_override).unwrap_or(Value::Null),
6124 );
6125 let resolved = match resolved {
6126 Some(resolved) => resolved,
6127 None => {
6128 let reason = format!(
6129 "Tool '{}' became unavailable after approval",
6130 request.requested_name
6131 );
6132 let record = self.record_from_parts_at(
6133 &request,
6134 request.requested_name.clone(),
6135 executed_arguments,
6136 started_at,
6137 start,
6138 false,
6139 false,
6140 reason.clone(),
6141 metadata,
6142 ToolPolicyDecisionRecord::unavailable(reason),
6143 approval_record,
6144 false,
6145 false,
6146 versions,
6147 );
6148 self.finish_tool_record(&record).await;
6149 return Ok(record);
6150 }
6151 };
6152
6153 let canonical_id = resolved.identity.canonical_id.clone();
6154 if let Some(reason) =
6155 fallback_state.final_rejection_reason(&admitted_canonical_id, &canonical_id)
6156 {
6157 metadata.insert(
6161 "fallback_chain".to_string(),
6162 serde_json::to_value(&fallback_state.visited_canonical_ids).unwrap_or(Value::Null),
6163 );
6164 metadata.insert(
6165 "final_resolved_canonical_id".to_string(),
6166 Value::String(canonical_id),
6167 );
6168 let record = self.record_from_parts_at(
6169 &request,
6170 admitted_canonical_id,
6171 executed_arguments,
6172 started_at,
6173 start,
6174 false,
6175 false,
6176 format!("Denied: {reason}"),
6177 metadata,
6178 ToolPolicyDecisionRecord::deny(reason),
6179 approval_record,
6180 false,
6181 false,
6182 versions,
6183 );
6184 self.finish_tool_record(&record).await;
6185 return Ok(record);
6186 }
6187 let bindings = resolved.tool.policy_bindings();
6188 let final_arguments = control_snapshot
6189 .tool_security
6190 .prepare_tool_arguments_with_bindings(&canonical_id, &executed_arguments, &bindings);
6191 if let Some(record) = approval_record.as_mut()
6192 && matches!(record.status, ToolApprovalStatus::Modified)
6193 {
6194 record.modified_arguments = Some(final_arguments.clone());
6195 }
6196 let classification = resolved.tool.classify_call(&final_arguments);
6197 let safety = resolved.tool.safety_metadata();
6198 let security_engine = control_snapshot.tool_security;
6199 let tool_config = self.recovery_manager.get_tool_config(&canonical_id).clone();
6200 let recovery_timeout_ms = self.recovery_manager.get_tool_timeout(&canonical_id);
6201 metadata.insert(
6202 "classification".to_string(),
6203 serde_json::to_value(&classification).unwrap_or(Value::Null),
6204 );
6205 let (limits, timeout) = match Self::effective_tool_limits(
6209 &security_engine,
6210 &canonical_id,
6211 &safety,
6212 &classification,
6213 recovery_timeout_ms,
6214 ) {
6215 Ok(effective) => effective,
6216 Err(error) => {
6217 let reason = error.to_string();
6218 metadata.insert(
6219 "configuration_error".to_string(),
6220 Value::String(reason.clone()),
6221 );
6222 let record = self.record_from_parts_at(
6223 &request,
6224 canonical_id,
6225 final_arguments,
6226 started_at,
6227 start,
6228 false,
6229 false,
6230 format!("Denied: {reason}"),
6231 metadata,
6232 ToolPolicyDecisionRecord::deny(reason),
6233 approval_record,
6234 false,
6235 false,
6236 versions,
6237 );
6238 self.finish_tool_record(&record).await;
6239 return Ok(record);
6240 }
6241 };
6242 let policy_snapshot = security_engine.policy_snapshot(&canonical_id);
6243 let resource_lock_keys =
6244 tool_resource_lock_keys(&canonical_id, &final_arguments, &bindings, &classification);
6245 metadata.insert(
6246 "effective_limits".to_string(),
6247 serde_json::to_value(&limits).unwrap_or(Value::Null),
6248 );
6249 metadata.insert(
6250 "resource_lock_keys".to_string(),
6251 serde_json::to_value(&resource_lock_keys).unwrap_or(Value::Null),
6252 );
6253 if policy_snapshot.is_null() {
6254 metadata.remove("policy_snapshot");
6255 } else {
6256 metadata.insert("policy_snapshot".to_string(), policy_snapshot.clone());
6257 }
6258
6259 let final_denial = |canonical_id: String,
6260 output: String,
6261 policy: ToolPolicyDecisionRecord,
6262 metadata: HashMap<String, Value>,
6263 decision_versions: ToolDecisionVersions| {
6264 self.record_from_parts_at(
6265 &request,
6266 canonical_id,
6267 final_arguments.clone(),
6268 started_at,
6269 start,
6270 false,
6271 false,
6272 output,
6273 metadata,
6274 policy,
6275 approval_record.clone(),
6276 false,
6277 false,
6278 decision_versions,
6279 )
6280 };
6281
6282 if control_snapshot.emergency_deny {
6283 let reason = "Tool execution is disabled by runtime control".to_string();
6284 let record = final_denial(
6285 canonical_id,
6286 reason.clone(),
6287 ToolPolicyDecisionRecord::deny(reason),
6288 metadata,
6289 versions,
6290 );
6291 self.finish_tool_record(&record).await;
6292 return Ok(record);
6293 }
6294
6295 let available_snapshot = self
6300 .get_available_tool_ids_snapshot_for_scope(
6301 control_snapshot.tool_scope_override.as_deref(),
6302 )
6303 .await?;
6304 versions.state = available_snapshot.state_generation;
6305 metadata.insert(
6306 "available_tool_ids_snapshot".to_string(),
6307 serde_json::to_value(&available_snapshot.tool_ids).unwrap_or(Value::Null),
6308 );
6309 metadata.insert(
6310 "state_generation_snapshot".to_string(),
6311 serde_json::to_value(available_snapshot.state_generation).unwrap_or(Value::Null),
6312 );
6313 if !available_snapshot
6314 .tool_ids
6315 .iter()
6316 .any(|tool_id| tool_id == &canonical_id)
6317 {
6318 let reason = format!(
6319 "Tool '{}' is not available in the final runtime scope",
6320 canonical_id
6321 );
6322 let record = final_denial(
6323 canonical_id,
6324 reason.clone(),
6325 ToolPolicyDecisionRecord::deny(reason),
6326 metadata,
6327 versions,
6328 );
6329 self.finish_tool_record(&record).await;
6330 return Ok(record);
6331 }
6332
6333 let final_security_result = security_engine
6338 .validate_tool_execution_with_bindings(&canonical_id, &final_arguments, &bindings)
6339 .await?;
6340 match &final_security_result {
6341 SecurityCheckResult::Block { reason } => {
6342 let record = final_denial(
6343 canonical_id,
6344 format!("Denied: {}", reason),
6345 ToolPolicyDecisionRecord::deny(reason.clone()),
6346 metadata,
6347 versions,
6348 );
6349 self.finish_tool_record(&record).await;
6350 return Ok(record);
6351 }
6352 SecurityCheckResult::Unavailable { reason } => {
6353 let record = final_denial(
6354 canonical_id,
6355 format!("Unavailable: {}", reason),
6356 ToolPolicyDecisionRecord::unavailable(reason.clone()),
6357 metadata,
6358 versions,
6359 );
6360 self.finish_tool_record(&record).await;
6361 return Ok(record);
6362 }
6363 SecurityCheckResult::Warn { message } => {
6364 warn!(tool = %canonical_id, message = %message, "Tool security warning after approval");
6365 }
6366 SecurityCheckResult::Allow | SecurityCheckResult::RequireConfirmation { .. } => {}
6367 }
6368 let final_confirmation_required = matches!(
6369 final_security_result,
6370 SecurityCheckResult::RequireConfirmation { .. }
6371 ) || security_engine
6372 .classification_approval_message(&canonical_id, &classification)
6373 .is_some();
6374 let stale_approval = approval_binding.as_ref().is_some_and(|binding| {
6375 binding.is_stale(
6376 &canonical_id,
6377 &final_arguments,
6378 final_confirmation_required,
6379 versions,
6380 &resolved.tool,
6381 )
6382 });
6383 if stale_approval {
6384 let reason = "Approval became stale before final admission".to_string();
6385 let record = final_denial(
6386 canonical_id,
6387 reason.clone(),
6388 ToolPolicyDecisionRecord::deny(reason),
6389 metadata,
6390 versions,
6391 );
6392 self.finish_tool_record(&record).await;
6393 return Ok(record);
6394 }
6395 if final_confirmation_required && approval_binding.is_none() {
6396 let reason = "Final policy requires fresh approval".to_string();
6397 let record = final_denial(
6398 canonical_id,
6399 reason.clone(),
6400 ToolPolicyDecisionRecord::approval(reason),
6401 metadata,
6402 versions,
6403 );
6404 self.finish_tool_record(&record).await;
6405 return Ok(record);
6406 }
6407
6408 if let Some((_, reason)) = self.host_tool_unavailability(&canonical_id) {
6409 let record = final_denial(
6410 canonical_id,
6411 reason.to_string(),
6412 ToolPolicyDecisionRecord::unavailable(reason),
6413 metadata,
6414 versions,
6415 );
6416 self.finish_tool_record(&record).await;
6417 return Ok(record);
6418 }
6419
6420 let Some(resource_guards) = self.acquire_tool_resource_locks(&resource_lock_keys).await
6425 else {
6426 let reason = "Tool execution cancelled while waiting for resource locks".to_string();
6430 let mut record = final_denial(
6431 canonical_id,
6432 reason.clone(),
6433 ToolPolicyDecisionRecord::deny(reason),
6434 metadata,
6435 versions,
6436 );
6437 record.cancelled = true;
6438 record.cancellation_reason = Some("runtime control cancellation".to_string());
6439 self.finish_tool_record(&record).await;
6440 return Ok(record);
6441 };
6442
6443 let admission = self.admit_tool_execution(
6448 versions.runtime_control,
6449 versions.policy,
6450 versions.state,
6451 &canonical_id,
6452 );
6453 if !matches!(admission, SecurityCheckResult::Allow) {
6454 let latest_control = self.runtime_safety_snapshot();
6455 let reason = admission
6456 .reason()
6457 .unwrap_or("tool admission was denied")
6458 .to_string();
6459 let policy = if admission.is_unavailable() {
6460 ToolPolicyDecisionRecord::unavailable(reason.clone())
6461 } else {
6462 ToolPolicyDecisionRecord::deny(reason.clone())
6463 };
6464 let record = self.record_from_parts_at(
6465 &request,
6466 canonical_id,
6467 final_arguments,
6468 started_at,
6469 start,
6470 false,
6471 false,
6472 reason,
6473 metadata,
6474 policy,
6475 approval_record,
6476 false,
6477 false,
6478 ToolDecisionVersions {
6479 policy: latest_control.tool_security.policy_version(),
6480 registry: versions.registry,
6481 runtime_control: latest_control.version,
6482 state: self
6483 .state_machine
6484 .as_ref()
6485 .map(|state_machine| state_machine.generation()),
6486 },
6487 );
6488 self.finish_tool_record_after_resource_guards(resource_guards, &record)
6489 .await;
6490 return Ok(record);
6491 }
6492 let executed_arguments = final_arguments;
6493
6494 let turn_actor = current_turn_actor_context();
6495 let actor = ToolActorContext {
6496 actor_id: turn_actor
6497 .as_ref()
6498 .and_then(|context| context.effective_actor_id().map(str::to_string))
6499 .or_else(|| self.actor_id()),
6500 origin_actor_id: turn_actor
6501 .as_ref()
6502 .and_then(|context| context.origin_actor_id.clone()),
6503 sender_agent_id: turn_actor
6504 .as_ref()
6505 .and_then(|context| context.sender_agent_id.clone()),
6506 };
6507 let tool_context = ToolExecutionContext {
6508 requested_name: request.requested_name.clone(),
6509 canonical_id: canonical_id.clone(),
6510 display_name: resolved.identity.display_name.clone(),
6511 provider_id: resolved.identity.provider_id.clone(),
6512 registry_version: versions.registry,
6513 policy_version: versions.policy,
6514 runtime_control_version: versions.runtime_control,
6515 call_id: request.call_id.clone(),
6516 source: request.source.clone(),
6517 actor,
6518 cancellation: ToolCancellationToken::new(
6519 Arc::clone(&self.runtime_control.emergency_deny),
6520 Some("runtime control cancellation".to_string()),
6521 ),
6522 started_at,
6523 deadline: None,
6524 permission: ToolPolicyDecisionRecord::allow(),
6525 approval: approval_record.clone(),
6526 classification: classification.clone(),
6527 safety,
6528 limits: limits.clone(),
6529 policy_snapshot,
6530 custom_config: security_engine.custom_config(&canonical_id),
6531 };
6532 let (mut result, timed_out, cancelled, invoked) = self
6533 .run_tool_with_retries(
6534 &canonical_id,
6535 resolved.tool.clone(),
6536 executed_arguments.clone(),
6537 tool_context,
6538 timeout,
6539 tool_config.max_retries,
6540 )
6541 .await?;
6542
6543 let fallback_tool = if !result.success && !cancelled {
6547 match &tool_config.on_failure {
6548 ToolFailureAction::Skip => {
6549 result = ToolResult::ok(format!(
6550 "{{\"skipped\": true, \"reason\": \"Tool '{}' was skipped after failure\"}}",
6551 canonical_id
6552 ));
6553 None
6554 }
6555 ToolFailureAction::Fallback { fallback_tool } => Some(fallback_tool.clone()),
6556 ToolFailureAction::ReportError => None,
6557 }
6558 } else {
6559 None
6560 };
6561
6562 let output_cap = limits.max_output_chars;
6563 let (output, output_truncated) =
6564 Self::truncate_tool_output(result.output.clone(), output_cap);
6565 if let Some(result_metadata) = result.metadata {
6566 metadata.extend(result_metadata);
6567 }
6568 let mut record = self.record_from_parts_at(
6569 &request,
6570 canonical_id,
6571 executed_arguments,
6572 started_at,
6573 start,
6574 invoked,
6575 result.success,
6576 output,
6577 metadata,
6578 ToolPolicyDecisionRecord::allow(),
6579 approval_record,
6580 timed_out,
6581 output_truncated,
6582 versions,
6583 );
6584 record.cancelled = cancelled;
6585 if cancelled {
6586 record.cancellation_reason = Some("runtime control cancellation".to_string());
6587 }
6588 if let Some(fallback_tool) = fallback_tool {
6589 let fallback_arguments = record.executed_arguments.clone();
6590 let original_tool = record.canonical_id.clone();
6591 self.finish_tool_record_after_resource_guards(resource_guards, &record)
6595 .await;
6596 let fallback_request = ToolExecutionRequest::new(
6597 request.call_id.clone(),
6598 fallback_tool,
6599 fallback_arguments,
6600 ToolCallSource::Fallback { original_tool },
6601 );
6602 return Box::pin(self.execute_tool_record_inner(fallback_request, fallback_state))
6603 .await;
6604 }
6605 self.finish_tool_record_after_resource_guards(resource_guards, &record)
6606 .await;
6607 Ok(record)
6608 }
6609
6610 #[instrument(skip(self, tool_call), fields(tool = %tool_call.name))]
6611 async fn execute_tool_smart(&self, tool_call: &ToolCall) -> Result<String> {
6612 let record = self
6613 .execute_tool_record(ToolExecutionRequest::new(
6614 tool_call.id.clone(),
6615 tool_call.name.clone(),
6616 tool_call.arguments.clone(),
6617 ToolCallSource::Model,
6618 ))
6619 .await?;
6620 if record.success {
6621 Ok(record.model_output_string())
6622 } else if matches!(record.policy.outcome, PermissionOutcome::RequiresApproval) {
6623 Err(AgentError::HITLRejected(record.model_output_string()))
6624 } else {
6625 Err(AgentError::Tool(record.model_output_string()))
6626 }
6627 }
6628
6629 async fn select_skill_candidate(&self, input: &str) -> Result<Option<SkillCandidate>> {
6635 let Some(ref router) = self.skill_router else {
6636 return Ok(None);
6637 };
6638 let available_skills = self.get_available_skills();
6639 if available_skills.is_empty() {
6640 return Ok(None);
6641 }
6642 let skill_ids: Vec<&str> = available_skills.iter().map(|s| s.id.as_str()).collect();
6643 let Some(skill_id) = self
6644 .observe_purpose(
6645 ObservationPurpose::SkillRouting,
6646 router.select_skill_filtered(input, &skill_ids),
6647 )
6648 .await?
6649 else {
6650 return Ok(None);
6651 };
6652 let skill = router
6653 .get_skill(&skill_id)
6654 .cloned()
6655 .ok_or_else(|| AgentError::Skill(format!("Skill not found: {}", skill_id)))?;
6656 info!(skill_id = %skill_id, "Skill selected");
6657 Ok(Some(SkillCandidate::new(skill_id, skill)))
6658 }
6659
6660 async fn commit_skill_candidate_route_result(
6665 &self,
6666 candidate: SkillCandidate,
6667 input: &str,
6668 ) -> Result<SkillRouteResult> {
6669 let skill_id = candidate.skill_id;
6670 let skill = candidate.skill;
6671 let expected_state_generation = self
6672 .state_machine
6673 .as_ref()
6674 .map(|state_machine| state_machine.generation());
6675 let expected_disambiguation_epoch = self.disambiguation_epoch.load(Ordering::SeqCst);
6676 if let Some(ref skill_disambig) = skill.disambiguation
6677 && skill_disambig.enabled.unwrap_or(false)
6678 && let Some(ref disambiguator) = self.disambiguation_manager
6679 {
6680 let context = self.build_disambiguation_context().await?;
6681 let state_override = self
6682 .state_machine
6683 .as_ref()
6684 .and_then(|sm| sm.current_definition())
6685 .and_then(|def| def.disambiguation.clone());
6686
6687 let disambiguation_result = self
6688 .observe_purpose(
6689 ObservationPurpose::DisambiguationDetection,
6690 disambiguator.process_input_with_override(
6691 input,
6692 &context,
6693 state_override.as_ref(),
6694 Some(skill_disambig),
6695 ),
6696 )
6697 .await?;
6698 let current_state_generation = self
6699 .state_machine
6700 .as_ref()
6701 .map(|state_machine| state_machine.generation());
6702 if current_state_generation != expected_state_generation
6703 || self.disambiguation_epoch.load(Ordering::SeqCst) != expected_disambiguation_epoch
6704 {
6705 disambiguator.clear_pending().await;
6706 *self.pending_skill_id.write() = None;
6707 return Err(AgentError::Other(
6708 "State or reset ownership changed during skill disambiguation".to_string(),
6709 ));
6710 }
6711 match disambiguation_result {
6712 DisambiguationResult::Clear => {
6713 debug!(skill_id = %skill_id, "Skill disambiguation: clear");
6714 }
6715 DisambiguationResult::NeedsClarification {
6716 question,
6717 detection,
6718 } => {
6719 let admission = self
6720 .admit_disambiguation_redispatch(
6721 expected_disambiguation_epoch,
6722 expected_state_generation,
6723 )
6724 .await?;
6725 let awaiting_confirmation = disambiguator.has_pending_confirmation().await;
6726 info!(
6727 skill_id = %skill_id,
6728 ambiguity_type = ?detection.ambiguity_type,
6729 confidence = detection.confidence,
6730 "Skill requires clarification before execution"
6731 );
6732 *self.pending_skill_id.write() = Some(skill_id.clone());
6733 let response = AgentResponse::new(&question.question).with_metadata(
6734 "disambiguation",
6735 serde_json::json!({
6736 "status": if awaiting_confirmation { "awaiting_confirmation" } else { "awaiting_clarification" },
6737 "skill_id": skill_id,
6738 "options": question.options,
6739 "clarifying": question.clarifying,
6740 "detection": {
6741 "type": detection.ambiguity_type,
6742 "confidence": detection.confidence,
6743 "what_is_unclear": detection.what_is_unclear,
6744 }
6745 }),
6746 );
6747 drop(admission);
6748 return Ok(SkillRouteResult::NeedsClarification {
6749 response,
6750 ownership: Some(DisambiguationOwnership {
6751 epoch: expected_disambiguation_epoch,
6752 state_generation: expected_state_generation,
6753 }),
6754 });
6755 }
6756 DisambiguationResult::Clarified { enriched_input, .. } => {
6757 info!(skill_id = %skill_id, enriched = %enriched_input, "Skill disambiguation clarified");
6758 let admission = self
6759 .admit_disambiguation_redispatch(
6760 expected_disambiguation_epoch,
6761 expected_state_generation,
6762 )
6763 .await?;
6764 drop(admission);
6765 let content = self.execute_skill(&skill, &enriched_input).await?;
6766 return Ok(SkillRouteResult::Response { skill_id, content });
6767 }
6768 DisambiguationResult::ProceedWithBestGuess { enriched_input } => {
6769 info!(skill_id = %skill_id, "Skill disambiguation best guess");
6770 let admission = self
6771 .admit_disambiguation_redispatch(
6772 expected_disambiguation_epoch,
6773 expected_state_generation,
6774 )
6775 .await?;
6776 drop(admission);
6777 let content = self.execute_skill(&skill, &enriched_input).await?;
6778 return Ok(SkillRouteResult::Response { skill_id, content });
6779 }
6780 DisambiguationResult::GiveUp { reason } => {
6781 warn!(skill_id = %skill_id, reason = %reason, "Skill disambiguation gave up");
6782 let apology = self
6783 .generate_localized_apology(
6784 "Generate a brief, polite apology saying you couldn't understand the request. Be concise.",
6785 &reason,
6786 )
6787 .await
6788 .unwrap_or_else(|_| {
6789 format!("I'm sorry, I couldn't understand your request: {}", reason)
6790 });
6791 return Ok(SkillRouteResult::NeedsClarification {
6792 response: AgentResponse::new(&apology),
6793 ownership: None,
6794 });
6795 }
6796 DisambiguationResult::Escalate { reason } => {
6797 info!(skill_id = %skill_id, reason = %reason, "Skill disambiguation escalating");
6798 let apology = self
6799 .generate_localized_apology(
6800 "Explain briefly that you're transferring the user to a human agent for help.",
6801 &reason,
6802 )
6803 .await
6804 .unwrap_or_else(|_| {
6805 format!("I need human assistance to help with your request: {}", reason)
6806 });
6807 return Ok(SkillRouteResult::NeedsClarification {
6808 response: AgentResponse::new(&apology),
6809 ownership: None,
6810 });
6811 }
6812 DisambiguationResult::Abandoned { .. } => {
6813 debug!(skill_id = %skill_id, "Skill disambiguation abandoned");
6814 return Ok(SkillRouteResult::NoMatch);
6815 }
6816 }
6817 }
6818 let admission = self
6819 .admit_disambiguation_redispatch(
6820 expected_disambiguation_epoch,
6821 expected_state_generation,
6822 )
6823 .await?;
6824 drop(admission);
6825 let content = self.execute_skill(&skill, input).await?;
6826 Ok(SkillRouteResult::Response { skill_id, content })
6827 }
6828
6829 async fn try_skill_route(&self, input: &str) -> Result<SkillRouteResult> {
6831 if let Some(candidate) = self.select_skill_candidate(input).await? {
6832 self.commit_skill_candidate_route_result(candidate, input)
6833 .await
6834 } else {
6835 Ok(SkillRouteResult::NoMatch)
6836 }
6837 }
6838
6839 fn skill_clarification_needs_memory_record(response: &AgentResponse) -> bool {
6842 response
6843 .metadata
6844 .as_ref()
6845 .and_then(|m| m.get("disambiguation"))
6846 .and_then(|d| d.get("status"))
6847 .and_then(|s| s.as_str())
6848 == Some("awaiting_clarification")
6849 }
6850
6851 async fn commit_winning_skill_candidate(
6858 &self,
6859 candidate: SkillCandidate,
6860 processed_input: &str,
6861 input_context: &HashMap<String, Value>,
6862 ) -> Result<Option<AgentResponse>> {
6863 self.commit_root_user_message(processed_input).await?;
6864 match self
6865 .commit_skill_candidate_route_result(candidate, processed_input)
6866 .await?
6867 {
6868 SkillRouteResult::Response { skill_id, content } => self
6869 .handle_skill_response(processed_input, &skill_id, content, input_context)
6870 .await
6871 .map(Some),
6872 SkillRouteResult::NeedsClarification {
6873 response,
6874 ownership,
6875 } => {
6876 let admission = self
6877 .admit_optional_disambiguation_ownership(ownership)
6878 .await?;
6879 if Self::skill_clarification_needs_memory_record(&response) {
6880 self.memory
6881 .add_message(ChatMessage::assistant(&response.content))
6882 .await?;
6883 }
6884 drop(admission);
6885 self.finish_turn_if_root(&response).await?;
6886 Ok(Some(response))
6887 }
6888 SkillRouteResult::NoMatch => Ok(None),
6889 }
6890 }
6891
6892 async fn execute_skill(&self, skill: &SkillDefinition, input: &str) -> Result<String> {
6894 if let Some(ref executor) = self.skill_executor {
6895 let skill_reasoning = self.get_skill_reasoning_config(skill);
6896 let skill_reflection = self.get_skill_reflection_config(skill);
6897
6898 debug!(
6899 skill_id = %skill.id,
6900 reasoning_mode = ?skill_reasoning.mode,
6901 reflection_enabled = ?skill_reflection.enabled,
6902 "Skill reasoning/reflection config"
6903 );
6904
6905 let response = self
6906 .observe_purpose(
6907 ObservationPurpose::SkillPrompt,
6908 executor.execute_with_invoker(skill, input, serde_json::json!({}), self),
6909 )
6910 .await?;
6911
6912 if skill_reflection.requires_evaluation() && skill_reflection.is_enabled() {
6913 let should_reflect = self
6914 .should_reflect_with_config(input, &response, &skill_reflection)
6915 .await?;
6916 if should_reflect {
6917 let evaluated = self
6918 .evaluate_and_retry_with_config(input, response, &skill_reflection)
6919 .await?;
6920 return Ok(evaluated);
6921 }
6922 }
6923
6924 return Ok(response);
6925 }
6926 Err(AgentError::Skill(
6927 "No skill executor configured".to_string(),
6928 ))
6929 }
6930
6931 async fn execute_skill_by_id(&self, skill_id: &str, input: &str) -> Result<String> {
6934 let skill = self
6935 .skill_router
6936 .as_ref()
6937 .and_then(|r| r.get_skill(skill_id).cloned())
6938 .ok_or_else(|| AgentError::Skill(format!("Skill not found: {}", skill_id)))?;
6939 self.execute_skill(&skill, input).await
6940 }
6941
6942 async fn should_reflect_with_config(
6945 &self,
6946 input: &str,
6947 response: &str,
6948 config: &ReflectionConfig,
6949 ) -> Result<bool> {
6950 if !config.requires_evaluation() {
6951 return Ok(false);
6952 }
6953
6954 if config.is_enabled() {
6955 return Ok(true);
6956 }
6957
6958 let evaluator_llm = self.optional_role_llm(
6959 ai_agents_llm::LLMRole::ReasoningReflectionDecision,
6960 config.evaluator_llm.as_deref(),
6961 || {
6962 config
6963 .evaluator_llm
6964 .as_ref()
6965 .and_then(|alias| self.llm_registry.get(alias).ok())
6966 .or_else(|| self.llm_registry.router().ok())
6967 .or_else(|| self.llm_registry.default().ok())
6968 },
6969 )?;
6970
6971 let Some(llm) = evaluator_llm else {
6972 return Ok(false);
6973 };
6974
6975 let response_preview: String = response.chars().take(500).collect();
6976 let prompt = format!(
6977 r#"Should this response be evaluated for quality? Consider if it's a complex or important response.
6978
6979User query: "{}"
6980Response: "{}"
6981
6982Answer YES or NO only."#,
6983 input, response_preview
6984 );
6985
6986 let messages = vec![ChatMessage::user(&prompt)];
6987 let result = self
6988 .observe_purpose(
6989 ObservationPurpose::ReflectionDecision,
6990 llm.complete(&messages, None),
6991 )
6992 .await;
6993
6994 match result {
6995 Ok(resp) => Ok(resp.content.trim().to_uppercase().contains("YES")),
6996 Err(_) => Ok(false),
6997 }
6998 }
6999
7000 async fn evaluate_and_retry_with_config(
7001 &self,
7002 input: &str,
7003 mut response: String,
7004 config: &ReflectionConfig,
7005 ) -> Result<String> {
7006 let llm = self.get_state_llm()?;
7007 let mut attempts = 0u32;
7008 let max_retries = config.max_retries;
7009
7010 loop {
7011 let evaluation = self
7012 .evaluate_response_with_config(input, &response, config)
7013 .await?;
7014
7015 if evaluation.passed || attempts >= max_retries {
7016 info!(
7017 passed = evaluation.passed,
7018 confidence = evaluation.confidence,
7019 attempts = attempts + 1,
7020 "Skill reflection evaluation complete"
7021 );
7022 return Ok(response);
7023 }
7024
7025 debug!(
7026 attempt = attempts + 1,
7027 failed_criteria = evaluation.failed_criteria().count(),
7028 "Skill response did not meet criteria, retrying"
7029 );
7030
7031 let feedback: Vec<String> = evaluation
7032 .failed_criteria()
7033 .map(|c| format!("- {}", c.criterion))
7034 .collect();
7035
7036 let retry_prompt = format!(
7037 "Your previous response did not meet these criteria:\n{}\n\nPlease provide an improved response to: {}",
7038 feedback.join("\n"),
7039 input
7040 );
7041
7042 let messages = vec![ChatMessage::user(&retry_prompt)];
7043 let retry_response = self
7044 .observe_purpose(
7045 ObservationPurpose::ReflectionEvaluation,
7046 llm.complete(&messages, None),
7047 )
7048 .await
7049 .map_err(|e| AgentError::LLM(e.to_string()))?;
7050
7051 response = retry_response.content.trim().to_string();
7052 attempts += 1;
7053 }
7054 }
7055
7056 async fn evaluate_response_with_config(
7058 &self,
7059 input: &str,
7060 response: &str,
7061 config: &ReflectionConfig,
7062 ) -> Result<EvaluationResult> {
7063 let evaluator_llm = self.role_llm(
7064 ai_agents_llm::LLMRole::ReasoningReflectionEvaluation,
7065 config.evaluator_llm.as_deref(),
7066 || {
7067 config
7068 .evaluator_llm
7069 .as_ref()
7070 .and_then(|alias| self.llm_registry.get(alias).ok())
7071 .or_else(|| self.llm_registry.router().ok())
7072 .or_else(|| self.llm_registry.default().ok())
7073 .ok_or_else(|| AgentError::Config("No LLM available for evaluation".into()))
7074 },
7075 )?;
7076
7077 let criteria = &config.criteria;
7078 let criteria_list = criteria
7079 .iter()
7080 .enumerate()
7081 .map(|(i, c)| format!("{}. {}", i + 1, c))
7082 .collect::<Vec<_>>()
7083 .join("\n");
7084
7085 let prompt = format!(
7086 r#"Evaluate this response against the criteria.
7087
7088User query: "{}"
7089
7090Response to evaluate: "{}"
7091
7092Criteria:
7093{}
7094
7095For each criterion, respond with:
7096- criterion number
7097- PASS or FAIL
7098- brief reason
7099
7100Then provide overall confidence (0.0 to 1.0) and whether it passes overall.
7101
7102Format:
71031. PASS/FAIL - reason
71042. PASS/FAIL - reason
7105...
7106CONFIDENCE: 0.X
7107OVERALL: PASS/FAIL"#,
7108 input, response, criteria_list
7109 );
7110
7111 let messages = vec![ChatMessage::user(&prompt)];
7112 let eval_response = self
7113 .observe_purpose(
7114 ObservationPurpose::ReflectionEvaluation,
7115 evaluator_llm.complete(&messages, None),
7116 )
7117 .await
7118 .map_err(|e| AgentError::LLM(format!("Evaluation failed: {}", e)))?;
7119
7120 let content = eval_response.content.to_uppercase();
7121 let llm_pass = content.contains("OVERALL: PASS");
7122
7123 let confidence = content
7124 .lines()
7125 .find(|l| l.contains("CONFIDENCE:"))
7126 .and_then(|l| {
7127 l.split(':')
7128 .nth(1)
7129 .and_then(|v| v.trim().parse::<f32>().ok())
7130 })
7131 .unwrap_or(if llm_pass { 0.8 } else { 0.4 });
7132
7133 let overall_pass = llm_pass && confidence >= config.pass_threshold;
7136
7137 let mut criteria_results = Vec::new();
7138 for (i, criterion) in criteria.iter().enumerate() {
7139 let line_marker = format!("{}.", i + 1);
7140 let passed = eval_response
7141 .content
7142 .lines()
7143 .find(|l| l.contains(&line_marker))
7144 .map(|l| l.to_uppercase().contains("PASS"))
7145 .unwrap_or(overall_pass);
7146
7147 if passed {
7148 criteria_results.push(CriterionResult::pass(criterion));
7149 } else {
7150 criteria_results.push(CriterionResult::fail(criterion, "Did not meet criterion"));
7151 }
7152 }
7153
7154 Ok(EvaluationResult::new(overall_pass, confidence).with_criteria(criteria_results))
7155 }
7156
7157 async fn process_input(&self, input: &str) -> Result<ProcessData> {
7159 if let Some(processor) = self.get_state_process_processor() {
7160 let purpose = observation_purpose_for_process(processor.input_purpose_hint());
7161 return self
7162 .observe_purpose(purpose, processor.process_input(input))
7163 .await;
7164 }
7165 if let Some(ref processor) = self.process_processor {
7166 let purpose = observation_purpose_for_process(processor.input_purpose_hint());
7167 self.observe_purpose(purpose, processor.process_input(input))
7168 .await
7169 } else {
7170 Ok(ProcessData::new(input))
7171 }
7172 }
7173
7174 async fn process_output(
7176 &self,
7177 output: &str,
7178 input_context: &std::collections::HashMap<String, serde_json::Value>,
7179 ) -> Result<ProcessData> {
7180 if let Some(processor) = self.get_state_process_processor() {
7181 let purpose = observation_purpose_for_process(processor.output_purpose_hint());
7182 return self
7183 .observe_purpose(purpose, processor.process_output(output, input_context))
7184 .await;
7185 }
7186 if let Some(ref processor) = self.process_processor {
7187 let purpose = observation_purpose_for_process(processor.output_purpose_hint());
7188 self.observe_purpose(purpose, processor.process_output(output, input_context))
7189 .await
7190 } else {
7191 Ok(ProcessData::new(output))
7192 }
7193 }
7194
7195 fn get_state_process_processor(&self) -> Option<ProcessProcessor> {
7197 let sm = self.state_machine.as_ref()?;
7198 let def = sm.current_definition()?;
7199 let config = def.process.as_ref()?;
7200 let mut processor = ProcessProcessor::new(config.clone());
7201 if let Some(ref registry) = Some(self.llm_registry.clone()) {
7202 processor = processor.with_llm_registry(registry.clone());
7203 }
7204 processor = processor.with_stage_observer(Arc::new(ObservabilityProcessStageObserver));
7205 Some(processor)
7206 }
7207
7208 async fn check_turn_timeout(&self) -> Result<()> {
7210 let Some(ref sm) = self.state_machine else {
7211 return Ok(());
7212 };
7213 let Some(timeout_state) = sm.check_timeout() else {
7214 return Ok(());
7215 };
7216 let claim_admission = self.disambiguation_admission.write().await;
7217 if sm.check_timeout().as_deref() != Some(timeout_state.as_str()) {
7218 return Ok(());
7219 }
7220 let Some(reservation) = self.reserve_state_transition() else {
7221 return Ok(());
7222 };
7223 let from_state = sm.current();
7224 let expected_state_generation = sm.generation();
7225 let expected_disambiguation_epoch = self.disambiguation_epoch.load(Ordering::SeqCst);
7226 let history_before = sm.history();
7227 drop(claim_admission);
7228
7229 self.execute_state_exit_actions(&from_state).await;
7230
7231 let admission = self.disambiguation_admission.write().await;
7232 if sm.current() != from_state
7233 || sm.generation() != expected_state_generation
7234 || self.disambiguation_epoch.load(Ordering::SeqCst) != expected_disambiguation_epoch
7235 || sm.check_timeout().as_deref() != Some(timeout_state.as_str())
7236 {
7237 return Ok(());
7238 }
7239 sm.transition_to(&timeout_state, "max_turns exceeded")?;
7240 self.invalidate_pending_confirmation("state_timeout").await;
7241 let entered = sm.current();
7242 let is_reentry = Self::state_was_previously_entered(&entered, &from_state, &history_before);
7243 drop(admission);
7244
7245 self.execute_state_enter_actions(&entered, is_reentry).await;
7246 drop(reservation);
7247 info!(to = %entered, "Timeout transition");
7248 Ok(())
7249 }
7250
7251 fn increment_turn(&self) {
7252 if let Some(ref sm) = self.state_machine {
7253 sm.increment_turn();
7254 }
7255 }
7256
7257 fn transitions_available_for_commit(&self) -> Option<(Vec<Transition>, String)> {
7258 let sm = self.state_machine.as_ref()?;
7259 let current = sm.current();
7260 let transitions: Vec<_> = sm
7261 .auto_transitions()
7262 .into_iter()
7263 .filter(|t| match t.cooldown_turns {
7264 Some(cd) if cd > 0 => {
7265 let resolved = sm.config().resolve_full_path(¤t, &t.to);
7266 !sm.is_on_cooldown(&resolved, cd)
7267 }
7268 _ => true,
7269 })
7270 .collect();
7271 Some((transitions, current))
7272 }
7273
7274 fn transition_reason(transition: &Transition) -> String {
7275 if transition.when.is_empty() {
7276 "guard condition met".to_string()
7277 } else {
7278 transition.when.clone()
7279 }
7280 }
7281
7282 fn build_transition_context(
7284 &self,
7285 user_message: &str,
7286 response: &str,
7287 current_state: &str,
7288 staged: Option<&HashMap<String, Value>>,
7289 ) -> TransitionContext {
7290 let context_map = staged
7291 .map(|writes| self.build_context_with_staged(writes))
7292 .unwrap_or_else(|| self.build_context_with_overlays());
7293 TransitionContext::new(user_message, response, current_state).with_context(context_map)
7294 }
7295
7296 async fn select_transition_candidate(
7298 &self,
7299 user_message: &str,
7300 response: &str,
7301 ) -> Result<Option<TransitionCandidate>> {
7302 let Some((transitions, current_state)) = self.transitions_available_for_commit() else {
7303 return Ok(None);
7304 };
7305 let transitions: Vec<Transition> = transitions
7306 .into_iter()
7307 .filter(|transition| matches!(transition.timing, TransitionTiming::PostResponse))
7308 .collect();
7309 if transitions.is_empty() {
7310 return Ok(None);
7311 }
7312 let Some(evaluator) = self.transition_evaluator.as_ref() else {
7313 return Ok(None);
7314 };
7315 let context = self.build_transition_context(user_message, response, ¤t_state, None);
7316 let selected = self
7317 .observe_purpose(
7318 ObservationPurpose::StateTransitionEvaluation,
7319 evaluator.select_transition(&transitions, &context),
7320 )
7321 .await?;
7322 Ok(selected.map(|index| {
7323 let transition = transitions[index].clone();
7324 TransitionCandidate::new(
7325 current_state,
7326 transition.clone(),
7327 Self::transition_reason(&transition),
7328 )
7329 }))
7330 }
7331
7332 fn select_deterministic_transition_candidate(
7334 &self,
7335 user_message: &str,
7336 current_state: &str,
7337 transitions: &[Transition],
7338 staged: &HashMap<String, Value>,
7339 ) -> Option<TransitionCandidate> {
7340 let context = self.build_transition_context(user_message, "", current_state, Some(staged));
7341
7342 for transition in transitions {
7343 if let Some(guard) = transition.guard.as_ref()
7344 && evaluate_guard(guard, &context)
7345 {
7346 return Some(TransitionCandidate::new(
7347 current_state,
7348 transition.clone(),
7349 Self::transition_reason(transition),
7350 ));
7351 }
7352 }
7353
7354 let resolved_intent = context
7355 .context
7356 .get("resolved_intent")
7357 .and_then(Value::as_str)
7358 .filter(|value| !value.is_empty());
7359 if let Some(resolved_intent) = resolved_intent {
7360 for transition in transitions {
7361 if transition.intent.as_deref() == Some(resolved_intent) {
7362 return Some(TransitionCandidate::new(
7363 current_state,
7364 transition.clone(),
7365 Self::transition_reason(transition),
7366 ));
7367 }
7368 }
7369 }
7370
7371 None
7372 }
7373
7374 async fn commit_transition_candidate(&self, candidate: &TransitionCandidate) -> Result<bool> {
7376 self.commit_transition_target(&candidate.from_state, candidate.target(), &candidate.reason)
7377 .await
7378 }
7379
7380 async fn approve_transition_target(&self, from_state: &str, target: &str) -> Result<bool> {
7382 let approved = self.check_state_hitl(Some(from_state), target).await?;
7383 if !approved {
7384 info!(to = %target, "State transition rejected by HITL");
7385 }
7386 Ok(approved)
7387 }
7388
7389 async fn apply_transition_target(
7391 &self,
7392 from_state: &str,
7393 target: &str,
7394 reason: &str,
7395 staged: Option<&HashMap<String, Value>>,
7396 ) -> Result<bool> {
7397 let Some(ref sm) = self.state_machine else {
7398 return Ok(false);
7399 };
7400 let claim_admission = self.disambiguation_admission.write().await;
7401 if sm.current() != from_state {
7402 return Ok(false);
7403 }
7404 let Some(reservation) = self.reserve_state_transition() else {
7405 return Ok(false);
7406 };
7407 let expected_state_generation = sm.generation();
7408 let expected_disambiguation_epoch = self.disambiguation_epoch.load(Ordering::SeqCst);
7409 let history_before = sm.history();
7410 drop(claim_admission);
7411
7412 self.execute_state_exit_actions(from_state).await;
7413
7414 let admission = self.disambiguation_admission.write().await;
7415 if sm.current() != from_state
7416 || sm.generation() != expected_state_generation
7417 || self.disambiguation_epoch.load(Ordering::SeqCst) != expected_disambiguation_epoch
7418 {
7419 return Ok(false);
7420 }
7421 sm.transition_to(target, reason)?;
7422 self.invalidate_pending_confirmation("state_transition")
7423 .await;
7424 sm.reset_no_transition();
7425 if let Some(staged) = staged {
7426 self.commit_staged_context_writes(staged);
7427 }
7428 let entered = sm.current();
7429 let is_reentry = Self::state_was_previously_entered(&entered, from_state, &history_before);
7430 drop(admission);
7431
7432 self.execute_state_enter_actions(&entered, is_reentry).await;
7433 drop(reservation);
7434 self.hooks
7435 .on_state_transition(Some(from_state), &entered, reason)
7436 .await;
7437 info!(from = %from_state, to = %entered, "State transition");
7438 Ok(true)
7439 }
7440
7441 async fn commit_transition_target(
7443 &self,
7444 from_state: &str,
7445 target: &str,
7446 reason: &str,
7447 ) -> Result<bool> {
7448 if !self.approve_transition_target(from_state, target).await? {
7449 return Ok(false);
7450 }
7451 self.apply_transition_target(from_state, target, reason, None)
7452 .await
7453 }
7454
7455 async fn apply_pre_response_transition_candidate(
7457 &self,
7458 candidate: &TransitionCandidate,
7459 staged: &HashMap<String, Value>,
7460 processed_input: &str,
7461 ) -> Result<bool> {
7462 self.commit_root_user_message(processed_input).await?;
7463 self.apply_transition_target(
7464 &candidate.from_state,
7465 candidate.target(),
7466 &candidate.reason,
7467 Some(staged),
7468 )
7469 .await
7470 }
7471
7472 async fn commit_pre_response_transition_candidate(
7474 &self,
7475 candidate: &TransitionCandidate,
7476 staged: &HashMap<String, Value>,
7477 processed_input: &str,
7478 ) -> Result<bool> {
7479 if !self
7480 .approve_transition_target(&candidate.from_state, candidate.target())
7481 .await?
7482 {
7483 return Ok(false);
7484 }
7485 self.apply_pre_response_transition_candidate(candidate, staged, processed_input)
7486 .await
7487 }
7488
7489 async fn handle_transition_miss(&self, current_state: &str) -> Result<bool> {
7491 let Some(ref sm) = self.state_machine else {
7492 return Ok(false);
7493 };
7494 sm.increment_no_transition();
7495 let Some(fallback) = sm.check_fallback() else {
7496 return Ok(false);
7497 };
7498 self.commit_transition_target(current_state, &fallback, "fallback after no transitions")
7499 .await
7500 }
7501
7502 async fn evaluate_transitions(&self, user_message: &str, response: &str) -> Result<bool> {
7504 let Some((transitions, current_state)) = self.transitions_available_for_commit() else {
7505 return Ok(false);
7506 };
7507 if transitions.is_empty() {
7508 return Ok(false);
7509 }
7510 if let Some(candidate) = self
7511 .select_transition_candidate(user_message, response)
7512 .await?
7513 {
7514 return self.commit_transition_candidate(&candidate).await;
7515 }
7516 self.handle_transition_miss(¤t_state).await
7517 }
7518
7519 async fn try_pre_response_transition(
7521 &self,
7522 processed_input: &str,
7523 ) -> Result<Option<AgentResponse>> {
7524 let optimization = &self.runtime_config.optimization;
7525 if !optimization.enabled || !optimization.pre_response_deterministic_transitions {
7526 return Ok(None);
7527 }
7528 let Some((transitions, current_state)) = self.transitions_available_for_commit() else {
7529 return Ok(None);
7530 };
7531 let eligible: Vec<Transition> = transitions
7532 .into_iter()
7533 .filter(|transition| !transition.requires_response)
7534 .filter(|transition| matches!(transition.timing, TransitionTiming::PreResponse))
7535 .collect();
7536 if eligible.is_empty() {
7537 return Ok(None);
7538 }
7539
7540 let empty_staged = HashMap::new();
7541 let mut extracted_staged: Option<HashMap<String, Value>> = None;
7542 let mut selected: Option<(TransitionCandidate, HashMap<String, Value>)> = None;
7543
7544 for transition in &eligible {
7545 let use_extractors = optimization.pre_response_extractors || transition.run_extractors;
7546 let staged_for_eval = if use_extractors {
7547 if extracted_staged.is_none() {
7548 extracted_staged =
7549 Some(self.run_context_extractors_staged(processed_input).await?);
7550 }
7551 extracted_staged.as_ref().unwrap_or(&empty_staged)
7552 } else {
7553 &empty_staged
7554 };
7555
7556 if let Some(candidate) = self.select_deterministic_transition_candidate(
7557 processed_input,
7558 ¤t_state,
7559 std::slice::from_ref(transition),
7560 staged_for_eval,
7561 ) {
7562 let staged_for_commit = if use_extractors {
7563 staged_for_eval.clone()
7564 } else {
7565 HashMap::new()
7566 };
7567 selected = Some((candidate, staged_for_commit));
7568 break;
7569 }
7570 }
7571
7572 let Some((candidate, staged)) = selected else {
7573 return Ok(None);
7574 };
7575
7576 if !self
7577 .commit_pre_response_transition_candidate(&candidate, &staged, processed_input)
7578 .await?
7579 {
7580 return Ok(None);
7581 }
7582 self.redispatch_current_state(processed_input)
7583 .await
7584 .map(Some)
7585 }
7586
7587 async fn try_speculative_branches(
7592 &self,
7593 processed_input: &str,
7594 input_context: &HashMap<String, Value>,
7595 ) -> Result<Option<AgentResponse>> {
7596 let optimization = &self.runtime_config.optimization;
7597 if !optimization.enabled {
7598 return Ok(None);
7599 }
7600
7601 let effective_reasoning_mode = self.get_effective_reasoning_config().mode.clone();
7602 if !matches!(
7603 effective_reasoning_mode,
7604 ReasoningMode::None | ReasoningMode::Auto
7605 ) {
7606 return Ok(None);
7607 }
7608
7609 let mut transition_enabled =
7610 optimization.speculative_state_transitions && self.has_parallel_transition_candidates();
7611 let mut skill_enabled = optimization.speculative_skill_routing
7612 && self.skill_router.is_some()
7613 && self.pending_skill_id.read().is_none();
7614 let mut reasoning_enabled = optimization.speculative_reasoning_auto
7615 && matches!(effective_reasoning_mode, ReasoningMode::Auto);
7616
7617 if matches!(effective_reasoning_mode, ReasoningMode::Auto)
7618 && (!reasoning_enabled || optimization.max_speculative_llm_calls_per_turn < 2)
7619 {
7620 return Ok(None);
7621 }
7622
7623 if !transition_enabled && !skill_enabled && !reasoning_enabled {
7624 return Ok(None);
7625 }
7626
7627 let mut optional_slots = optimization.max_parallel_runtime_tasks.saturating_sub(1);
7628 let mut speculative_call_slots = optimization
7629 .max_speculative_llm_calls_per_turn
7630 .saturating_sub(1);
7631 if reasoning_enabled {
7632 if optional_slots == 0 || speculative_call_slots == 0 {
7633 return Ok(None);
7634 }
7635 optional_slots -= 1;
7636 speculative_call_slots -= 1;
7637 }
7638 if transition_enabled {
7639 if optional_slots == 0 {
7640 transition_enabled = false;
7641 } else {
7642 optional_slots -= 1;
7643 }
7644 }
7645 if skill_enabled && (optional_slots == 0 || speculative_call_slots == 0) {
7646 skill_enabled = false;
7647 }
7648
7649 if !transition_enabled && !skill_enabled && !reasoning_enabled {
7650 return Ok(None);
7651 }
7652
7653 let main_kind = if transition_enabled {
7654 RuntimeOptimizationKind::ParallelStateTransition
7655 } else if skill_enabled {
7656 RuntimeOptimizationKind::SpeculativeSkillRouting
7657 } else {
7658 RuntimeOptimizationKind::SpeculativeReasoningAuto
7659 };
7660 if !self.reserve_active_speculative_llm_call(main_kind) {
7661 return Ok(None);
7662 }
7663
7664 let mut branch_set = ScheduledBranchSet::new(optimization.max_parallel_runtime_tasks)?;
7665 let main_branch = RuntimeBranch::new(
7666 RuntimeTaskPurpose::MainResponse,
7667 main_kind,
7668 RuntimeTaskPriority::Normal,
7669 RuntimeCommitBehavior::FinalResponse,
7670 );
7671 let transition_branch = RuntimeBranch::new(
7672 RuntimeTaskPurpose::StateTransition,
7673 RuntimeOptimizationKind::ParallelStateTransition,
7674 RuntimeTaskPriority::Critical,
7675 RuntimeCommitBehavior::TransitionDecision,
7676 );
7677 let skill_branch = RuntimeBranch::new(
7678 RuntimeTaskPurpose::SkillRouting,
7679 RuntimeOptimizationKind::SpeculativeSkillRouting,
7680 RuntimeTaskPriority::High,
7681 RuntimeCommitBehavior::SkillSelection,
7682 );
7683 let reasoning_branch = RuntimeBranch::new(
7684 RuntimeTaskPurpose::ReasoningJudge,
7685 RuntimeOptimizationKind::SpeculativeReasoningAuto,
7686 RuntimeTaskPriority::Normal,
7687 RuntimeCommitBehavior::ReasoningDecision,
7688 );
7689 let main_id = main_branch.branch_id();
7690 let transition_id = transition_branch.branch_id();
7691 let skill_id = skill_branch.branch_id();
7692 let reasoning_id = reasoning_branch.branch_id();
7693
7694 let main_id_for_future = main_id.clone();
7695 if !branch_set.schedule(
7696 main_branch,
7697 Box::pin(async move {
7698 match crate::optimization::observability::with_branch_observation(
7699 &main_id_for_future,
7700 main_kind,
7701 RuntimeCommitBehavior::FinalResponse,
7702 self.generate_main_response_draft(processed_input, &ReasoningMode::None),
7703 )
7704 .await
7705 {
7706 Ok(draft) => RuntimeBranchResult::MainDraft(draft),
7707 Err(error) => RuntimeBranchResult::Failed(error),
7708 }
7709 }),
7710 ) {
7711 return Ok(None);
7712 }
7713
7714 if transition_enabled {
7715 let transition_id_for_future = transition_id.clone();
7716 if !branch_set.schedule(
7717 transition_branch,
7718 Box::pin(async move {
7719 match crate::optimization::observability::with_branch_observation(
7720 &transition_id_for_future,
7721 RuntimeOptimizationKind::ParallelStateTransition,
7722 RuntimeCommitBehavior::TransitionDecision,
7723 self.select_parallel_transition_candidate(processed_input),
7724 )
7725 .await
7726 {
7727 Ok(ParallelTransitionSelection::Candidate(candidate)) => {
7728 RuntimeBranchResult::Transition(Some(candidate))
7729 }
7730 Ok(ParallelTransitionSelection::NoMatch) => {
7731 RuntimeBranchResult::Transition(None)
7732 }
7733 Ok(ParallelTransitionSelection::ReservationExhausted) => {
7734 RuntimeBranchResult::Cancelled
7735 }
7736 Err(error) => RuntimeBranchResult::Failed(error),
7737 }
7738 }),
7739 ) {
7740 transition_enabled = false;
7741 }
7742 }
7743
7744 if skill_enabled {
7745 let skill_id_for_future = skill_id.clone();
7746 if !branch_set.schedule(
7747 skill_branch,
7748 Box::pin(async move {
7749 if !self.reserve_active_speculative_llm_call(
7750 RuntimeOptimizationKind::SpeculativeSkillRouting,
7751 ) {
7752 return RuntimeBranchResult::Cancelled;
7753 }
7754 match crate::optimization::observability::with_branch_observation(
7755 &skill_id_for_future,
7756 RuntimeOptimizationKind::SpeculativeSkillRouting,
7757 RuntimeCommitBehavior::SkillSelection,
7758 self.select_skill_candidate(processed_input),
7759 )
7760 .await
7761 {
7762 Ok(candidate) => RuntimeBranchResult::Skill(candidate),
7763 Err(error) => RuntimeBranchResult::Failed(error),
7764 }
7765 }),
7766 ) {
7767 skill_enabled = false;
7768 }
7769 }
7770
7771 if reasoning_enabled {
7772 let reasoning_id_for_future = reasoning_id.clone();
7773 if !branch_set.schedule(
7774 reasoning_branch,
7775 Box::pin(async move {
7776 if !self.reserve_active_speculative_llm_call(
7777 RuntimeOptimizationKind::SpeculativeReasoningAuto,
7778 ) {
7779 return RuntimeBranchResult::Cancelled;
7780 }
7781 match crate::optimization::observability::with_branch_observation(
7782 &reasoning_id_for_future,
7783 RuntimeOptimizationKind::SpeculativeReasoningAuto,
7784 RuntimeCommitBehavior::ReasoningDecision,
7785 self.determine_reasoning_mode_strict(processed_input),
7786 )
7787 .await
7788 {
7789 Ok(mode) => RuntimeBranchResult::Reasoning(mode),
7790 Err(error) => RuntimeBranchResult::Failed(error),
7791 }
7792 }),
7793 ) {
7794 reasoning_enabled = false;
7795 }
7796 }
7797
7798 if matches!(effective_reasoning_mode, ReasoningMode::Auto) && !reasoning_enabled {
7799 self.finalize_pending_branches(branch_set.cancel_pending());
7800 return Ok(None);
7801 }
7802
7803 if !transition_enabled && !skill_enabled && !reasoning_enabled {
7804 self.finalize_pending_branches(branch_set.cancel_pending());
7805 return Ok(None);
7806 }
7807
7808 let mut main_pending = true;
7809 let mut skill_pending = skill_enabled;
7810 let mut reasoning_pending = reasoning_enabled;
7811 let mut transition_finalized = !transition_enabled;
7812 let mut skill_finalized = !skill_enabled && self.skill_router.is_none();
7815 let mut reasoning_finalized = !reasoning_enabled;
7816 let mut main_result: Option<Result<MainResponseDraft>> = None;
7817 let mut transition_candidate: Option<TransitionCandidate> = None;
7818 let mut skill_candidate: Option<SkillCandidate> = None;
7819 let mut reasoning_decision: Option<ReasoningMode> = None;
7820 let mut transition_fallback_required = false;
7821 let mut skill_fallback_required = false;
7822 let mut reasoning_fallback_required = false;
7823
7824 loop {
7825 if let Some(candidate) = transition_candidate.take() {
7826 if self
7827 .approve_transition_target(&candidate.from_state, candidate.target())
7828 .await?
7829 {
7830 self.finalize_pending_branches(branch_set.cancel_pending());
7832 if !main_pending {
7833 self.finalize_branch_loss(
7834 &main_id,
7835 main_kind,
7836 RuntimeCommitBehavior::FinalResponse,
7837 false,
7838 main_result.as_ref().map(|result| result.is_err()),
7839 );
7840 }
7841 if skill_enabled && !skill_pending {
7842 self.finalize_branch_loss(
7843 &skill_id,
7844 RuntimeOptimizationKind::SpeculativeSkillRouting,
7845 RuntimeCommitBehavior::SkillSelection,
7846 false,
7847 Some(false),
7848 );
7849 }
7850 if reasoning_enabled && !reasoning_pending {
7851 self.finalize_branch_loss(
7852 &reasoning_id,
7853 RuntimeOptimizationKind::SpeculativeReasoningAuto,
7854 RuntimeCommitBehavior::ReasoningDecision,
7855 false,
7856 Some(false),
7857 );
7858 }
7859 if !self
7860 .apply_pre_response_transition_candidate(
7861 &candidate,
7862 &HashMap::new(),
7863 processed_input,
7864 )
7865 .await?
7866 {
7867 self.finalize_optional_branch(
7868 &transition_id,
7869 RuntimeOptimizationKind::ParallelStateTransition,
7870 RuntimeCommitBehavior::TransitionDecision,
7871 "discarded",
7872 false,
7873 );
7874 return Ok(None);
7875 }
7876 self.finalize_optional_branch(
7877 &transition_id,
7878 RuntimeOptimizationKind::ParallelStateTransition,
7879 RuntimeCommitBehavior::TransitionDecision,
7880 "committed",
7881 true,
7882 );
7883 return self
7884 .redispatch_current_state(processed_input)
7885 .await
7886 .map(Some);
7887 }
7888 self.finalize_optional_branch(
7889 &transition_id,
7890 RuntimeOptimizationKind::ParallelStateTransition,
7891 RuntimeCommitBehavior::TransitionDecision,
7892 "discarded",
7893 false,
7894 );
7895 transition_finalized = true;
7896 }
7897
7898 if transition_finalized
7907 && !skill_finalized
7908 && !skill_enabled
7909 && self.skill_router.is_some()
7910 {
7911 match self.select_skill_candidate(processed_input).await {
7912 Ok(Some(candidate)) => skill_candidate = Some(candidate),
7913 Ok(None) => {}
7914 Err(error) => {
7915 self.finalize_pending_branches(branch_set.cancel_pending());
7917 return Err(error);
7918 }
7919 }
7920 skill_finalized = true;
7921 }
7922
7923 if transition_finalized && skill_candidate.is_some() {
7924 let candidate = skill_candidate.take().unwrap();
7925 if skill_enabled {
7927 self.finalize_optional_branch(
7928 &skill_id,
7929 RuntimeOptimizationKind::SpeculativeSkillRouting,
7930 RuntimeCommitBehavior::SkillSelection,
7931 "committed",
7932 true,
7933 );
7934 }
7935 if !main_pending {
7936 self.finalize_branch_loss(
7937 &main_id,
7938 main_kind,
7939 RuntimeCommitBehavior::FinalResponse,
7940 false,
7941 main_result.as_ref().map(|result| result.is_err()),
7942 );
7943 }
7944 if reasoning_enabled && !reasoning_pending {
7945 self.finalize_branch_loss(
7946 &reasoning_id,
7947 RuntimeOptimizationKind::SpeculativeReasoningAuto,
7948 RuntimeCommitBehavior::ReasoningDecision,
7949 false,
7950 Some(false),
7951 );
7952 }
7953 self.finalize_pending_branches(branch_set.cancel_pending());
7954 return self
7955 .commit_winning_skill_candidate(candidate, processed_input, input_context)
7956 .await;
7957 }
7958
7959 if transition_finalized
7960 && skill_finalized
7961 && let Some(reasoning_mode) = reasoning_decision.take()
7962 {
7963 if !matches!(reasoning_mode, ReasoningMode::None) {
7964 self.finalize_optional_branch(
7965 &reasoning_id,
7966 RuntimeOptimizationKind::SpeculativeReasoningAuto,
7967 RuntimeCommitBehavior::ReasoningDecision,
7968 "committed",
7969 true,
7970 );
7971 if !main_pending {
7972 self.finalize_branch_loss(
7973 &main_id,
7974 main_kind,
7975 RuntimeCommitBehavior::FinalResponse,
7976 false,
7977 main_result.as_ref().map(|result| result.is_err()),
7978 );
7979 }
7980 self.finalize_pending_branches(branch_set.cancel_pending());
7981 self.commit_root_user_message(processed_input).await?;
7982 return if matches!(reasoning_mode, ReasoningMode::PlanAndExecute) {
7983 self.handle_plan_and_execute(processed_input, input_context, true)
7984 .await
7985 .map(Some)
7986 } else {
7987 self.run_committed_response_loop_with_reasoning(
7988 processed_input,
7989 input_context,
7990 reasoning_mode,
7991 true,
7992 )
7993 .await
7994 .map(Some)
7995 };
7996 }
7997 self.finalize_optional_branch(
7998 &reasoning_id,
7999 RuntimeOptimizationKind::SpeculativeReasoningAuto,
8000 RuntimeCommitBehavior::ReasoningDecision,
8001 "committed",
8002 true,
8003 );
8004 reasoning_finalized = true;
8005 }
8006
8007 if transition_finalized && skill_finalized && reasoning_finalized {
8008 if transition_fallback_required
8009 || skill_fallback_required
8010 || reasoning_fallback_required
8011 {
8012 if !main_pending {
8013 self.finalize_branch_loss(
8014 &main_id,
8015 main_kind,
8016 RuntimeCommitBehavior::FinalResponse,
8017 false,
8018 main_result.as_ref().map(|result| result.is_err()),
8019 );
8020 }
8021 self.finalize_pending_branches(branch_set.cancel_pending());
8022 return Ok(None);
8023 }
8024
8025 if let Some(result) = main_result.take() {
8026 let draft = match result {
8027 Ok(draft) => draft,
8028 Err(error) => {
8029 self.finalize_optional_branch(
8030 &main_id,
8031 main_kind,
8032 RuntimeCommitBehavior::FinalResponse,
8033 "failed",
8034 false,
8035 );
8036 self.finalize_pending_branches(branch_set.cancel_pending());
8037 return Err(error);
8038 }
8039 };
8040 self.finalize_optional_branch(
8041 &main_id,
8042 main_kind,
8043 RuntimeCommitBehavior::FinalResponse,
8044 "committed",
8045 true,
8046 );
8047 self.finalize_pending_branches(branch_set.cancel_pending());
8048 return self
8049 .commit_main_response_draft(
8050 processed_input,
8051 input_context,
8052 draft,
8053 ReasoningMode::None,
8054 reasoning_enabled,
8055 )
8056 .await
8057 .map(Some);
8058 }
8059 }
8060
8061 if branch_set.is_empty() {
8062 return Ok(None);
8063 }
8064
8065 let Some(outcome) = branch_set.next_completed().await else {
8066 return Ok(None);
8067 };
8068 let branch_id = outcome.branch.branch_id();
8069 match outcome.result {
8070 RuntimeBranchResult::MainDraft(draft) => {
8071 main_pending = false;
8072 main_result = Some(Ok(draft));
8073 }
8074 RuntimeBranchResult::Transition(candidate) => {
8075 if let Some(candidate) = candidate {
8076 transition_candidate = Some(candidate);
8077 } else {
8078 self.finalize_optional_branch(
8079 &transition_id,
8080 RuntimeOptimizationKind::ParallelStateTransition,
8081 RuntimeCommitBehavior::TransitionDecision,
8082 "discarded",
8083 false,
8084 );
8085 transition_finalized = true;
8086 }
8087 }
8088 RuntimeBranchResult::Skill(candidate) => {
8089 skill_pending = false;
8090 if let Some(candidate) = candidate {
8091 skill_candidate = Some(candidate);
8092 } else {
8093 self.finalize_optional_branch(
8094 &skill_id,
8095 RuntimeOptimizationKind::SpeculativeSkillRouting,
8096 RuntimeCommitBehavior::SkillSelection,
8097 "discarded",
8098 false,
8099 );
8100 skill_finalized = true;
8101 }
8102 }
8103 RuntimeBranchResult::Reasoning(mode) => {
8104 reasoning_pending = false;
8105 reasoning_decision = Some(mode);
8106 }
8107 RuntimeBranchResult::Failed(error) => {
8108 if branch_id == main_id {
8109 main_pending = false;
8110 main_result = Some(Err(error));
8111 } else if branch_id == transition_id {
8112 self.finalize_optional_branch(
8113 &transition_id,
8114 RuntimeOptimizationKind::ParallelStateTransition,
8115 RuntimeCommitBehavior::TransitionDecision,
8116 "failed",
8117 false,
8118 );
8119 transition_finalized = true;
8120 } else if branch_id == skill_id {
8121 skill_pending = false;
8122 self.finalize_optional_branch(
8123 &skill_id,
8124 RuntimeOptimizationKind::SpeculativeSkillRouting,
8125 RuntimeCommitBehavior::SkillSelection,
8126 "failed",
8127 false,
8128 );
8129 skill_finalized = true;
8130 } else if branch_id == reasoning_id {
8131 reasoning_pending = false;
8132 self.finalize_optional_branch(
8133 &reasoning_id,
8134 RuntimeOptimizationKind::SpeculativeReasoningAuto,
8135 RuntimeCommitBehavior::ReasoningDecision,
8136 "failed",
8137 false,
8138 );
8139 reasoning_finalized = true;
8140 }
8141 }
8142 RuntimeBranchResult::Cancelled => {
8143 self.finalize_optional_branch(
8144 &branch_id,
8145 outcome.branch.optimization,
8146 outcome.branch.commit_behavior,
8147 "cancelled",
8148 false,
8149 );
8150 if branch_id == main_id {
8151 main_pending = false;
8152 main_result =
8153 Some(Err(AgentError::Other("main branch cancelled".to_string())));
8154 } else if branch_id == transition_id {
8155 transition_finalized = true;
8156 transition_fallback_required = true;
8157 } else if branch_id == skill_id {
8158 skill_pending = false;
8159 skill_finalized = true;
8160 skill_fallback_required = true;
8161 } else if branch_id == reasoning_id {
8162 reasoning_pending = false;
8163 reasoning_finalized = true;
8164 reasoning_fallback_required = true;
8165 }
8166 }
8167 }
8168 }
8169 }
8170
8171 fn finalize_pending_branches(&self, branches: Vec<RuntimeBranch>) {
8172 for branch in branches {
8173 self.finalize_optional_branch(
8174 &branch.branch_id(),
8175 branch.optimization,
8176 branch.commit_behavior,
8177 "cancelled",
8178 false,
8179 );
8180 }
8181 }
8182
8183 fn finalize_branch_loss(
8188 &self,
8189 branch_id: &str,
8190 optimization: RuntimeOptimizationKind,
8191 commit_behavior: RuntimeCommitBehavior,
8192 pending: bool,
8193 completed_failed: Option<bool>,
8194 ) {
8195 let status = if pending {
8196 "cancelled"
8197 } else if completed_failed.unwrap_or(false) {
8198 "failed"
8199 } else {
8200 "discarded"
8201 };
8202 self.finalize_optional_branch(branch_id, optimization, commit_behavior, status, false);
8203 }
8204
8205 fn finalize_optional_branch(
8210 &self,
8211 branch_id: &str,
8212 optimization: RuntimeOptimizationKind,
8213 commit_behavior: RuntimeCommitBehavior,
8214 status: &str,
8215 winner: bool,
8216 ) {
8217 crate::optimization::observability::finalize_branch(
8218 self.observability_manager.as_ref(),
8219 branch_id,
8220 status,
8221 winner,
8222 optimization,
8223 commit_behavior,
8224 );
8225 }
8226
8227 fn has_parallel_transition_candidates(&self) -> bool {
8232 self.transitions_available_for_commit()
8233 .map(|(transitions, _)| {
8234 transitions
8235 .iter()
8236 .any(|transition| matches!(transition.timing, TransitionTiming::Parallel))
8237 })
8238 .unwrap_or(false)
8239 }
8240
8241 async fn select_parallel_transition_candidate(
8247 &self,
8248 processed_input: &str,
8249 ) -> Result<ParallelTransitionSelection> {
8250 let Some((transitions, current_state)) = self.transitions_available_for_commit() else {
8251 return Ok(ParallelTransitionSelection::NoMatch);
8252 };
8253 let parallel: Vec<Transition> = transitions
8254 .into_iter()
8255 .filter(|transition| matches!(transition.timing, TransitionTiming::Parallel))
8256 .filter(|transition| !transition.requires_response)
8257 .collect();
8258 if parallel.is_empty() {
8259 return Ok(ParallelTransitionSelection::NoMatch);
8260 }
8261 let empty_staged = HashMap::new();
8262 if let Some(candidate) = self.select_deterministic_transition_candidate(
8263 processed_input,
8264 ¤t_state,
8265 ¶llel,
8266 &empty_staged,
8267 ) {
8268 return Ok(ParallelTransitionSelection::Candidate(candidate));
8269 }
8270 let when_transitions: Vec<(usize, &Transition)> = parallel
8271 .iter()
8272 .enumerate()
8273 .filter(|(_, transition)| !transition.when.trim().is_empty())
8274 .collect();
8275 if when_transitions.is_empty() {
8276 return Ok(ParallelTransitionSelection::NoMatch);
8277 }
8278 let llm = self.role_llm(ai_agents_llm::LLMRole::StateTransition, None, || {
8279 self.llm_registry
8280 .router()
8281 .or_else(|_| self.llm_registry.default())
8282 .map_err(|e| AgentError::Config(e.to_string()))
8283 })?;
8284 let conditions = when_transitions
8285 .iter()
8286 .enumerate()
8287 .map(|(display_idx, (_, transition))| {
8288 format!("{}. {}", display_idx + 1, transition.when)
8289 })
8290 .collect::<Vec<_>>()
8291 .join("\n");
8292 if !self
8293 .reserve_active_speculative_llm_call(RuntimeOptimizationKind::ParallelStateTransition)
8294 {
8295 return Ok(ParallelTransitionSelection::ReservationExhausted);
8296 }
8297 let context_preview = self.branch_context_preview();
8298 let prompt = format!(
8299 "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-{}).",
8300 current_state,
8301 processed_input,
8302 context_preview,
8303 conditions,
8304 when_transitions.len()
8305 );
8306 let response = self
8307 .observe_purpose(
8308 ObservationPurpose::StateTransitionEvaluation,
8309 llm.complete(&[ChatMessage::user(prompt)], None),
8310 )
8311 .await
8312 .map_err(|e| AgentError::LLM(e.to_string()))?;
8313 let choice = response.content.trim().parse::<usize>().unwrap_or(0);
8314 if choice == 0 || choice > when_transitions.len() {
8315 return Ok(ParallelTransitionSelection::NoMatch);
8316 }
8317 let transition = when_transitions[choice - 1].1.clone();
8318 Ok(ParallelTransitionSelection::Candidate(
8319 TransitionCandidate::new(
8320 current_state,
8321 transition.clone(),
8322 Self::transition_reason(&transition),
8323 ),
8324 ))
8325 }
8326
8327 async fn redispatch_current_state(&self, processed_input: &str) -> Result<AgentResponse> {
8329 const MAX_REDISPATCH_DEPTH: u32 = 3;
8330 let current_depth = *self.redispatch_depth.read();
8331 if current_depth >= MAX_REDISPATCH_DEPTH {
8332 warn!(depth = current_depth, "Re-dispatch depth limit reached");
8333 let response = AgentResponse::new("");
8334 self.finish_turn_if_root(&response).await?;
8335 return Ok(response);
8336 }
8337 *self.redispatch_depth.write() += 1;
8338 if let Some(context) = self.active_turn_context.write().as_mut() {
8339 context.enter_redispatch();
8340 }
8341 let result = Box::pin(self.run_loop_internal(processed_input)).await;
8342 *self.redispatch_depth.write() -= 1;
8343 if let Some(context) = self.active_turn_context.write().as_mut() {
8344 context.exit_redispatch();
8345 }
8346 let response = result?;
8347 self.finish_turn_if_root(&response).await?;
8348 Ok(response)
8349 }
8350
8351 async fn finish_turn_if_root(&self, response: &AgentResponse) -> Result<()> {
8353 if *self.redispatch_depth.read() == 0 {
8354 self.post_turn_session_lifecycle().await?;
8355 if let Some(context) = self.active_turn_context.write().as_mut() {
8356 context.mark_post_turn_lifecycle_completed();
8357 }
8358 self.hooks.on_response(response).await;
8359 self.end_root_turn();
8360 }
8361 Ok(())
8362 }
8363
8364 async fn execute_state_exit_actions(&self, state_path: &str) {
8366 if let Some(ref sm) = self.state_machine
8367 && let Some(def) = sm.get_definition(state_path)
8368 && !def.on_exit.is_empty()
8369 {
8370 debug!(state = %state_path, count = def.on_exit.len(), "Executing on_exit actions");
8371 self.execute_state_actions(&def.on_exit).await;
8372 }
8373 }
8374
8375 fn state_was_previously_entered(
8377 state_path: &str,
8378 from_state: &str,
8379 history_before: &[StateTransitionEvent],
8380 ) -> bool {
8381 state_path == from_state
8382 || history_before
8383 .iter()
8384 .any(|event| event.from == state_path || event.to == state_path)
8385 }
8386
8387 async fn execute_state_enter_actions(&self, state_path: &str, is_reentry: bool) {
8389 if let Some(ref sm) = self.state_machine
8390 && let Some(def) = sm.get_definition(state_path)
8391 {
8392 if is_reentry && !def.on_reenter.is_empty() {
8393 debug!(state = %state_path, count = def.on_reenter.len(), "Executing on_reenter actions");
8394 self.execute_state_actions(&def.on_reenter).await;
8395 } else if !def.on_enter.is_empty() {
8396 debug!(state = %state_path, count = def.on_enter.len(), "Executing on_enter actions");
8397 self.execute_state_actions(&def.on_enter).await;
8398 }
8399 }
8400 }
8401
8402 async fn execute_state_actions(&self, actions: &[StateAction]) {
8404 for (action_index, action) in actions.iter().enumerate() {
8405 match action {
8406 StateAction::Tool { tool, args } => {
8407 let raw_args = args.clone().unwrap_or(Value::Object(Default::default()));
8408 let args_value = self.render_action_args(&raw_args);
8409 let state = self.state_machine.as_ref().map(|sm| sm.current());
8410 let request = ToolExecutionRequest::new(
8411 uuid::Uuid::new_v4().to_string(),
8412 tool.clone(),
8413 args_value,
8414 ToolCallSource::StateAction {
8415 state,
8416 action_index,
8417 },
8418 );
8419 match self.execute_tool_record(request).await {
8420 Ok(record) if record.success => {
8421 debug!(tool = %record.canonical_id, "State action: tool executed");
8422 let _ = self.context_manager.set(
8423 "last_tool_result",
8424 serde_json::Value::String(record.model_output_string()),
8425 );
8426 let _ = self.context_manager.set(
8427 "last_tool_record",
8428 serde_json::to_value(record).unwrap_or(Value::Null),
8429 );
8430 }
8431 Ok(record) => {
8432 warn!(tool = %record.canonical_id, error = %record.output, "State action: tool failed");
8433 }
8434 Err(e) => {
8435 warn!(tool = %tool, error = %e, "State action: tool failed")
8436 }
8437 }
8438 }
8439 StateAction::Skill { skill } => {
8440 if let Some(ref executor) = self.skill_executor {
8441 if let Some(def) = self.skills.iter().find(|s| s.id == *skill) {
8442 match executor
8443 .execute_with_invoker(def, "", serde_json::json!({}), self)
8444 .await
8445 {
8446 Ok(_) => debug!(skill = %skill, "State action: skill executed"),
8447 Err(e) => {
8448 warn!(skill = %skill, error = %e, "State action: skill failed")
8449 }
8450 }
8451 } else {
8452 warn!(skill = %skill, "State action: skill not found");
8453 }
8454 }
8455 }
8456 StateAction::SetContext { set_context } => {
8457 for (key, value) in set_context {
8458 if let Err(e) = self.context_manager.set(key, value.clone()) {
8459 warn!(key = %key, error = %e, "State action: set_context failed");
8460 } else {
8461 debug!(key = %key, "State action: context set");
8462 }
8463 }
8464 }
8465 StateAction::Prompt {
8466 prompt,
8467 llm,
8468 store_as,
8469 } => {
8470 let llm_result = if let Some(alias) = llm {
8471 self.llm_registry.get(alias)
8472 } else {
8473 self.llm_registry.default()
8474 };
8475 match llm_result {
8476 Ok(llm_provider) => {
8477 let context = self.build_context_with_overlays();
8479 let rendered_prompt = self
8480 .template_renderer
8481 .render(prompt, &context)
8482 .unwrap_or_else(|_| prompt.clone());
8483 let recent =
8484 self.memory.get_messages(Some(5)).await.unwrap_or_default();
8485 let mut messages: Vec<ChatMessage> = recent;
8486 messages.push(ChatMessage::user(&rendered_prompt));
8487 match self
8488 .observe_purpose(
8489 ObservationPurpose::StateAction,
8490 llm_provider.complete(&messages, None),
8491 )
8492 .await
8493 {
8494 Ok(response) => {
8495 if let Some(key) = store_as {
8496 let _ = self
8497 .context_manager
8498 .set(key, Value::String(response.content));
8499 debug!(key = %key, "State action: prompt result stored");
8500 }
8501 }
8502 Err(e) => {
8503 warn!(error = %e, "State action: prompt LLM call failed");
8504 }
8505 }
8506 }
8507 Err(e) => {
8508 warn!(error = %e, "State action: LLM not found for prompt");
8509 }
8510 }
8511 }
8512 }
8513 }
8514 }
8515
8516 async fn run_context_extractors_staged(
8518 &self,
8519 user_message: &str,
8520 ) -> Result<HashMap<String, Value>> {
8521 let extractors = match &self.state_machine {
8522 Some(sm) => match sm.current_definition() {
8523 Some(def) if !def.extract.is_empty() => def.extract.clone(),
8524 _ => return Ok(HashMap::new()),
8525 },
8526 None => return Ok(HashMap::new()),
8527 };
8528
8529 let mut staged = HashMap::new();
8530 for extractor in &extractors {
8531 let prompt = if let Some(ref custom) = extractor.llm_extract {
8532 format!(
8533 "User message:\n\"{}\"\n\nInstruction:\n{}",
8534 user_message, custom
8535 )
8536 } else if let Some(ref desc) = extractor.description {
8537 format!(
8538 "From the following message, extract: {}\n\n\
8539 Message: \"{}\"\n\n\
8540 If the information is present, return ONLY the extracted value.\n\
8541 If NOT present, return exactly: __NONE__",
8542 desc, user_message
8543 )
8544 } else {
8545 continue;
8546 };
8547
8548 let llm = match self.role_llm(
8549 ai_agents_llm::LLMRole::StateExtract,
8550 extractor.llm.as_deref(),
8551 || {
8552 self.llm_registry
8553 .get(extractor.llm.as_deref().unwrap_or("router"))
8554 .or_else(|_| self.llm_registry.get("router"))
8555 .or_else(|_| self.llm_registry.get("default"))
8556 .map_err(|e| AgentError::Config(e.to_string()))
8557 },
8558 ) {
8559 Ok(llm) => llm,
8560 Err(e) => {
8561 if self.llm_registry.router_roles().is_some() {
8562 return Err(e);
8563 }
8564 warn!(key = %extractor.key, error = %e, "Extractor LLM not found");
8565 continue;
8566 }
8567 };
8568
8569 let messages = vec![ChatMessage::user(&prompt)];
8570 match self
8571 .observe_purpose(
8572 ObservationPurpose::ContextExtraction,
8573 llm.complete(&messages, None),
8574 )
8575 .await
8576 {
8577 Ok(response) => {
8578 let value = response.content.trim().to_string();
8579 if value != "__NONE__" && !value.is_empty() {
8580 staged.insert(
8581 extractor.key.clone(),
8582 serde_json::Value::String(value.clone()),
8583 );
8584 debug!(key = %extractor.key, value = %value, "Context extracted");
8585 } else if extractor.required {
8586 warn!(key = %extractor.key, "Required extraction returned no value");
8587 }
8588 }
8589 Err(e) => {
8590 warn!(key = %extractor.key, error = %e, "Context extraction LLM call failed");
8591 }
8592 }
8593 }
8594 Ok(staged)
8595 }
8596
8597 fn commit_staged_context_writes(&self, staged: &HashMap<String, Value>) {
8598 for (key, value) in staged {
8599 if let Err(error) = self.context_manager.update(key, value.clone()) {
8600 warn!(key = %key, error = %error, "staged context write failed");
8601 }
8602 }
8603 }
8604
8605 async fn run_context_extractors(&self, user_message: &str) -> Result<()> {
8608 let staged = self.run_context_extractors_staged(user_message).await?;
8609 self.commit_staged_context_writes(&staged);
8610 Ok(())
8611 }
8612
8613 async fn check_memory_compression(&self) -> Result<()> {
8614 if self.memory.needs_compression() {
8615 let result = self.memory.compress(None).await?;
8616 if let CompressResult::Compressed {
8617 messages_summarized,
8618 new_summary_length,
8619 tokens_saved,
8620 } = result
8621 {
8622 let event = MemoryCompressEvent::new(
8623 messages_summarized,
8624 tokens_saved,
8625 new_summary_length as u32,
8626 );
8627 self.hooks.on_memory_compress(&event).await;
8628 debug!(
8629 messages = messages_summarized,
8630 tokens_saved = tokens_saved,
8631 "Memory compressed"
8632 );
8633 }
8634 }
8635
8636 self.handle_memory_overflow().await?;
8638 self.check_memory_budget().await;
8639
8640 Ok(())
8641 }
8642
8643 async fn check_memory_budget(&self) {
8644 let Some(ref budget) = self.memory_token_budget else {
8645 return;
8646 };
8647
8648 let context = match self.memory.get_context().await {
8649 Ok(ctx) => ctx,
8650 Err(_) => return,
8651 };
8652
8653 let used_tokens = context.estimated_tokens();
8655 if budget.is_over_warn_threshold(used_tokens) {
8656 let event = MemoryBudgetEvent::new("memory", used_tokens, budget.total);
8657 self.hooks.on_memory_budget_warning(&event).await;
8658 debug!(
8659 used = used_tokens,
8660 total = budget.total,
8661 percent = event.usage_percent,
8662 "Memory budget warning"
8663 );
8664 }
8665
8666 if let Some(ref summary) = context.summary {
8668 let summary_tokens = ai_agents_memory::estimate_tokens(summary);
8669 let summary_budget = budget.allocation.summary;
8670 if summary_budget > 0 {
8671 let warn_threshold =
8672 (summary_budget as f64 * budget.warn_at_percent as f64 / 100.0) as u32;
8673 if summary_tokens >= warn_threshold {
8674 let event = MemoryBudgetEvent::new("summary", summary_tokens, summary_budget);
8675 self.hooks.on_memory_budget_warning(&event).await;
8676 }
8677 }
8678 }
8679
8680 let recent_tokens: u32 = context
8682 .messages
8683 .iter()
8684 .map(ai_agents_memory::estimate_message_tokens)
8685 .sum();
8686 let recent_budget = budget.allocation.recent_messages;
8687 if recent_budget > 0 {
8688 let warn_threshold =
8689 (recent_budget as f64 * budget.warn_at_percent as f64 / 100.0) as u32;
8690 if recent_tokens >= warn_threshold {
8691 let event = MemoryBudgetEvent::new("recent_messages", recent_tokens, recent_budget);
8692 self.hooks.on_memory_budget_warning(&event).await;
8693 }
8694 }
8695
8696 let relationship_budget = budget.allocation.relationships;
8697 if relationship_budget > 0 {
8698 let relationship_tokens = self
8699 .relationship_memory_text()
8700 .map(|text| ai_agents_memory::estimate_tokens(&text))
8701 .unwrap_or(0);
8702 let warn_threshold =
8703 (relationship_budget as f64 * budget.warn_at_percent as f64 / 100.0) as u32;
8704 if relationship_tokens >= warn_threshold {
8705 let event = MemoryBudgetEvent::new(
8706 "relationships",
8707 relationship_tokens,
8708 relationship_budget,
8709 );
8710 self.hooks.on_memory_budget_warning(&event).await;
8711 }
8712 }
8713 }
8714
8715 async fn handle_memory_overflow(&self) -> Result<()> {
8716 let Some(ref budget) = self.memory_token_budget else {
8717 return Ok(());
8718 };
8719
8720 let context = self.memory.get_context().await?;
8721 let used_tokens = context.estimated_tokens();
8722
8723 if used_tokens <= budget.total {
8724 return Ok(());
8725 }
8726
8727 match budget.overflow_strategy {
8728 OverflowStrategy::TruncateOldest => {
8729 let tokens_to_free = used_tokens - budget.total;
8730 let messages_to_evict = self.calculate_eviction_count(tokens_to_free);
8731 if messages_to_evict > 0 {
8732 self.evict_messages(messages_to_evict, EvictionReason::TokenBudgetExceeded)
8733 .await?;
8734 }
8735 }
8736 OverflowStrategy::SummarizeMore => {
8737 let max_attempts = context.total_messages.max(1);
8738 for _ in 0..max_attempts {
8739 match self.memory.compress(None).await? {
8740 CompressResult::Compressed {
8741 messages_summarized,
8742 ..
8743 } if messages_summarized > 0 => {
8744 let context = self.memory.get_context().await?;
8745 if context.estimated_tokens() <= budget.total {
8746 return Ok(());
8747 }
8748 }
8749 _ => break,
8750 }
8751 }
8752 let context = self.memory.get_context().await?;
8753 let used_tokens = context.estimated_tokens();
8754 if used_tokens > budget.total {
8755 return Err(AgentError::MemoryBudgetExceeded {
8756 used: used_tokens,
8757 budget: budget.total,
8758 });
8759 }
8760 }
8761 OverflowStrategy::Error => {
8762 return Err(AgentError::MemoryBudgetExceeded {
8763 used: used_tokens,
8764 budget: budget.total,
8765 });
8766 }
8767 }
8768 Ok(())
8769 }
8770
8771 fn calculate_eviction_count(&self, tokens_to_free: u32) -> usize {
8772 ((tokens_to_free as f64 / 50.0).ceil() as usize).max(1)
8774 }
8775
8776 async fn evict_messages(&self, count: usize, reason: EvictionReason) -> Result<()> {
8777 let evicted = self.memory.evict_oldest(count).await?;
8778 if !evicted.is_empty() {
8779 let event = MemoryEvictEvent {
8780 reason,
8781 messages_evicted: evicted.len(),
8782 importance_scores: vec![],
8783 };
8784 self.hooks.on_memory_evict(&event).await;
8785 debug!(count = evicted.len(), "Messages evicted from memory");
8786 }
8787 Ok(())
8788 }
8789
8790 #[instrument(skip(self, input), fields(agent = %self.info.name))]
8791 async fn determine_reasoning_mode(&self, input: &str) -> Result<ReasoningMode> {
8792 match self.determine_reasoning_mode_strict(input).await {
8793 Ok(mode) => Ok(mode),
8794 Err(error @ AgentError::Config(_)) if self.llm_registry.router_roles().is_some() => {
8795 Err(error)
8796 }
8797 Err(_) => Ok(ReasoningMode::None),
8798 }
8799 }
8800
8801 async fn determine_reasoning_mode_strict(&self, input: &str) -> Result<ReasoningMode> {
8803 let effective_config = self.get_effective_reasoning_config();
8804
8805 if !matches!(effective_config.mode, ReasoningMode::Auto) {
8806 return Ok(effective_config.mode.clone());
8807 }
8808
8809 let judge_llm = self.optional_role_llm(
8810 ai_agents_llm::LLMRole::ReasoningSelection,
8811 effective_config.judge_llm.as_deref(),
8812 || {
8813 effective_config
8814 .judge_llm
8815 .as_ref()
8816 .and_then(|alias| self.llm_registry.get(alias).ok())
8817 .or_else(|| self.llm_registry.router().ok())
8818 .or_else(|| self.llm_registry.default().ok())
8819 },
8820 )?;
8821
8822 let Some(llm) = judge_llm else {
8823 return Ok(ReasoningMode::None);
8824 };
8825
8826 let prompt = format!(
8827 r#"Analyze this user request and determine the appropriate reasoning mode.
8828
8829User request: "{}"
8830
8831Choose ONE of these modes:
8832- none: Simple queries, greetings, direct answers (fastest)
8833- cot: Complex analysis, multi-step reasoning, math problems
8834- react: Tasks requiring multiple tool calls with observation
8835- plan_and_execute: Complex multi-step tasks requiring coordination
8836
8837Respond with ONLY the mode name (none, cot, react, or plan_and_execute)."#,
8838 input
8839 );
8840
8841 let messages = vec![ChatMessage::user(&prompt)];
8842 let response = self
8843 .observe_purpose(
8844 ObservationPurpose::ReflectionDecision,
8845 llm.complete(&messages, None),
8846 )
8847 .await
8848 .map_err(|e| AgentError::LLM(e.to_string()))?;
8849
8850 let mode_str = response.content.trim().to_lowercase();
8851 Ok(match mode_str.as_str() {
8852 "cot" => ReasoningMode::CoT,
8853 "react" => ReasoningMode::React,
8854 "plan_and_execute" => ReasoningMode::PlanAndExecute,
8855 _ => ReasoningMode::None,
8856 })
8857 }
8858
8859 fn build_cot_system_prompt(&self, base_prompt: &str) -> String {
8860 format!(
8861 "{}\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>",
8862 base_prompt
8863 )
8864 }
8865
8866 fn build_react_system_prompt(&self, base_prompt: &str) -> String {
8867 format!(
8868 "{}\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>",
8869 base_prompt
8870 )
8871 }
8872
8873 async fn generate_plan(&self, input: &str) -> Result<Plan> {
8876 let effective = self.get_effective_reasoning_config();
8877 let planning_config = effective.get_planning();
8878
8879 let planner_llm = self.role_llm(
8880 ai_agents_llm::LLMRole::ReasoningPlanning,
8881 planning_config.and_then(|c| c.planner_llm.as_deref()),
8882 || {
8883 planning_config
8884 .and_then(|c| c.planner_llm.as_ref())
8885 .and_then(|alias| self.llm_registry.get(alias).ok())
8886 .or_else(|| self.llm_registry.router().ok())
8887 .or_else(|| self.llm_registry.default().ok())
8888 .ok_or_else(|| AgentError::Config("No LLM available for planning".into()))
8889 },
8890 )?;
8891
8892 let mut available_tool_ids = self.get_available_tool_ids().await?;
8893 let mut available_skills: Vec<String> = self.skills.iter().map(|s| s.id.clone()).collect();
8894
8895 if let Some(config) = planning_config {
8897 if !config.available.tools.is_all() {
8898 available_tool_ids.retain(|t| config.available.tools.allows(t));
8899 }
8900 if !config.available.skills.is_all() {
8901 available_skills.retain(|s| config.available.skills.allows(s));
8902 }
8903 }
8904
8905 let tool_descriptions: Vec<String> = available_tool_ids
8908 .iter()
8909 .filter_map(|id| {
8910 self.tools.get(id).map(|tool| {
8911 let schema = tool.input_schema();
8912 let args_desc = schema
8913 .get("properties")
8914 .and_then(|p| serde_json::to_string(p).ok())
8915 .unwrap_or_else(|| "{}".to_string());
8916 format!(
8917 "- {} ({}): {}\n Arguments: {}",
8918 id,
8919 tool.name(),
8920 tool.description(),
8921 args_desc
8922 )
8923 })
8924 })
8925 .collect();
8926
8927 let tools_section = if tool_descriptions.is_empty() {
8928 "Available tools: none".to_string()
8929 } else {
8930 format!("Available tools:\n{}", tool_descriptions.join("\n"))
8931 };
8932
8933 let skills_section = if available_skills.is_empty() {
8934 "Available skills: none".to_string()
8935 } else {
8936 format!("Available skills: {}", available_skills.join(", "))
8937 };
8938
8939 let prompt = format!(
8940 r#"Create a step-by-step plan to accomplish this goal.
8941
8942Goal: "{}"
8943
8944{}
8945
8946{}
8947
8948Create a plan with clear steps. For each step, specify:
8949- description: What this step accomplishes
8950- action_type: "tool", "skill", "think", or "respond"
8951- action_target: The tool/skill id (if applicable)
8952- args: The arguments object matching the tool's schema (if action_type is "tool")
8953- dependencies: List of step IDs this depends on (empty if none)
8954
8955Respond in JSON format:
8956{{
8957 "steps": [
8958 {{"id": "step1", "description": "...", "action_type": "tool", "action_target": "tool_id", "args": {{"required_field": "value"}}, "dependencies": []}},
8959 {{"id": "step2", "description": "...", "action_type": "think", "action_target": "...", "dependencies": ["step1"]}}
8960 ]
8961}}"#,
8962 input, tools_section, skills_section,
8963 );
8964
8965 let messages = vec![ChatMessage::user(&prompt)];
8966 let response = self
8967 .observe_purpose(
8968 ObservationPurpose::PlanGeneration,
8969 planner_llm.complete(&messages, None),
8970 )
8971 .await
8972 .map_err(|e| AgentError::LLM(format!("Planning failed: {}", e)))?;
8973
8974 let mut plan = Plan::new(input);
8975
8976 if let Some(json_start) = response.content.find('{')
8977 && let Some(json_end) = response.content.rfind('}')
8978 {
8979 let json_str = &response.content[json_start..=json_end];
8980 if let Ok(parsed) = serde_json::from_str::<serde_json::Value>(json_str)
8981 && let Some(steps) = parsed.get("steps").and_then(|s| s.as_array())
8982 {
8983 for step_value in steps {
8984 let id = step_value
8985 .get("id")
8986 .and_then(|v| v.as_str())
8987 .unwrap_or("step");
8988 let desc = step_value
8989 .get("description")
8990 .and_then(|v| v.as_str())
8991 .unwrap_or("");
8992 let action_type = step_value
8993 .get("action_type")
8994 .and_then(|v| v.as_str())
8995 .unwrap_or("think");
8996 let action_target = step_value
8997 .get("action_target")
8998 .and_then(|v| v.as_str())
8999 .unwrap_or("");
9000 let args = step_value
9001 .get("args")
9002 .cloned()
9003 .unwrap_or(serde_json::json!({}));
9004 let deps: Vec<String> = step_value
9005 .get("dependencies")
9006 .and_then(|v| v.as_array())
9007 .map(|arr| {
9008 arr.iter()
9009 .filter_map(|v| v.as_str().map(String::from))
9010 .collect()
9011 })
9012 .unwrap_or_default();
9013
9014 let action = match action_type {
9015 "tool" => PlanAction::tool(action_target, args),
9016 "skill" => PlanAction::skill(action_target),
9017 "respond" => PlanAction::respond(action_target),
9018 _ => PlanAction::think(desc),
9019 };
9020
9021 let step = PlanStep::new(desc, action)
9022 .with_id(id)
9023 .with_dependencies(deps);
9024 plan.add_step(step);
9025 }
9026 }
9027 }
9028
9029 if plan.steps.is_empty() {
9030 plan.add_step(PlanStep::new(
9031 "Process the request",
9032 PlanAction::think(input),
9033 ));
9034 plan.add_step(PlanStep::new(
9035 "Provide response",
9036 PlanAction::respond("Answer based on analysis"),
9037 ));
9038 }
9039
9040 Ok(plan)
9041 }
9042
9043 async fn execute_plan(&self, plan: &mut Plan) -> Result<String> {
9044 let llm = self.get_state_llm()?;
9045 let mut results: HashMap<String, serde_json::Value> = HashMap::new();
9046 let effective = self.get_effective_reasoning_config();
9047 let max_steps = effective.get_planning().map(|c| c.max_steps).unwrap_or(10);
9048
9049 plan.status = PlanStatus::InProgress;
9050
9051 for step_idx in 0..plan.steps.len().min(max_steps as usize) {
9052 let step = &plan.steps[step_idx];
9053
9054 let deps_satisfied = step.dependencies.iter().all(|dep| {
9055 plan.steps
9056 .iter()
9057 .find(|s| &s.id == dep)
9058 .map(|s| s.status.is_completed())
9059 .unwrap_or(false)
9060 });
9061
9062 if !deps_satisfied {
9063 continue;
9064 }
9065
9066 plan.steps[step_idx].mark_running();
9067
9068 let result = match &plan.steps[step_idx].action {
9069 PlanAction::Tool { tool, args } => {
9070 let has_dep_results = plan.steps[step_idx]
9076 .dependencies
9077 .iter()
9078 .any(|dep| results.contains_key(dep));
9079
9080 let final_args = if has_dep_results {
9081 let dep_context: String = plan.steps[step_idx]
9082 .dependencies
9083 .iter()
9084 .filter_map(|dep| results.get(dep).map(|r| format!("{}: {}", dep, r)))
9085 .collect::<Vec<_>>()
9086 .join("\n");
9087
9088 let tool_schema = self
9089 .tools
9090 .get(tool)
9091 .map(|t| {
9092 let schema = t.input_schema();
9093 let props = schema
9094 .get("properties")
9095 .and_then(|p| serde_json::to_string(p).ok())
9096 .unwrap_or_else(|| "{}".to_string());
9097 format!(
9098 "{}: {}\nArguments schema: {}",
9099 t.id(),
9100 t.description(),
9101 props
9102 )
9103 })
9104 .unwrap_or_default();
9105
9106 let step_desc = &plan.steps[step_idx].description;
9107 let arg_prompt = format!(
9108 "Generate the JSON arguments for a tool call.\n\n\
9109 Tool: {}\n\n\
9110 Task: {}\n\n\
9111 Previous step results:\n{}\n\n\
9112 Planner's draft arguments: {}\n\n\
9113 Produce ONLY a valid JSON object with the correct argument values.\n\
9114 Use actual values from the previous step results, not template references.",
9115 tool_schema,
9116 step_desc,
9117 dep_context,
9118 serde_json::to_string(args).unwrap_or_default()
9119 );
9120 let messages = vec![ChatMessage::user(&arg_prompt)];
9121 match self
9122 .observe_purpose(
9123 ObservationPurpose::PlanStep,
9124 llm.complete(&messages, None),
9125 )
9126 .await
9127 {
9128 Ok(resp) => {
9129 let content = resp.content.trim();
9130 let json_start = content.find('{');
9132 let json_end = content.rfind('}');
9133 if let (Some(start), Some(end)) = (json_start, json_end) {
9134 serde_json::from_str(&content[start..=end])
9135 .unwrap_or_else(|_| args.clone())
9136 } else {
9137 args.clone()
9138 }
9139 }
9140 Err(_) => args.clone(),
9141 }
9142 } else {
9143 args.clone()
9144 };
9145
9146 let request = ToolExecutionRequest::new(
9147 uuid::Uuid::new_v4().to_string(),
9148 tool.clone(),
9149 final_args,
9150 ToolCallSource::Plan {
9151 step_index: step_idx,
9152 },
9153 );
9154 match self.execute_tool_record(request).await {
9155 Ok(record) if record.success => {
9156 serde_json::json!({ "output": record.model_output_string() })
9157 }
9158 Ok(record) => {
9159 plan.steps[step_idx].mark_failed(record.model_output_string());
9160 continue;
9161 }
9162 Err(e) => {
9163 plan.steps[step_idx].mark_failed(e.to_string());
9164 continue;
9165 }
9166 }
9167 }
9168 PlanAction::Skill { skill } => {
9169 if let Some(skill_def) = self.skills.iter().find(|s| &s.id == skill) {
9170 if let Some(ref executor) = self.skill_executor {
9171 match executor
9172 .execute_with_invoker(skill_def, "", serde_json::json!({}), self)
9173 .await
9174 {
9175 Ok(output) => serde_json::json!({ "output": output }),
9176 Err(e) => {
9177 plan.steps[step_idx].mark_failed(e.to_string());
9178 continue;
9179 }
9180 }
9181 } else {
9182 serde_json::json!({ "output": "Skill executor not available" })
9183 }
9184 } else {
9185 plan.steps[step_idx].mark_failed("Skill not found");
9186 continue;
9187 }
9188 }
9189 PlanAction::Think { prompt } => {
9190 let context: String = results
9191 .iter()
9192 .map(|(k, v)| format!("{}: {}", k, v))
9193 .collect::<Vec<_>>()
9194 .join("\n");
9195
9196 let think_prompt = format!("Context:\n{}\n\nTask: {}", context, prompt);
9197 let messages = vec![ChatMessage::user(&think_prompt)];
9198
9199 match self
9200 .observe_purpose(
9201 ObservationPurpose::PlanStep,
9202 llm.complete(&messages, None),
9203 )
9204 .await
9205 {
9206 Ok(resp) => serde_json::json!({ "output": resp.content }),
9207 Err(e) => {
9208 plan.steps[step_idx].mark_failed(e.to_string());
9209 continue;
9210 }
9211 }
9212 }
9213 PlanAction::Respond { template } => {
9214 let context: String = results
9215 .iter()
9216 .map(|(k, v)| format!("{}: {}", k, v))
9217 .collect::<Vec<_>>()
9218 .join("\n");
9219
9220 let respond_prompt = format!(
9221 "Based on this context:\n{}\n\nGenerate a response following this template/instruction: {}",
9222 context, template
9223 );
9224 let messages = vec![ChatMessage::user(&respond_prompt)];
9225
9226 match self
9227 .observe_purpose(
9228 ObservationPurpose::PlanStep,
9229 llm.complete(&messages, None),
9230 )
9231 .await
9232 {
9233 Ok(resp) => serde_json::json!({ "output": resp.content }),
9234 Err(e) => {
9235 plan.steps[step_idx].mark_failed(e.to_string());
9236 continue;
9237 }
9238 }
9239 }
9240 };
9241
9242 results.insert(plan.steps[step_idx].id.clone(), result.clone());
9243 plan.steps[step_idx].mark_completed(Some(result));
9244 }
9245
9246 let has_failures = plan.steps.iter().any(|s| s.status.is_failed());
9248 if has_failures {
9249 let failed_ids: Vec<String> = plan
9250 .steps
9251 .iter()
9252 .filter(|s| s.status.is_failed())
9253 .map(|s| s.id.clone())
9254 .collect();
9255 plan.status = PlanStatus::Failed {
9256 error: format!("Steps failed: {}", failed_ids.join(", ")),
9257 };
9258 } else {
9259 plan.status = PlanStatus::Completed;
9260 }
9261
9262 let all_outputs: Vec<String> = plan
9264 .steps
9265 .iter()
9266 .filter(|s| s.status.is_completed())
9267 .filter_map(|s| {
9268 s.result
9269 .as_ref()
9270 .and_then(|r| r.get("output"))
9271 .and_then(|o| o.as_str())
9272 .map(|o| format!("{}: {}", s.description, o))
9273 })
9274 .collect();
9275
9276 if all_outputs.is_empty() {
9277 return Ok("Plan execution completed but produced no results.".to_string());
9278 }
9279
9280 if all_outputs.len() == 1 {
9281 return Ok(all_outputs.into_iter().next().unwrap());
9282 }
9283
9284 let context = all_outputs.join("\n\n");
9286 let prompt = format!(
9287 "You completed a multi-step plan for: \"{}\"\n\nStep results:\n{}\n\nProvide a coherent final response that synthesizes these results.",
9288 plan.goal, context
9289 );
9290 let messages = vec![ChatMessage::user(&prompt)];
9291 match self
9292 .observe_purpose(ObservationPurpose::PlanStep, llm.complete(&messages, None))
9293 .await
9294 {
9295 Ok(resp) => Ok(resp.content.trim().to_string()),
9296 Err(_) => Ok(context),
9297 }
9298 }
9299
9300 fn extract_thinking(&self, content: &str) -> (Option<String>, String) {
9301 if let Some(start) = content.find("<thinking>")
9302 && let Some(end) = content.find("</thinking>")
9303 {
9304 let thinking = content[start + 10..end].trim().to_string();
9305 let answer = content[end + 11..].trim().to_string();
9306 return (Some(thinking), answer);
9307 }
9308 (None, content.to_string())
9309 }
9310
9311 fn format_response_with_thinking(&self, thinking: Option<&str>, answer: &str) -> String {
9312 match self.get_effective_reasoning_config().output {
9313 ReasoningOutput::Hidden => answer.to_string(),
9314 ReasoningOutput::Visible => {
9315 if let Some(t) = thinking {
9316 format!("Thinking:\n{}\n\nAnswer:\n{}", t, answer)
9317 } else {
9318 answer.to_string()
9319 }
9320 }
9321 ReasoningOutput::Tagged => {
9322 if let Some(t) = thinking {
9323 format!("<thinking>{}</thinking>\n{}", t, answer)
9324 } else {
9325 answer.to_string()
9326 }
9327 }
9328 }
9329 }
9330
9331 fn disambiguation_question_response(
9334 question: &ClarificationQuestion,
9335 detection: &AmbiguityDetectionResult,
9336 awaiting_confirmation: bool,
9337 ) -> AgentResponse {
9338 let status = if awaiting_confirmation {
9339 "awaiting_confirmation"
9340 } else {
9341 "awaiting_clarification"
9342 };
9343 AgentResponse::new(&question.question).with_metadata(
9344 "disambiguation",
9345 serde_json::json!({
9346 "status": status,
9347 "options": question.options,
9348 "clarifying": question.clarifying,
9349 "detection": {
9350 "type": detection.ambiguity_type,
9351 "confidence": detection.confidence,
9352 "what_is_unclear": detection.what_is_unclear,
9353 }
9354 }),
9355 )
9356 }
9357
9358 async fn resolve_disambiguation(&self, input: &str) -> Result<DisambiguationDispatch> {
9370 let Some(ref disambiguator) = self.disambiguation_manager else {
9371 return Ok(DisambiguationDispatch::Proceed(input.to_string()));
9372 };
9373 let disambiguation_context = self.build_disambiguation_context().await?;
9374
9375 let state_override = self
9377 .state_machine
9378 .as_ref()
9379 .and_then(|sm| sm.current_definition())
9380 .and_then(|def| def.disambiguation.clone());
9381
9382 let state_generation = self
9383 .state_machine
9384 .as_ref()
9385 .map(|state_machine| state_machine.generation());
9386 let disambiguation_epoch = self.disambiguation_epoch.load(Ordering::SeqCst);
9387 let mut disambiguation_result = self
9388 .observe_purpose(
9389 ObservationPurpose::DisambiguationDetection,
9390 disambiguator.process_input_with_override(
9391 input,
9392 &disambiguation_context,
9393 state_override.as_ref(),
9394 None,
9395 ),
9396 )
9397 .await?;
9398 let current_state_generation = self
9399 .state_machine
9400 .as_ref()
9401 .map(|state_machine| state_machine.generation());
9402 if current_state_generation != state_generation
9403 || self.disambiguation_epoch.load(Ordering::SeqCst) != disambiguation_epoch
9404 {
9405 disambiguator.clear_pending().await;
9406 *self.pending_skill_id.write() = None;
9407 disambiguation_result = DisambiguationResult::Abandoned { new_input: None };
9408 info!(
9409 confirmation_event = "invalidated",
9410 invalidation_reason = "state_generation_changed",
9411 "Disambiguation result invalidated before redispatch"
9412 );
9413 }
9414 match disambiguation_result {
9415 DisambiguationResult::Clear => {
9416 debug!("Input is clear, proceeding normally");
9417 Ok(DisambiguationDispatch::Proceed(input.to_string()))
9418 }
9419 DisambiguationResult::NeedsClarification {
9420 question,
9421 detection,
9422 } => {
9423 let admission = match self
9424 .admit_disambiguation_redispatch(disambiguation_epoch, state_generation)
9425 .await
9426 {
9427 Ok(admission) => admission,
9428 Err(error) => {
9429 *self.pending_skill_id.write() = None;
9430 return Err(error);
9431 }
9432 };
9433 let awaiting_confirmation = disambiguator.has_pending_confirmation().await;
9434 info!(
9435 ambiguity_type = ?detection.ambiguity_type,
9436 confidence = detection.confidence,
9437 "Input requires clarification"
9438 );
9439
9440 self.commit_root_user_message(input).await?;
9443 self.memory
9444 .add_message(ChatMessage::assistant(&question.question))
9445 .await?;
9446
9447 let response = Self::disambiguation_question_response(
9448 &question,
9449 &detection,
9450 awaiting_confirmation,
9451 );
9452 drop(admission);
9453 self.finish_turn_if_root(&response).await?;
9454 Ok(DisambiguationDispatch::Terminal(response))
9455 }
9456 DisambiguationResult::Clarified {
9457 enriched_input,
9458 resolved,
9459 ..
9460 } => {
9461 let admission = match self
9462 .admit_disambiguation_redispatch(disambiguation_epoch, state_generation)
9463 .await
9464 {
9465 Ok(admission) => admission,
9466 Err(error) => {
9467 *self.pending_skill_id.write() = None;
9468 return Err(error);
9469 }
9470 };
9471 info!(
9472 resolved_count = resolved.len(),
9473 enriched = %enriched_input,
9474 "Input clarified, injecting resolved intent into context"
9475 );
9476
9477 for (key, value) in &resolved {
9480 let context_key = format!("disambiguation.{}", key);
9481 let _ = self.context_manager.set(&context_key, value.clone());
9482 }
9483
9484 if let Some(intent) = resolved.get("intent") {
9485 let _ = self.context_manager.set("resolved_intent", intent.clone());
9486 }
9487
9488 let _ = self
9489 .context_manager
9490 .set("disambiguation.resolved", serde_json::Value::Bool(true));
9491
9492 let skill_id = self.pending_skill_id.read().clone();
9496 drop(admission);
9497 if let Some(skill_id) = skill_id {
9498 info!(skill_id = %skill_id, "Re-checking skill disambiguation on clarified input");
9499 return Ok(DisambiguationDispatch::RecheckSkill {
9500 skill_id,
9501 enriched_input,
9502 disambiguation_epoch,
9503 state_generation,
9504 });
9505 }
9506 Ok(DisambiguationDispatch::Proceed(enriched_input))
9507 }
9508 DisambiguationResult::ProceedWithBestGuess { enriched_input } => {
9509 info!("Proceeding with best guess interpretation");
9510
9511 let skill_id = self.pending_skill_id.read().clone();
9513 if let Some(skill_id) = skill_id {
9514 info!(skill_id = %skill_id, "Re-checking skill disambiguation on best-guess input");
9515 return Ok(DisambiguationDispatch::RecheckSkill {
9516 skill_id,
9517 enriched_input,
9518 disambiguation_epoch,
9519 state_generation,
9520 });
9521 }
9522 Ok(DisambiguationDispatch::Proceed(enriched_input))
9523 }
9524 DisambiguationResult::GiveUp { reason } => {
9525 *self.pending_skill_id.write() = None;
9526 warn!(reason = %reason, "Disambiguation gave up");
9527 let apology = self
9528 .generate_localized_apology(
9529 "Generate a brief, polite apology saying you couldn't understand the request. Be concise.",
9530 &reason,
9531 )
9532 .await
9533 .unwrap_or_else(|_| {
9534 format!("I'm sorry, I couldn't understand your request: {}", reason)
9535 });
9536 let response = AgentResponse::new(&apology);
9537 self.finish_turn_if_root(&response).await?;
9538 Ok(DisambiguationDispatch::Terminal(response))
9539 }
9540 DisambiguationResult::Escalate { reason } => {
9541 *self.pending_skill_id.write() = None;
9542 info!(reason = %reason, "Escalating to human");
9543 if let Some(ref hitl) = self.hitl_engine {
9544 let trigger =
9545 ApprovalTrigger::condition("disambiguation_escalation", reason.clone());
9546 let mut context_map = HashMap::new();
9547 context_map.insert("original_input".to_string(), serde_json::json!(input));
9548 context_map.insert("reason".to_string(), serde_json::json!(&reason));
9549 let check_result = HITLCheckResult::required(
9550 trigger,
9551 context_map,
9552 format!("User request needs human assistance: {}", reason),
9553 Some(hitl.config().default_timeout_seconds),
9554 );
9555 let result = self.request_hitl_approval(check_result).await?;
9556 if matches!(
9557 result,
9558 ApprovalResult::Approved | ApprovalResult::Modified { .. }
9559 ) {
9560 return Ok(DisambiguationDispatch::Proceed(input.to_string()));
9562 }
9563 }
9564 let apology = self
9565 .generate_localized_apology(
9566 "Explain briefly that you're transferring the user to a human agent for help.",
9567 &reason,
9568 )
9569 .await
9570 .unwrap_or_else(|_| {
9571 format!("I need human assistance to help with your request: {}", reason)
9572 });
9573 let response = AgentResponse::new(&apology);
9574 self.finish_turn_if_root(&response).await?;
9575 Ok(DisambiguationDispatch::Terminal(response))
9576 }
9577 DisambiguationResult::Abandoned { new_input } => {
9578 *self.pending_skill_id.write() = None;
9579
9580 info!(
9581 has_new_input = new_input.is_some(),
9582 "Clarification abandoned by user"
9583 );
9584
9585 self.commit_root_user_message(input).await?;
9586
9587 match new_input {
9588 Some(fresh_input) => {
9589 Ok(DisambiguationDispatch::Proceed(fresh_input))
9592 }
9593 None => {
9594 let ack = self
9596 .generate_localized_apology(
9597 "The user changed their mind about their previous request. \
9598 Generate a brief, friendly acknowledgment (e.g. 'OK, no problem. What else can I help with?'). \
9599 Do NOT apologize excessively. Be concise.",
9600 "User abandoned clarification",
9601 )
9602 .await
9603 .unwrap_or_else(|_| {
9604 "OK, no problem. What else can I help with?".to_string()
9605 });
9606
9607 self.memory
9608 .add_message(ChatMessage::assistant(&ack))
9609 .await?;
9610
9611 let response = AgentResponse::new(&ack);
9612 self.finish_turn_if_root(&response).await?;
9613 Ok(DisambiguationDispatch::Terminal(response))
9614 }
9615 }
9616 }
9617 }
9618 }
9619
9620 async fn prepare_turn_context(&self) -> Result<()> {
9624 if !self.context_initialized.load(Ordering::SeqCst) {
9625 self.context_manager.initialize().await?;
9626 self.context_initialized.store(true, Ordering::SeqCst);
9627 debug!("Context manager initialized (defaults, env, builtins)");
9628 }
9629
9630 self.check_turn_timeout().await?;
9631 self.context_manager.refresh_per_turn().await?;
9632 self.context_manager.validate()
9633 }
9634
9635 async fn run_loop(&self, input: &str) -> Result<AgentResponse> {
9638 self.init_storage().await?;
9642 self.begin_root_turn();
9643 let _root_cleanup = RootTurnCleanup::new(self);
9644 info!(input_len = input.len(), "Starting chat");
9645
9646 self.hooks.on_message_received(input).await;
9647
9648 self.prepare_turn_context().await?;
9649
9650 self.clear_disambiguation_context();
9653
9654 let input_to_run = match self.resolve_disambiguation(input).await? {
9657 DisambiguationDispatch::Terminal(response) => return Ok(response),
9658 DisambiguationDispatch::RecheckSkill {
9659 skill_id,
9660 enriched_input,
9661 disambiguation_epoch,
9662 state_generation,
9663 } => {
9664 return self
9665 .recheck_skill_disambiguation(
9666 &skill_id,
9667 &enriched_input,
9668 disambiguation_epoch,
9669 state_generation,
9670 )
9671 .await;
9672 }
9673 DisambiguationDispatch::Proceed(input) => input,
9674 };
9675
9676 self.run_loop_internal(&input_to_run).await
9677 }
9678
9679 async fn generate_localized_apology(&self, instruction: &str, reason: &str) -> Result<String> {
9682 let llm = self.role_llm(ai_agents_llm::LLMRole::DisambiguationResponse, None, || {
9683 self.llm_registry.router().map_err(|e| {
9684 AgentError::LLM(format!(
9685 "Router LLM not available for localized response: {}",
9686 e
9687 ))
9688 })
9689 })?;
9690
9691 let recent: Vec<String> = self
9692 .memory
9693 .get_messages(Some(3))
9694 .await?
9695 .iter()
9696 .map(|m| m.content.clone())
9697 .collect();
9698
9699 let context_hint = if recent.is_empty() {
9700 String::new()
9701 } else {
9702 format!(
9703 "\nRecent conversation (detect the user's language from this):\n{}\n",
9704 recent.join("\n")
9705 )
9706 };
9707
9708 let prompt = format!(
9709 "{}\nReason: {}\n{}Respond in the same language as the user. Output ONLY the message, nothing else.",
9710 instruction, reason, context_hint
9711 );
9712
9713 let messages = vec![ChatMessage::user(&prompt)];
9714 let response = self
9715 .observe_purpose(
9716 ObservationPurpose::DisambiguationClarification,
9717 llm.complete(&messages, None),
9718 )
9719 .await
9720 .map_err(|e| AgentError::LLM(format!("Localized response generation failed: {}", e)))?;
9721
9722 Ok(response.content.trim().to_string())
9723 }
9724
9725 fn render_action_args(&self, args: &Value) -> Value {
9729 let context = self.build_context_with_overlays();
9730 match args {
9731 Value::Object(map) => {
9732 let mut rendered = serde_json::Map::new();
9733 for (k, v) in map {
9734 match v {
9735 Value::String(s) if s.contains("{{") => {
9736 match self.template_renderer.render(s, &context) {
9737 Ok(rendered_str) => {
9738 rendered.insert(k.clone(), Value::String(rendered_str));
9739 }
9740 Err(_) => {
9741 rendered.insert(k.clone(), v.clone());
9742 }
9743 }
9744 }
9745 _ => {
9746 rendered.insert(k.clone(), v.clone());
9747 }
9748 }
9749 }
9750 Value::Object(rendered)
9751 }
9752 _ => args.clone(),
9753 }
9754 }
9755
9756 fn clear_disambiguation_context(&self) {
9758 let _ = self
9759 .context_manager
9760 .set("resolved_intent", serde_json::Value::Null);
9761
9762 let all = self.context_manager.get_all();
9763 for key in all.keys() {
9764 if key.starts_with("disambiguation.") {
9765 let _ = self.context_manager.set(key, serde_json::Value::Null);
9766 }
9767 }
9768 }
9769
9770 async fn recheck_skill_disambiguation(
9776 &self,
9777 skill_id: &str,
9778 enriched_input: &str,
9779 expected_disambiguation_epoch: u64,
9780 expected_state_generation: Option<u64>,
9781 ) -> Result<AgentResponse> {
9782 let skill = self
9783 .skill_router
9784 .as_ref()
9785 .and_then(|r| r.get_skill(skill_id).cloned());
9786
9787 if let Some(ref skill) = skill
9789 && let Some(ref skill_disambig) = skill.disambiguation
9790 && skill_disambig.enabled.unwrap_or(false)
9791 && let Some(ref disambiguator) = self.disambiguation_manager
9792 {
9793 let context = self.build_disambiguation_context().await?;
9794 let state_override = self
9795 .state_machine
9796 .as_ref()
9797 .and_then(|sm| sm.current_definition())
9798 .and_then(|def| def.disambiguation.clone());
9799
9800 let disambiguation_result = self
9801 .observe_purpose(
9802 ObservationPurpose::DisambiguationDetection,
9803 disambiguator.process_input_with_override(
9804 enriched_input,
9805 &context,
9806 state_override.as_ref(),
9807 Some(skill_disambig),
9808 ),
9809 )
9810 .await?;
9811 let current_state_generation = self
9812 .state_machine
9813 .as_ref()
9814 .map(|state_machine| state_machine.generation());
9815 if current_state_generation != expected_state_generation
9816 || self.disambiguation_epoch.load(Ordering::SeqCst) != expected_disambiguation_epoch
9817 {
9818 disambiguator.clear_pending().await;
9819 *self.pending_skill_id.write() = None;
9820 return Err(AgentError::Other(
9821 "State or reset ownership changed during skill disambiguation recheck"
9822 .to_string(),
9823 ));
9824 }
9825 match disambiguation_result {
9826 DisambiguationResult::Clear => {
9827 debug!(skill_id = %skill_id, "Skill re-check: all fields present");
9828 }
9829 DisambiguationResult::NeedsClarification {
9830 question,
9831 detection,
9832 } => {
9833 let admission = self
9834 .admit_disambiguation_redispatch(
9835 expected_disambiguation_epoch,
9836 expected_state_generation,
9837 )
9838 .await?;
9839 let awaiting_confirmation = disambiguator.has_pending_confirmation().await;
9840 info!(
9841 skill_id = %skill_id,
9842 ambiguity_type = ?detection.ambiguity_type,
9843 what_is_unclear = ?detection.what_is_unclear,
9844 "Skill re-check: still missing fields, asking again"
9845 );
9846 self.memory
9850 .add_message(ChatMessage::user(enriched_input))
9851 .await?;
9852 self.memory
9853 .add_message(ChatMessage::assistant(&question.question))
9854 .await?;
9855
9856 let response = AgentResponse::new(&question.question).with_metadata(
9857 "disambiguation",
9858 serde_json::json!({
9859 "status": if awaiting_confirmation { "awaiting_confirmation" } else { "awaiting_clarification" },
9860 "skill_id": skill_id,
9861 "options": question.options,
9862 "clarifying": question.clarifying,
9863 "detection": {
9864 "type": detection.ambiguity_type,
9865 "confidence": detection.confidence,
9866 "what_is_unclear": detection.what_is_unclear,
9867 }
9868 }),
9869 );
9870 drop(admission);
9871 self.finish_turn_if_root(&response).await?;
9872 return Ok(response);
9873 }
9874 DisambiguationResult::Clarified {
9875 enriched_input: re_enriched,
9876 ..
9877 } => {
9878 debug!(skill_id = %skill_id, "Skill re-check: clarified immediately, executing");
9879 let admission = self
9880 .admit_disambiguation_redispatch(
9881 expected_disambiguation_epoch,
9882 expected_state_generation,
9883 )
9884 .await?;
9885 *self.pending_skill_id.write() = None;
9886 drop(admission);
9887 let skill_response = self.execute_skill_by_id(skill_id, &re_enriched).await?;
9888 self.memory
9889 .add_message(ChatMessage::user(&re_enriched))
9890 .await?;
9891 return self
9892 .handle_skill_response(
9893 &re_enriched,
9894 skill_id,
9895 skill_response,
9896 &HashMap::new(),
9897 )
9898 .await;
9899 }
9900 DisambiguationResult::ProceedWithBestGuess {
9901 enriched_input: re_enriched,
9902 } => {
9903 debug!(skill_id = %skill_id, "Skill re-check: proceeding with best guess");
9904 let admission = self
9905 .admit_disambiguation_redispatch(
9906 expected_disambiguation_epoch,
9907 expected_state_generation,
9908 )
9909 .await?;
9910 *self.pending_skill_id.write() = None;
9911 drop(admission);
9912 let skill_response = self.execute_skill_by_id(skill_id, &re_enriched).await?;
9913 self.memory
9914 .add_message(ChatMessage::user(&re_enriched))
9915 .await?;
9916 return self
9917 .handle_skill_response(
9918 &re_enriched,
9919 skill_id,
9920 skill_response,
9921 &HashMap::new(),
9922 )
9923 .await;
9924 }
9925 DisambiguationResult::GiveUp { reason } => {
9926 *self.pending_skill_id.write() = None;
9927 let apology = self
9928 .generate_localized_apology(
9929 "Generate a brief, polite apology saying you couldn't understand the request. Be concise.",
9930 &reason,
9931 )
9932 .await
9933 .unwrap_or_else(|_| {
9934 format!("I'm sorry, I couldn't understand your request: {}", reason)
9935 });
9936 let response = AgentResponse::new(&apology);
9937 self.finish_turn_if_root(&response).await?;
9938 return Ok(response);
9939 }
9940 DisambiguationResult::Escalate { reason } => {
9941 *self.pending_skill_id.write() = None;
9942 let apology = self
9943 .generate_localized_apology(
9944 "Explain briefly that you're transferring the user to a human agent for help.",
9945 &reason,
9946 )
9947 .await
9948 .unwrap_or_else(|_| {
9949 format!("I need human assistance to help with your request: {}", reason)
9950 });
9951 let response = AgentResponse::new(&apology);
9952 self.finish_turn_if_root(&response).await?;
9953 return Ok(response);
9954 }
9955 DisambiguationResult::Abandoned { new_input } => {
9956 *self.pending_skill_id.write() = None;
9959 debug!(skill_id = %skill_id, "Skill re-check: abandoned by user");
9960 if let Some(fresh) = new_input {
9961 return self.run_loop_internal(&fresh).await;
9962 }
9963 let ack = self
9964 .generate_localized_apology(
9965 "The user changed their mind about their previous request. \
9966 Generate a brief, friendly acknowledgment (e.g. 'OK, no problem. What else can I help with?'). \
9967 Do NOT apologize excessively. Be concise.",
9968 "User abandoned clarification",
9969 )
9970 .await
9971 .unwrap_or_else(|_| {
9972 "OK, no problem. What else can I help with?".to_string()
9973 });
9974 self.memory
9975 .add_message(ChatMessage::assistant(&ack))
9976 .await?;
9977 let response = AgentResponse::new(&ack);
9978 self.finish_turn_if_root(&response).await?;
9979 return Ok(response);
9980 }
9981 }
9982 }
9983
9984 let admission = self
9986 .admit_disambiguation_redispatch(
9987 expected_disambiguation_epoch,
9988 expected_state_generation,
9989 )
9990 .await?;
9991 *self.pending_skill_id.write() = None;
9992 drop(admission);
9993 let skill_response = self.execute_skill_by_id(skill_id, enriched_input).await?;
9994 self.memory
9995 .add_message(ChatMessage::user(enriched_input))
9996 .await?;
9997 self.handle_skill_response(enriched_input, skill_id, skill_response, &HashMap::new())
9998 .await
9999 }
10000
10001 async fn handle_skill_response(
10004 &self,
10005 processed_input: &str,
10006 skill_id: &str,
10007 skill_response: String,
10008 input_context: &HashMap<String, Value>,
10009 ) -> Result<AgentResponse> {
10010 let output_data = self.process_output(&skill_response, input_context).await?;
10011 let final_response = output_data.content;
10012
10013 self.memory
10014 .add_message(ChatMessage::assistant(&final_response))
10015 .await?;
10016
10017 self.check_memory_compression().await?;
10018
10019 self.increment_turn();
10020 self.evaluate_transitions(processed_input, &final_response)
10021 .await?;
10022
10023 let response = AgentResponse::new(final_response)
10024 .with_metadata("skill_id", serde_json::json!(skill_id));
10025 self.finish_turn_if_root(&response).await?;
10026 Ok(response)
10027 }
10028
10029 async fn handle_plan_and_execute(
10032 &self,
10033 processed_input: &str,
10034 input_context: &HashMap<String, Value>,
10035 auto_detected: bool,
10036 ) -> Result<AgentResponse> {
10037 let effective = self.get_effective_reasoning_config();
10038 let plan_reflection = effective
10039 .get_planning()
10040 .map(|c| c.reflection.clone())
10041 .unwrap_or_default();
10042
10043 let max_attempts = if plan_reflection.enabled {
10044 1 + plan_reflection.max_replans
10045 } else {
10046 1
10047 };
10048
10049 let mut plan = self.generate_plan(processed_input).await?;
10050 info!(
10051 plan_id = %plan.id,
10052 steps = plan.steps.len(),
10053 "Plan generated"
10054 );
10055
10056 let mut plan_result = String::new();
10057
10058 for attempt in 0..max_attempts {
10059 *self.current_plan.write() = Some(plan.clone());
10060 plan_result = self.execute_plan(&mut plan).await?;
10061
10062 info!(
10063 plan_status = ?plan.status,
10064 completed_steps = plan.completed_steps().count(),
10065 attempt = attempt + 1,
10066 "Plan execution completed"
10067 );
10068
10069 if !plan_reflection.enabled {
10070 break;
10071 }
10072
10073 let has_failures = plan.steps.iter().any(|s| s.status.is_failed());
10074 if !has_failures {
10075 break;
10076 }
10077
10078 if attempt + 1 >= max_attempts {
10079 break;
10080 }
10081
10082 match plan_reflection.on_step_failure {
10083 StepFailureAction::Replan => {
10084 info!(attempt = attempt + 1, "Plan had failures, replanning");
10085 plan = self.generate_plan(processed_input).await?;
10086 }
10087 StepFailureAction::Abort => {
10088 warn!("Plan step failed, aborting");
10089 break;
10090 }
10091 StepFailureAction::Skip | StepFailureAction::Continue => {
10092 break;
10093 }
10094 }
10095 }
10096
10097 *self.current_plan.write() = Some(plan);
10098
10099 let output_data = self.process_output(&plan_result, input_context).await?;
10100 let final_content = output_data.content;
10101
10102 self.memory
10103 .add_message(ChatMessage::assistant(&final_content))
10104 .await?;
10105
10106 self.check_memory_compression().await?;
10107 self.increment_turn();
10108 self.evaluate_transitions(processed_input, &final_content)
10109 .await?;
10110
10111 let reasoning_metadata =
10112 ReasoningMetadata::new(ReasoningMode::PlanAndExecute).with_auto_detected(auto_detected);
10113
10114 let response = AgentResponse::new(&final_content).with_metadata(
10115 "reasoning",
10116 serde_json::to_value(&reasoning_metadata).unwrap_or_default(),
10117 );
10118
10119 self.finish_turn_if_root(&response).await?;
10120 Ok(response)
10121 }
10122
10123 fn inject_reasoning_prompt(
10125 &self,
10126 messages: &mut [ChatMessage],
10127 reasoning_mode: &ReasoningMode,
10128 is_first_iteration: bool,
10129 ) {
10130 if !is_first_iteration {
10131 return;
10132 }
10133 match reasoning_mode {
10134 ReasoningMode::CoT => {
10135 if let Some(msg) = messages.first_mut()
10136 && matches!(msg.role, ai_agents_core::Role::System)
10137 {
10138 msg.content = self.build_cot_system_prompt(&msg.content);
10139 debug!("Applied Chain-of-Thought system prompt");
10140 }
10141 }
10142 ReasoningMode::React => {
10143 if let Some(msg) = messages.first_mut()
10144 && matches!(msg.role, ai_agents_core::Role::System)
10145 {
10146 msg.content = self.build_react_system_prompt(&msg.content);
10147 debug!("Applied ReAct system prompt");
10148 }
10149 }
10150 _ => {}
10151 }
10152 }
10153
10154 async fn generate_main_response_draft(
10159 &self,
10160 processed_input: &str,
10161 reasoning_mode: &ReasoningMode,
10162 ) -> Result<MainResponseDraft> {
10163 let llm = self.get_state_llm()?;
10164 let protocol = self.main_tool_protocol(llm.as_ref(), true).await?;
10165 let mut messages = self
10166 .build_messages_internal(false, Some(processed_input), protocol.choice.is_none())
10167 .await?;
10168 self.inject_reasoning_prompt(&mut messages, reasoning_mode, true);
10169 let response = self
10170 .complete_main_llm_with_recovery(llm, &messages, &protocol)
10171 .await?;
10172 let content = response.content.trim().to_string();
10173 let (thinking, answer) = self.extract_thinking(&content);
10174 if let Some(calls) = self.parse_main_tool_calls(&content, &protocol)? {
10175 return Ok(MainResponseDraft::ToolCalls {
10176 raw_content: content,
10177 calls,
10178 thinking,
10179 });
10180 }
10181 Ok(MainResponseDraft::Text {
10182 raw_content: answer,
10183 thinking,
10184 })
10185 }
10186
10187 async fn commit_main_response_draft(
10192 &self,
10193 processed_input: &str,
10194 input_context: &HashMap<String, Value>,
10195 draft: MainResponseDraft,
10196 reasoning_mode: ReasoningMode,
10197 auto_detected: bool,
10198 ) -> Result<AgentResponse> {
10199 self.commit_root_user_message(processed_input).await?;
10200 match draft {
10201 MainResponseDraft::Text {
10202 raw_content,
10203 thinking,
10204 } => {
10205 self.finish_text_response_from_model(CommittedTextResponse {
10206 processed_input,
10207 input_context,
10208 answer: raw_content,
10209 reasoning_mode,
10210 auto_detected,
10211 iterations: 1,
10212 thinking_content: thinking,
10213 all_tool_calls: Vec::new(),
10214 })
10215 .await
10216 }
10217 MainResponseDraft::ToolCalls {
10218 raw_content,
10219 calls,
10220 thinking: _,
10221 } => {
10222 let mut all_tool_calls = Vec::new();
10223 match self
10224 .handle_tool_calls(
10225 processed_input,
10226 &raw_content,
10227 calls,
10228 &mut all_tool_calls,
10229 None,
10230 )
10231 .await?
10232 {
10233 ToolCallOutcome::Rejected(response) => {
10234 self.finish_turn_if_root(&response).await?;
10235 Ok(response)
10236 }
10237 ToolCallOutcome::Continue | ToolCallOutcome::TransitionFired => {
10238 self.continue_after_committed_tool_draft(processed_input)
10239 .await
10240 }
10241 }
10242 }
10243 }
10244 }
10245
10246 async fn continue_after_committed_tool_draft(
10251 &self,
10252 processed_input: &str,
10253 ) -> Result<AgentResponse> {
10254 *self.redispatch_depth.write() += 1;
10255 if let Some(context) = self.active_turn_context.write().as_mut() {
10256 context.enter_redispatch();
10257 }
10258 let result = Box::pin(self.run_loop_internal(processed_input)).await;
10259 *self.redispatch_depth.write() -= 1;
10260 if let Some(context) = self.active_turn_context.write().as_mut() {
10261 context.exit_redispatch();
10262 }
10263 let response = result?;
10264 self.finish_turn_if_root(&response).await?;
10265 Ok(response)
10266 }
10267
10268 async fn finish_text_response_from_model(
10273 &self,
10274 response: CommittedTextResponse<'_>,
10275 ) -> Result<AgentResponse> {
10276 let CommittedTextResponse {
10277 processed_input,
10278 input_context,
10279 answer,
10280 reasoning_mode,
10281 auto_detected,
10282 iterations,
10283 thinking_content,
10284 all_tool_calls,
10285 } = response;
10286 let output_data = self.process_output(&answer, input_context).await?;
10287 let mut final_content = if output_data.metadata.rejected {
10288 output_data
10289 .metadata
10290 .rejection_reason
10291 .unwrap_or_else(|| answer.to_string())
10292 } else {
10293 output_data.content
10294 };
10295 let llm = self.get_state_llm()?;
10296 let reflection_metadata;
10297 (final_content, reflection_metadata) = self
10298 .run_reflection(&*llm, processed_input, final_content)
10299 .await?;
10300 final_content =
10301 self.format_response_with_thinking(thinking_content.as_deref(), &final_content);
10302 let final_content = {
10303 let result = self
10304 .post_loop_processing(processed_input, final_content)
10305 .await?;
10306 self.apply_post_loop_result(processed_input, result)
10307 .await?
10308 .content
10309 };
10310 let response = self.build_agent_response(AgentResponseParts {
10311 content: final_content,
10312 all_tool_calls,
10313 reasoning_mode,
10314 auto_detected,
10315 iterations,
10316 thinking: thinking_content,
10317 reflection_metadata,
10318 });
10319 self.finish_turn_if_root(&response).await?;
10320 Ok(response)
10321 }
10322
10323 async fn run_committed_response_loop_with_reasoning(
10328 &self,
10329 processed_input: &str,
10330 input_context: &HashMap<String, Value>,
10331 reasoning_mode: ReasoningMode,
10332 auto_detected: bool,
10333 ) -> Result<AgentResponse> {
10334 self.commit_root_user_message(processed_input).await?;
10335 let llm = self.get_state_llm()?;
10336 let mut iterations = 0u32;
10337 let mut all_tool_calls = Vec::new();
10338 let mut thinking_content = None;
10339 loop {
10340 let effective_max = if reasoning_mode != ReasoningMode::None {
10341 let rc = self.get_effective_reasoning_config();
10342 self.max_iterations.min(rc.max_iterations)
10343 } else {
10344 self.max_iterations
10345 };
10346 if iterations >= effective_max {
10347 return Err(AgentError::Other(format!(
10348 "Max iterations ({}) exceeded",
10349 effective_max
10350 )));
10351 }
10352 iterations += 1;
10353 *self.iteration_count.write() = iterations;
10354 let protocol = self.main_tool_protocol(llm.as_ref(), false).await?;
10355 let mut messages = self
10356 .build_messages_internal(true, None, protocol.choice.is_none())
10357 .await?;
10358 self.inject_reasoning_prompt(&mut messages, &reasoning_mode, iterations == 1);
10359 self.hooks.on_llm_start(&messages).await;
10360 let llm_start = Instant::now();
10361 let response = self
10362 .complete_main_llm_with_recovery(Arc::clone(&llm), &messages, &protocol)
10363 .await?;
10364 let llm_duration_ms = llm_start.elapsed().as_millis() as u64;
10365 self.hooks.on_llm_complete(&response, llm_duration_ms).await;
10366 let content = response.content.trim();
10367 if let Some(tool_calls) = self.parse_main_tool_calls(content, &protocol)? {
10368 match self
10369 .handle_tool_calls(
10370 processed_input,
10371 content,
10372 tool_calls,
10373 &mut all_tool_calls,
10374 None,
10375 )
10376 .await?
10377 {
10378 ToolCallOutcome::Continue | ToolCallOutcome::TransitionFired => continue,
10379 ToolCallOutcome::Rejected(resp) => {
10380 self.finish_turn_if_root(&resp).await?;
10381 return Ok(resp);
10382 }
10383 }
10384 }
10385 let (extracted_thinking, answer) = self.extract_thinking(content);
10386 if extracted_thinking.is_some() {
10387 thinking_content = extracted_thinking;
10388 }
10389 return self
10390 .finish_text_response_from_model(CommittedTextResponse {
10391 processed_input,
10392 input_context,
10393 answer,
10394 reasoning_mode,
10395 auto_detected,
10396 iterations,
10397 thinking_content,
10398 all_tool_calls,
10399 })
10400 .await;
10401 }
10402 }
10403
10404 async fn handle_tool_calls(
10410 &self,
10411 processed_input: &str,
10412 content: &str,
10413 tool_calls: Vec<ToolCall>,
10414 all_tool_calls: &mut Vec<ToolCall>,
10415 mut events: Option<&mut Vec<StreamChunk>>,
10416 ) -> Result<ToolCallOutcome> {
10417 let include_tool_events = self.streaming.include_tool_events;
10418 let transition_content = native_readable_projection(content)
10422 .map_err(|error| AgentError::LLM(error.to_string()))?;
10423 let transition_fired = self
10424 .evaluate_transitions(processed_input, &transition_content)
10425 .await?;
10426 if transition_fired {
10427 self.memory
10428 .add_message(ChatMessage::assistant(
10429 "(Transitioned to new state — tool call handled by workflow)",
10430 ))
10431 .await?;
10432 if let Some(events) = events.as_deref_mut()
10433 && self.streaming.include_state_events
10434 && let Some(state) = self.current_state()
10435 {
10436 events.push(StreamChunk::state_transition(None, state));
10437 }
10438 return Ok(ToolCallOutcome::TransitionFired);
10439 }
10440
10441 self.memory
10443 .add_message(ChatMessage::assistant(content))
10444 .await?;
10445 self.remember_committed_native_exchange(content).await?;
10446 let native_tool_call = Self::is_native_tool_call_content(content)?;
10447
10448 if let Some(events) = events.as_deref_mut()
10449 && include_tool_events
10450 {
10451 for tool_call in &tool_calls {
10452 events.push(StreamChunk::tool_start(&tool_call.id, &tool_call.name));
10453 }
10454 }
10455 let results = self.execute_tools_parallel(&tool_calls).await;
10456 let mut rejection = None;
10457
10458 for ((_id, result), tool_call) in results.into_iter().zip(tool_calls.iter()) {
10459 match result {
10460 Ok(output) => {
10461 if let Some(events) = events.as_deref_mut()
10462 && include_tool_events
10463 {
10464 events.push(StreamChunk::tool_result(
10465 &tool_call.id,
10466 &tool_call.name,
10467 &output,
10468 true,
10469 ));
10470 }
10471 self.memory
10472 .add_message(Self::tool_result_message(
10473 tool_call,
10474 &output,
10475 native_tool_call,
10476 )?)
10477 .await?;
10478 }
10479 Err(e) => {
10480 if matches!(e, AgentError::HITLRejected(_)) {
10481 if !native_tool_call {
10482 self.memory
10483 .add_message(ChatMessage::assistant(format!(
10484 "The operation was rejected by the approver: {e}"
10485 )))
10486 .await?;
10487 return Ok(ToolCallOutcome::Rejected(AgentResponse {
10488 content: format!("Operation cancelled: {e}"),
10489 metadata: None,
10490 tool_calls: Some(all_tool_calls.clone()),
10491 }));
10492 }
10493 if rejection.is_none() {
10494 rejection = Some(e.to_string());
10495 }
10496 }
10497 if let Some(events) = events.as_deref_mut()
10498 && include_tool_events
10499 {
10500 events.push(StreamChunk::tool_result(
10501 &tool_call.id,
10502 &tool_call.name,
10503 e.to_string(),
10504 false,
10505 ));
10506 }
10507 self.memory
10508 .add_message(Self::tool_result_message(
10509 tool_call,
10510 &format!("Error: {}", e),
10511 native_tool_call,
10512 )?)
10513 .await?;
10514 }
10515 }
10516 all_tool_calls.push(tool_call.clone());
10517 if let Some(events) = events.as_deref_mut()
10518 && include_tool_events
10519 {
10520 events.push(StreamChunk::tool_end(&tool_call.id));
10521 }
10522 }
10523 if let Some(rejection) = rejection {
10524 self.memory
10525 .add_message(ChatMessage::assistant(format!(
10526 "The operation was rejected by the approver: {rejection}"
10527 )))
10528 .await?;
10529 return Ok(ToolCallOutcome::Rejected(AgentResponse {
10530 content: format!("Operation cancelled: {rejection}"),
10531 metadata: None,
10532 tool_calls: Some(all_tool_calls.clone()),
10533 }));
10534 }
10535 Ok(ToolCallOutcome::Continue)
10536 }
10537
10538 async fn run_reflection(
10540 &self,
10541 llm: &dyn LLMProvider,
10542 processed_input: &str,
10543 mut content: String,
10544 ) -> Result<(String, Option<ReflectionMetadata>)> {
10545 let config = self.get_effective_reflection_config();
10546 let should_reflect = self
10547 .should_reflect_with_config(processed_input, &content, &config)
10548 .await?;
10549 if !should_reflect {
10550 return Ok((content, None));
10551 }
10552
10553 info!("Starting response reflection evaluation");
10554 let mut attempts = 0u32;
10555 let max_retries = config.max_retries;
10556 let mut history: Vec<ReflectionAttempt> = Vec::new();
10557
10558 loop {
10559 let evaluation = self
10560 .evaluate_response_with_config(processed_input, &content, &config)
10561 .await?;
10562
10563 if evaluation.passed || attempts >= max_retries {
10564 info!(
10565 passed = evaluation.passed,
10566 confidence = evaluation.confidence,
10567 attempts = attempts + 1,
10568 "Reflection evaluation complete"
10569 );
10570 let reflection_metadata = Some(
10571 ReflectionMetadata::new(evaluation)
10572 .with_attempts(attempts + 1)
10573 .with_history(history),
10574 );
10575 return Ok((content, reflection_metadata));
10576 }
10577
10578 debug!(
10579 attempt = attempts + 1,
10580 failed_criteria = evaluation.failed_criteria().count(),
10581 "Response did not meet criteria, retrying"
10582 );
10583
10584 history.push(
10585 ReflectionAttempt::new(&content, evaluation.clone())
10586 .with_feedback("Response did not meet quality criteria"),
10587 );
10588
10589 let feedback: Vec<String> = evaluation
10590 .failed_criteria()
10591 .map(|c| format!("- {}", c.criterion))
10592 .collect();
10593
10594 let retry_prompt = format!(
10595 "Your previous response did not meet these criteria:\n{}\n\nPlease provide an improved response.",
10596 feedback.join("\n")
10597 );
10598
10599 self.memory
10600 .add_message(ChatMessage::user(&retry_prompt))
10601 .await?;
10602
10603 let retry_messages = self.build_messages().await?;
10604 let retry_response = self
10605 .observe_purpose(
10606 ObservationPurpose::ReflectionEvaluation,
10607 llm.complete(&retry_messages, None),
10608 )
10609 .await
10610 .map_err(|e| AgentError::LLM(e.to_string()))?;
10611
10612 content = retry_response.content.trim().to_string();
10613 attempts += 1;
10614 }
10615 }
10616
10617 async fn post_loop_processing(
10620 &self,
10621 processed_input: &str,
10622 content: String,
10623 ) -> Result<PostLoopResult> {
10624 self.increment_turn();
10629
10630 self.run_context_extractors(processed_input).await?;
10632
10633 let transitioned = self.evaluate_transitions(processed_input, &content).await?;
10634
10635 if !transitioned {
10636 self.memory
10637 .add_message(ChatMessage::assistant(&content))
10638 .await?;
10639 self.check_memory_compression().await?;
10640 return Ok(PostLoopResult::NoTransition(content));
10641 }
10642
10643 if !self.should_regenerate_after_transition() {
10645 self.memory
10646 .add_message(ChatMessage::assistant(&content))
10647 .await?;
10648 self.check_memory_compression().await?;
10649 return Ok(PostLoopResult::Transitioned {
10650 content,
10651 regenerated: false,
10652 });
10653 }
10654
10655 if self.needs_redispatch_for_new_state() {
10659 info!("Post-transition NeedsRedispatch: new state requires full dispatch");
10660 return Ok(PostLoopResult::NeedsRedispatch);
10663 }
10664
10665 self.memory
10668 .add_message(ChatMessage::assistant(&content))
10669 .await?;
10670 self.check_memory_compression().await?;
10671
10672 let new_llm = self.get_state_llm()?;
10678 let mut final_content;
10679
10680 for post_iter in 0..self.max_iterations {
10681 let protocol = self.main_tool_protocol(new_llm.as_ref(), false).await?;
10682 let new_messages = self
10683 .build_messages_internal(true, None, protocol.choice.is_none())
10684 .await?;
10685 if post_iter == 0
10686 && let Some(system_msg) = new_messages.first()
10687 && system_msg.role == ai_agents_core::Role::System
10688 {
10689 debug!(
10690 prompt_preview =
10691 &system_msg.content[system_msg.content.len().saturating_sub(200)..],
10692 "Post-transition system prompt (last 200 chars)"
10693 );
10694 }
10695
10696 let new_response = self
10697 .complete_main_llm_with_recovery(Arc::clone(&new_llm), &new_messages, &protocol)
10698 .await?;
10699 final_content = new_response.content.trim().to_string();
10700
10701 if let Some(tool_calls) = self.parse_main_tool_calls(&final_content, &protocol)? {
10704 let native_tool_call = Self::is_native_tool_call_content(&final_content)?;
10705 debug!(
10706 post_iter = post_iter,
10707 tools = tool_calls.len(),
10708 "Post-transition tool call detected, executing"
10709 );
10710
10711 self.memory
10712 .add_message(ChatMessage::assistant(&final_content))
10713 .await?;
10714 self.remember_committed_native_exchange(&final_content)
10715 .await?;
10716
10717 let results = self.execute_tools_parallel(&tool_calls).await;
10718 let mut rejection = None;
10719 for ((_id, result), tool_call) in results.into_iter().zip(tool_calls.iter()) {
10720 match result {
10721 Ok(output) => {
10722 self.memory
10723 .add_message(Self::tool_result_message(
10724 tool_call,
10725 &output,
10726 native_tool_call,
10727 )?)
10728 .await?;
10729 }
10730 Err(e) => {
10731 if native_tool_call
10732 && rejection.is_none()
10733 && matches!(e, AgentError::HITLRejected(_))
10734 {
10735 rejection = Some(e.to_string());
10736 }
10737 self.memory
10738 .add_message(Self::tool_result_message(
10739 tool_call,
10740 &format!("Error: {}", e),
10741 native_tool_call,
10742 )?)
10743 .await?;
10744 }
10745 }
10746 }
10747 if let Some(rejection) = rejection {
10748 self.memory
10749 .add_message(ChatMessage::assistant(format!(
10750 "The operation was rejected by the approver: {rejection}"
10751 )))
10752 .await?;
10753 return Err(AgentError::HITLRejected(rejection));
10754 }
10755 continue;
10757 }
10758
10759 self.memory
10761 .add_message(ChatMessage::assistant(&final_content))
10762 .await?;
10763 return Ok(PostLoopResult::Transitioned {
10764 content: final_content,
10765 regenerated: true,
10766 });
10767 }
10768
10769 final_content = "Post-transition processing completed.".to_string();
10771 self.memory
10772 .add_message(ChatMessage::assistant(&final_content))
10773 .await?;
10774
10775 Ok(PostLoopResult::Transitioned {
10776 content: final_content,
10777 regenerated: true,
10778 })
10779 }
10780
10781 fn should_regenerate_after_transition(&self) -> bool {
10784 if let Some(ref sm) = self.state_machine {
10785 if !sm.config().regenerate_on_transition {
10787 return false;
10788 }
10789 if let Some(def) = sm.current_definition()
10791 && let Some(regen) = def.regenerate_on_enter
10792 {
10793 return regen;
10794 }
10795 }
10796 true
10797 }
10798
10799 fn needs_redispatch_for_new_state(&self) -> bool {
10802 if let Some(ref sm) = self.state_machine
10803 && let Some(def) = sm.current_definition()
10804 {
10805 if def.concurrent.is_some()
10806 || def.group_chat.is_some()
10807 || def.pipeline.is_some()
10808 || def.handoff.is_some()
10809 || def.delegate.is_some()
10810 {
10811 return true;
10812 }
10813 let effective = self.get_effective_reasoning_config();
10815 if !matches!(effective.mode, ReasoningMode::None) {
10816 return true;
10817 }
10818 }
10819 false
10820 }
10821
10822 async fn apply_post_loop_result(
10828 &self,
10829 processed_input: &str,
10830 result: PostLoopResult,
10831 ) -> Result<AppliedPostLoop> {
10832 match result {
10833 PostLoopResult::NoTransition(content) => Ok(AppliedPostLoop {
10834 content,
10835 transitioned: false,
10836 regenerated: false,
10837 }),
10838 PostLoopResult::Transitioned {
10839 content,
10840 regenerated,
10841 } => Ok(AppliedPostLoop {
10842 content,
10843 transitioned: true,
10844 regenerated,
10845 }),
10846 PostLoopResult::NeedsRedispatch => {
10847 const MAX_REDISPATCH_DEPTH: u32 = 3;
10848 let current_depth = *self.redispatch_depth.read();
10849 if current_depth >= MAX_REDISPATCH_DEPTH {
10850 warn!(
10851 depth = current_depth,
10852 "Post-transition re-dispatch depth limit reached, returning empty response"
10853 );
10854 let content = String::new();
10855 self.memory
10856 .add_message(ChatMessage::assistant(&content))
10857 .await?;
10858 return Ok(AppliedPostLoop {
10860 content,
10861 transitioned: true,
10862 regenerated: false,
10863 });
10864 }
10865 *self.redispatch_depth.write() += 1;
10866 if let Some(context) = self.active_turn_context.write().as_mut() {
10867 context.enter_redispatch();
10868 }
10869 info!(
10870 depth = current_depth + 1,
10871 "Re-dispatching for new state after transition"
10872 );
10873 let resp = Box::pin(self.run_loop_internal(processed_input)).await;
10874 *self.redispatch_depth.write() -= 1;
10875 if let Some(context) = self.active_turn_context.write().as_mut() {
10876 context.exit_redispatch();
10877 }
10878 resp.map(|r| AppliedPostLoop {
10879 content: r.content,
10880 transitioned: true,
10881 regenerated: true,
10882 })
10883 }
10884 }
10885 }
10886
10887 fn build_agent_response(&self, parts: AgentResponseParts) -> AgentResponse {
10889 let AgentResponseParts {
10890 content,
10891 all_tool_calls,
10892 reasoning_mode,
10893 auto_detected,
10894 iterations,
10895 thinking,
10896 reflection_metadata,
10897 } = parts;
10898 let reasoning_metadata = ReasoningMetadata::new(reasoning_mode.clone())
10899 .with_thinking(thinking.clone().unwrap_or_default())
10900 .with_iterations(iterations)
10901 .with_auto_detected(auto_detected);
10902
10903 let mut response = AgentResponse::new(&content);
10904 if !all_tool_calls.is_empty() {
10905 response = response.with_tool_calls(all_tool_calls);
10906 }
10907
10908 if let Some(state) = self.current_state() {
10909 response = response.with_metadata("current_state", serde_json::json!(state));
10910 }
10911
10912 response = response.with_metadata(
10913 "reasoning",
10914 serde_json::to_value(&reasoning_metadata).unwrap_or_default(),
10915 );
10916
10917 if let Some(ref refl_meta) = reflection_metadata {
10918 response = response.with_metadata(
10919 "reflection",
10920 serde_json::to_value(refl_meta).unwrap_or_default(),
10921 );
10922 }
10923
10924 response
10925 }
10926
10927 async fn handle_delegated_state(
10930 &self,
10931 input: &str,
10932 delegate_id: &str,
10933 state_def: &ai_agents_state::StateDefinition,
10934 ) -> Result<AgentResponse> {
10935 use std::time::Instant;
10936
10937 let registry = self.spawner_registry.as_ref().ok_or_else(|| {
10938 AgentError::Config(format!(
10939 "State delegates to '{}' but no agent registry is configured. \
10940 Add a spawner section with auto_spawn to your YAML.",
10941 delegate_id
10942 ))
10943 })?;
10944
10945 let state_name = self
10946 .state_machine
10947 .as_ref()
10948 .map(|sm| sm.current())
10949 .unwrap_or_else(|| "unknown".to_string());
10950
10951 self.hooks.on_delegate_start(delegate_id, &state_name).await;
10952 let start = Instant::now();
10953
10954 let delegate = registry.get(delegate_id).ok_or_else(|| {
10955 AgentError::Other(format!(
10956 "State '{}' delegates to '{}' but no agent with that ID exists in the registry.",
10957 state_name, delegate_id
10958 ))
10959 })?;
10960
10961 let context_mode = state_def.delegate_context.clone().unwrap_or_default();
10963 let effective_input = self
10964 .observe_purpose(
10965 ObservationPurpose::OrchestrationRouting,
10966 crate::orchestration::context::prepare_delegate_input(
10967 input,
10968 &context_mode,
10969 &*self.memory,
10970 self.optional_role_llm(
10971 ai_agents_llm::LLMRole::OrchestrationSummary,
10972 None,
10973 || self.llm_registry.get("router").ok(),
10974 )?
10975 .as_deref(),
10976 ),
10977 )
10978 .await?;
10979
10980 let response = delegate
10981 .chat_with_actor_context(&effective_input, self.outbound_actor_context())
10982 .await?;
10983
10984 let duration_ms = start.elapsed().as_millis() as u64;
10985 self.hooks
10986 .on_delegate_complete(delegate_id, &state_name, duration_ms)
10987 .await;
10988
10989 let ctx_key = format!("delegation.{}.last_response", delegate_id);
10991 let _ = self.context_manager.set(
10992 &ctx_key,
10993 serde_json::Value::String(response.content.clone()),
10994 );
10995
10996 let _ = self.context_manager.set(
10998 "orchestration",
10999 serde_json::json!({
11000 "type": "delegate",
11001 "agent": delegate_id,
11002 "state": state_name,
11003 "response": response.content,
11004 "duration_ms": duration_ms,
11005 }),
11006 );
11007
11008 self.commit_root_user_message(input).await?;
11009
11010 let post_result = self
11013 .post_loop_processing(
11014 input,
11015 format!("[Delegated to {}]: {}", delegate_id, response.content),
11016 )
11017 .await?;
11018 let final_content = self
11019 .apply_post_loop_result(input, post_result)
11020 .await?
11021 .content;
11022
11023 let mut result = AgentResponse::new(final_content);
11024
11025 let metadata = serde_json::json!({
11026 "orchestration": {
11027 "type": "delegate",
11028 "agent": delegate_id,
11029 "state": state_name,
11030 "response": response.content,
11031 "duration_ms": duration_ms,
11032 }
11033 });
11034 result.metadata = Some(
11035 serde_json::from_value::<std::collections::HashMap<String, serde_json::Value>>(
11036 metadata,
11037 )
11038 .unwrap_or_default(),
11039 );
11040
11041 self.finish_turn_if_root(&result).await?;
11042 Ok(result)
11043 }
11044
11045 async fn handle_concurrent_state(
11048 &self,
11049 input: &str,
11050 config: &ai_agents_state::ConcurrentStateConfig,
11051 ) -> Result<AgentResponse> {
11052 use std::time::Instant;
11053
11054 let registry = self.spawner_registry.as_ref().ok_or_else(|| {
11055 AgentError::Config(
11056 "Concurrent state requires an agent registry. Add a spawner section.".into(),
11057 )
11058 })?;
11059
11060 let context_mode = config.context_mode.clone().unwrap_or_default();
11065 let context_input = self
11066 .observe_purpose(
11067 ObservationPurpose::OrchestrationRouting,
11068 crate::orchestration::context::prepare_delegate_input(
11069 input,
11070 &context_mode,
11071 &*self.memory,
11072 self.optional_role_llm(
11073 ai_agents_llm::LLMRole::OrchestrationSummary,
11074 None,
11075 || self.llm_registry.get("router").ok(),
11076 )?
11077 .as_deref(),
11078 ),
11079 )
11080 .await?;
11081
11082 let effective_input = if let Some(ref tmpl) = config.input {
11083 render_concurrent_template(tmpl, &context_input, &self.build_context_with_overlays())
11084 .unwrap_or_else(|_| context_input.clone())
11085 } else {
11086 context_input
11087 };
11088
11089 let start = Instant::now();
11090
11091 let providers = crate::orchestration::AggregationProviders::resolve(
11092 &self.llm_registry,
11093 config.aggregation.synthesizer_llm.as_deref(),
11094 )?;
11095
11096 let vote_parallelism = if self.runtime_config.optimization.enabled
11097 && self
11098 .runtime_config
11099 .optimization
11100 .parallel_orchestration_vote_extraction
11101 {
11102 Some(self.runtime_config.optimization.max_parallel_runtime_tasks)
11103 } else {
11104 None
11105 };
11106
11107 let result = self
11108 .observe_purpose(
11109 ObservationPurpose::OrchestrationAggregation,
11110 scope_actor_context(
11111 self.outbound_actor_context(),
11112 crate::orchestration::concurrent_with_llms(
11113 registry,
11114 &effective_input,
11115 &config.agents,
11116 &config.aggregation,
11117 providers.as_refs(),
11118 config.min_required,
11119 config.timeout_ms,
11120 config.on_partial_failure.clone(),
11121 vote_parallelism,
11122 ),
11123 ),
11124 )
11125 .await?;
11126
11127 let duration_ms = start.elapsed().as_millis() as u64;
11128 let agent_ids: Vec<String> = config.agents.iter().map(|a| a.id().to_string()).collect();
11129 let strategy = format!("{:?}", config.aggregation.strategy);
11130 self.hooks
11131 .on_concurrent_complete(&agent_ids, &strategy, duration_ms)
11132 .await;
11133
11134 let _ = self.context_manager.set(
11136 "concurrent.result",
11137 serde_json::Value::String(result.response.content.clone()),
11138 );
11139
11140 let agents_json: Vec<serde_json::Value> = result
11142 .agent_results
11143 .iter()
11144 .map(|ar| {
11145 serde_json::json!({
11146 "id": ar.agent_id,
11147 "response": ar.response.as_ref().map(|r| r.content.as_str()),
11148 "success": ar.success,
11149 "error": ar.error,
11150 "duration_ms": ar.duration_ms,
11151 })
11152 })
11153 .collect();
11154
11155 let _ = self.context_manager.set(
11157 "orchestration",
11158 serde_json::json!({
11159 "type": "concurrent",
11160 "result": result.response.content,
11161 "strategy": strategy,
11162 "agents": agents_json,
11163 "duration_ms": duration_ms,
11164 }),
11165 );
11166
11167 self.commit_root_user_message(input).await?;
11168
11169 let post_result = self
11170 .post_loop_processing(input, result.response.content.clone())
11171 .await?;
11172 let final_content = self
11173 .apply_post_loop_result(input, post_result)
11174 .await?
11175 .content;
11176
11177 let mut response = AgentResponse::new(final_content);
11178 let metadata = serde_json::json!({
11179 "orchestration": {
11180 "type": "concurrent",
11181 "result": result.response.content,
11182 "strategy": strategy,
11183 "agents": agents_json,
11184 "duration_ms": duration_ms,
11185 }
11186 });
11187 response.metadata = Some(
11188 serde_json::from_value::<std::collections::HashMap<String, serde_json::Value>>(
11189 metadata,
11190 )
11191 .unwrap_or_default(),
11192 );
11193
11194 self.finish_turn_if_root(&response).await?;
11195 Ok(response)
11196 }
11197
11198 async fn handle_group_chat_state(
11201 &self,
11202 input: &str,
11203 config: &ai_agents_state::GroupChatStateConfig,
11204 ) -> Result<AgentResponse> {
11205 use std::time::Instant;
11206
11207 let registry = self.spawner_registry.as_ref().ok_or_else(|| {
11208 AgentError::Config(
11209 "Group chat state requires an agent registry. Add a spawner section.".into(),
11210 )
11211 })?;
11212
11213 let start = Instant::now();
11214
11215 let speaker = crate::orchestration::role_provider(
11216 &self.llm_registry,
11217 ai_agents_llm::LLMRole::OrchestrationSpeaker,
11218 None,
11219 )?;
11220 let consensus = crate::orchestration::role_provider(
11221 &self.llm_registry,
11222 ai_agents_llm::LLMRole::OrchestrationConsensus,
11223 None,
11224 )?;
11225
11226 let context_mode = config.context_mode.clone().unwrap_or_default();
11228 let context_input = self
11229 .observe_purpose(
11230 ObservationPurpose::OrchestrationRouting,
11231 crate::orchestration::context::prepare_delegate_input(
11232 input,
11233 &context_mode,
11234 &*self.memory,
11235 self.optional_role_llm(
11236 ai_agents_llm::LLMRole::OrchestrationSummary,
11237 None,
11238 || self.llm_registry.get("router").ok(),
11239 )?
11240 .as_deref(),
11241 ),
11242 )
11243 .await?;
11244
11245 let effective_topic = if let Some(ref tmpl) = config.input {
11247 render_concurrent_template(tmpl, &context_input, &self.build_context_with_overlays())
11248 .unwrap_or_else(|_| context_input.clone())
11249 } else {
11250 context_input
11251 };
11252
11253 let result = self
11254 .observe_purpose(
11255 ObservationPurpose::OrchestrationConversation,
11256 scope_actor_context(
11257 self.outbound_actor_context(),
11258 crate::orchestration::group_chat_with_llms(
11259 registry,
11260 &effective_topic,
11261 config,
11262 crate::orchestration::GroupLLMs {
11263 speaker: speaker.as_deref(),
11264 consensus: consensus.as_deref(),
11265 },
11266 Some(&*self.hooks),
11267 ),
11268 ),
11269 )
11270 .await?;
11271
11272 let duration_ms = start.elapsed().as_millis() as u64;
11273
11274 let _ = self.context_manager.set(
11276 "group_chat.conclusion",
11277 serde_json::Value::String(result.response.content.clone()),
11278 );
11279
11280 let transcript_json: Vec<serde_json::Value> = result
11282 .transcript
11283 .iter()
11284 .map(|t| {
11285 serde_json::json!({
11286 "speaker": t.speaker,
11287 "round": t.round,
11288 "content": t.content,
11289 })
11290 })
11291 .collect();
11292
11293 let _ = self.context_manager.set(
11295 "orchestration",
11296 serde_json::json!({
11297 "type": "group_chat",
11298 "conclusion": result.response.content,
11299 "transcript": transcript_json,
11300 "rounds": result.rounds_completed,
11301 "termination": result.termination_reason,
11302 "duration_ms": duration_ms,
11303 }),
11304 );
11305
11306 self.commit_root_user_message(input).await?;
11307
11308 let post_result = self
11309 .post_loop_processing(input, result.response.content.clone())
11310 .await?;
11311 let final_content = self
11312 .apply_post_loop_result(input, post_result)
11313 .await?
11314 .content;
11315
11316 let mut response = AgentResponse::new(final_content);
11317 let metadata = serde_json::json!({
11318 "orchestration": {
11319 "type": "group_chat",
11320 "conclusion": result.response.content,
11321 "transcript": transcript_json,
11322 "rounds": result.rounds_completed,
11323 "termination": result.termination_reason,
11324 "duration_ms": duration_ms,
11325 }
11326 });
11327 response.metadata = Some(
11328 serde_json::from_value::<std::collections::HashMap<String, serde_json::Value>>(
11329 metadata,
11330 )
11331 .unwrap_or_default(),
11332 );
11333
11334 self.finish_turn_if_root(&response).await?;
11335 Ok(response)
11336 }
11337
11338 async fn handle_pipeline_state(
11341 &self,
11342 input: &str,
11343 config: &ai_agents_state::PipelineStateConfig,
11344 ) -> Result<AgentResponse> {
11345 use std::time::Instant;
11346
11347 let registry = self.spawner_registry.as_ref().ok_or_else(|| {
11348 AgentError::Config(
11349 "Pipeline state requires an agent registry. Add a spawner section.".into(),
11350 )
11351 })?;
11352
11353 let start = Instant::now();
11354
11355 let stages: Vec<crate::orchestration::PipelineStage> = config
11356 .stages
11357 .iter()
11358 .map(|entry| {
11359 let mut stage = crate::orchestration::PipelineStage::id(entry.id());
11360 if let Some(tmpl) = entry.input() {
11361 stage = stage.with_input(tmpl);
11362 }
11363 stage
11364 })
11365 .collect();
11366
11367 let context_mode = config.context_mode.clone().unwrap_or_default();
11369 let context_input = self
11370 .observe_purpose(
11371 ObservationPurpose::OrchestrationRouting,
11372 crate::orchestration::context::prepare_delegate_input(
11373 input,
11374 &context_mode,
11375 &*self.memory,
11376 self.optional_role_llm(
11377 ai_agents_llm::LLMRole::OrchestrationSummary,
11378 None,
11379 || self.llm_registry.get("router").ok(),
11380 )?
11381 .as_deref(),
11382 ),
11383 )
11384 .await?;
11385
11386 let context_values = self.build_context_with_overlays();
11387 let result = self
11388 .observe_purpose(
11389 ObservationPurpose::OrchestrationRouting,
11390 scope_actor_context(
11391 self.outbound_actor_context(),
11392 crate::orchestration::pipeline(
11393 registry,
11394 &context_input,
11395 &stages,
11396 config.timeout_ms,
11397 Some(&*self.hooks),
11398 Some(&context_values),
11399 ),
11400 ),
11401 )
11402 .await?;
11403
11404 let duration_ms = start.elapsed().as_millis() as u64;
11405
11406 let _ = self.context_manager.set(
11408 "pipeline.result",
11409 serde_json::Value::String(result.response.content.clone()),
11410 );
11411
11412 let stages_json: Vec<serde_json::Value> = result
11414 .stage_outputs
11415 .iter()
11416 .map(|s| {
11417 serde_json::json!({
11418 "agent_id": s.agent_id,
11419 "output": s.output,
11420 "duration_ms": s.duration_ms,
11421 "skipped": s.skipped,
11422 })
11423 })
11424 .collect();
11425
11426 let _ = self.context_manager.set(
11428 "orchestration",
11429 serde_json::json!({
11430 "type": "pipeline",
11431 "result": result.response.content,
11432 "stages": stages_json,
11433 "duration_ms": duration_ms,
11434 }),
11435 );
11436
11437 self.commit_root_user_message(input).await?;
11438
11439 let post_result = self
11440 .post_loop_processing(input, result.response.content.clone())
11441 .await?;
11442 let final_content = self
11443 .apply_post_loop_result(input, post_result)
11444 .await?
11445 .content;
11446
11447 let mut response = AgentResponse::new(final_content);
11448 let metadata = serde_json::json!({
11449 "orchestration": {
11450 "type": "pipeline",
11451 "result": result.response.content,
11452 "stages": stages_json,
11453 "duration_ms": duration_ms,
11454 }
11455 });
11456 response.metadata = Some(
11457 serde_json::from_value::<std::collections::HashMap<String, serde_json::Value>>(
11458 metadata,
11459 )
11460 .unwrap_or_default(),
11461 );
11462
11463 self.finish_turn_if_root(&response).await?;
11464 Ok(response)
11465 }
11466
11467 async fn handle_handoff_state(
11470 &self,
11471 input: &str,
11472 config: &ai_agents_state::HandoffStateConfig,
11473 ) -> Result<AgentResponse> {
11474 use std::time::Instant;
11475
11476 let registry = self.spawner_registry.as_ref().ok_or_else(|| {
11477 AgentError::Config(
11478 "Handoff state requires an agent registry. Add a spawner section.".into(),
11479 )
11480 })?;
11481
11482 let llm = self.role_llm(ai_agents_llm::LLMRole::OrchestrationHandoff, None, || {
11483 self.llm_registry
11484 .get("router")
11485 .map_err(|_| AgentError::Config("Handoff state requires a router LLM.".into()))
11486 })?;
11487
11488 let start = Instant::now();
11489
11490 let context_mode = config.context_mode.clone().unwrap_or_default();
11492 let context_input = self
11493 .observe_purpose(
11494 ObservationPurpose::OrchestrationRouting,
11495 crate::orchestration::context::prepare_delegate_input(
11496 input,
11497 &context_mode,
11498 &*self.memory,
11499 self.optional_role_llm(
11500 ai_agents_llm::LLMRole::OrchestrationSummary,
11501 None,
11502 || self.llm_registry.get("router").ok(),
11503 )?
11504 .as_deref(),
11505 ),
11506 )
11507 .await?;
11508
11509 let effective_input = if let Some(ref tmpl) = config.input {
11511 render_concurrent_template(tmpl, &context_input, &self.build_context_with_overlays())
11512 .unwrap_or_else(|_| context_input.clone())
11513 } else {
11514 context_input
11515 };
11516
11517 let result = self
11518 .observe_purpose(
11519 ObservationPurpose::OrchestrationRouting,
11520 scope_actor_context(
11521 self.outbound_actor_context(),
11522 crate::orchestration::handoff(
11523 registry,
11524 &effective_input,
11525 &config.initial_agent,
11526 &config.available_agents,
11527 config.max_handoffs,
11528 llm.as_ref(),
11529 Some(&*self.hooks),
11530 ),
11531 ),
11532 )
11533 .await?;
11534
11535 let duration_ms = start.elapsed().as_millis() as u64;
11536
11537 let _ = self.context_manager.set(
11539 "handoff.result",
11540 serde_json::Value::String(result.response.content.clone()),
11541 );
11542
11543 let chain_json: Vec<serde_json::Value> = result
11545 .handoff_chain
11546 .iter()
11547 .map(|h| {
11548 serde_json::json!({
11549 "from": h.from_agent,
11550 "to": h.to_agent,
11551 "reason": h.reason,
11552 })
11553 })
11554 .collect();
11555
11556 let _ = self.context_manager.set(
11558 "orchestration",
11559 serde_json::json!({
11560 "type": "handoff",
11561 "result": result.response.content,
11562 "final_agent": result.final_agent,
11563 "handoff_chain": chain_json,
11564 "duration_ms": duration_ms,
11565 }),
11566 );
11567
11568 self.commit_root_user_message(input).await?;
11569
11570 let post_result = self
11571 .post_loop_processing(input, result.response.content.clone())
11572 .await?;
11573 let final_content = self
11574 .apply_post_loop_result(input, post_result)
11575 .await?
11576 .content;
11577
11578 let mut response = AgentResponse::new(final_content);
11579 let metadata = serde_json::json!({
11580 "orchestration": {
11581 "type": "handoff",
11582 "result": result.response.content,
11583 "final_agent": result.final_agent,
11584 "handoff_chain": chain_json,
11585 "duration_ms": duration_ms,
11586 }
11587 });
11588 response.metadata = Some(
11589 serde_json::from_value::<std::collections::HashMap<String, serde_json::Value>>(
11590 metadata,
11591 )
11592 .unwrap_or_default(),
11593 );
11594
11595 self.finish_turn_if_root(&response).await?;
11596 Ok(response)
11597 }
11598
11599 async fn run_loop_internal(&self, input: &str) -> Result<AgentResponse> {
11601 self.begin_root_turn();
11602 self.pre_turn_session_lifecycle().await;
11604
11605 let input_data = self.process_input(input).await?;
11606 self.update_active_turn_context(&input_data.content, input_data.context.clone());
11607
11608 for (key, value) in &input_data.context {
11611 let _ = self.context_manager.set(key, value.clone());
11612 }
11613
11614 if input_data.metadata.rejected {
11615 let reason = input_data
11616 .metadata
11617 .rejection_reason
11618 .unwrap_or_else(|| "Input rejected".to_string());
11619 warn!(reason = %reason, "Input rejected");
11620 let response = AgentResponse::new(reason);
11621 self.finish_turn_if_root(&response).await?;
11622 return Ok(response);
11623 }
11624
11625 let processed_input = &input_data.content;
11626
11627 if let Some(response) = self.try_pre_response_transition(processed_input).await? {
11628 return Ok(response);
11629 }
11630
11631 if let Some(ref sm) = self.state_machine
11633 && let Some(def) = sm.current_definition()
11634 {
11635 if let Some(ref delegate_id) = def.delegate {
11636 return self
11637 .handle_delegated_state(processed_input, delegate_id, &def)
11638 .await;
11639 }
11640 if let Some(ref concurrent_config) = def.concurrent {
11641 return self
11642 .handle_concurrent_state(processed_input, concurrent_config)
11643 .await;
11644 }
11645 if let Some(ref group_chat_config) = def.group_chat {
11646 return self
11647 .handle_group_chat_state(processed_input, group_chat_config)
11648 .await;
11649 }
11650 if let Some(ref pipeline_config) = def.pipeline {
11651 return self
11652 .handle_pipeline_state(processed_input, pipeline_config)
11653 .await;
11654 }
11655 if let Some(ref handoff_config) = def.handoff {
11656 return self
11657 .handle_handoff_state(processed_input, handoff_config)
11658 .await;
11659 }
11660 }
11661
11662 if let Some(response) =
11667 Box::pin(self.try_speculative_branches(processed_input, &input_data.context)).await?
11668 {
11669 return Ok(response);
11670 }
11671
11672 match self.try_skill_route(processed_input).await? {
11673 SkillRouteResult::Response { skill_id, content } => {
11674 self.commit_root_user_message(processed_input).await?;
11675 return self
11676 .handle_skill_response(processed_input, &skill_id, content, &input_data.context)
11677 .await;
11678 }
11679 SkillRouteResult::NeedsClarification {
11680 response,
11681 ownership,
11682 } => {
11683 let admission = self
11684 .admit_optional_disambiguation_ownership(ownership)
11685 .await?;
11686 self.commit_root_user_message(processed_input).await?;
11687 if Self::skill_clarification_needs_memory_record(&response) {
11688 self.memory
11691 .add_message(ChatMessage::assistant(&response.content))
11692 .await?;
11693 }
11694 drop(admission);
11695 self.finish_turn_if_root(&response).await?;
11696 return Ok(response);
11697 }
11698 SkillRouteResult::NoMatch => {} }
11700
11701 let effective_reasoning = self.get_effective_reasoning_config();
11702 let reasoning_mode = self.determine_reasoning_mode(processed_input).await?;
11703 let auto_detected = matches!(effective_reasoning.mode, ReasoningMode::Auto);
11704 let reflection_enabled = self.routing_reflection_mode();
11706
11707 info!(
11708 reasoning_mode = ?reasoning_mode,
11709 auto_detected = auto_detected,
11710 reflection_enabled = ?reflection_enabled,
11711 "Reasoning mode determined"
11712 );
11713
11714 if matches!(reasoning_mode, ReasoningMode::PlanAndExecute) {
11715 self.commit_root_user_message(processed_input).await?;
11716 return self
11717 .handle_plan_and_execute(processed_input, &input_data.context, auto_detected)
11718 .await;
11719 }
11720
11721 self.commit_root_user_message(processed_input).await?;
11722
11723 let mut iterations = 0u32;
11724 let mut all_tool_calls: Vec<ToolCall> = Vec::new();
11725 let mut thinking_content: Option<String> = None;
11726
11727 let llm = self.get_state_llm()?;
11728
11729 loop {
11730 let effective_max = if reasoning_mode != ReasoningMode::None {
11732 let rc = self.get_effective_reasoning_config();
11733 self.max_iterations.min(rc.max_iterations)
11734 } else {
11735 self.max_iterations
11736 };
11737
11738 if iterations >= effective_max {
11739 let err = AgentError::Other(format!("Max iterations ({}) exceeded", effective_max));
11740 self.hooks.on_error(&err).await;
11741 error!(iterations = iterations, "Max iterations exceeded");
11742 return Err(err);
11743 }
11744 iterations += 1;
11745 *self.iteration_count.write() = iterations;
11746
11747 debug!(iteration = iterations, max = effective_max, "LLM call");
11748
11749 let protocol = self.main_tool_protocol(llm.as_ref(), false).await?;
11750 let mut messages = self
11751 .build_messages_internal(true, None, protocol.choice.is_none())
11752 .await?;
11753 self.inject_reasoning_prompt(&mut messages, &reasoning_mode, iterations == 1);
11754
11755 self.hooks.on_llm_start(&messages).await;
11756 let llm_start = Instant::now();
11757 let response = self
11758 .complete_main_llm_with_recovery(Arc::clone(&llm), &messages, &protocol)
11759 .await?;
11760
11761 let llm_duration_ms = llm_start.elapsed().as_millis() as u64;
11762 self.hooks.on_llm_complete(&response, llm_duration_ms).await;
11763
11764 let content = response.content.trim();
11765
11766 if let Some(tool_calls) = self.parse_main_tool_calls(content, &protocol)? {
11767 match self
11768 .handle_tool_calls(
11769 processed_input,
11770 content,
11771 tool_calls,
11772 &mut all_tool_calls,
11773 None,
11774 )
11775 .await?
11776 {
11777 ToolCallOutcome::Continue | ToolCallOutcome::TransitionFired => continue,
11778 ToolCallOutcome::Rejected(resp) => {
11779 self.finish_turn_if_root(&resp).await?;
11780 return Ok(resp);
11781 }
11782 }
11783 }
11784
11785 let (extracted_thinking, answer) = self.extract_thinking(content);
11786 if extracted_thinking.is_some() {
11787 thinking_content = extracted_thinking;
11788 }
11789
11790 let output_data = self.process_output(&answer, &input_data.context).await?;
11791
11792 let mut final_content = if output_data.metadata.rejected {
11793 output_data
11794 .metadata
11795 .rejection_reason
11796 .unwrap_or_else(|| answer.to_string())
11797 } else {
11798 output_data.content
11799 };
11800
11801 let reflection_metadata;
11803 (final_content, reflection_metadata) = self
11804 .run_reflection(&*llm, processed_input, final_content)
11805 .await?;
11806
11807 final_content =
11808 self.format_response_with_thinking(thinking_content.as_deref(), &final_content);
11809
11810 let final_content = {
11814 let result = self
11815 .post_loop_processing(processed_input, final_content)
11816 .await?;
11817 self.apply_post_loop_result(processed_input, result)
11818 .await?
11819 .content
11820 };
11821
11822 let reflected = reflection_metadata.is_some();
11823 let reasoning_mode_debug = format!("{:?}", reasoning_mode);
11824
11825 let response = self.build_agent_response(AgentResponseParts {
11826 content: final_content,
11827 all_tool_calls,
11828 reasoning_mode,
11829 auto_detected,
11830 iterations,
11831 thinking: thinking_content,
11832 reflection_metadata,
11833 });
11834
11835 self.finish_turn_if_root(&response).await?;
11836
11837 let tool_call_count = response.tool_calls.as_ref().map(|tc| tc.len()).unwrap_or(0);
11838 info!(
11839 tool_calls = tool_call_count,
11840 response_len = response.content.len(),
11841 reasoning_mode = %reasoning_mode_debug,
11842 reflected = reflected,
11843 "Chat completed"
11844 );
11845 return Ok(response);
11846 }
11847 }
11848
11849 async fn generate_buffered_streaming_draft(
11850 &self,
11851 processed_input: &str,
11852 routing_resolved: Arc<AtomicBool>,
11853 ) -> Result<StreamingDraftResult> {
11854 let llm = self.get_state_llm()?;
11855 if llm.configured_tool_choice().is_some() {
11856 let draft = self
11857 .generate_main_response_draft(processed_input, &ReasoningMode::None)
11858 .await?;
11859 return Ok(StreamingDraftResult::new(draft, Vec::new()));
11860 }
11861 let protocol = self.main_tool_protocol(llm.as_ref(), true).await?;
11863 let messages = self.build_messages_for_draft(processed_input).await?;
11864 let source = self
11865 .open_main_stream_with_recovery(Arc::clone(&llm), &messages, &protocol)
11866 .await?;
11867 let mut buffer = crate::optimization::StreamBranchBuffer::new(self.streaming.buffer_size)?;
11868 let mut chunks = Vec::new();
11869 let mut accumulated = String::new();
11870 match source {
11871 MainStreamSource::StaticResponse(text) => {
11872 accumulated.push_str(&text);
11873 let stream_chunk = StreamChunk::content(text);
11874 if routing_resolved.load(Ordering::SeqCst) {
11875 chunks.push(stream_chunk);
11876 } else {
11877 buffer.push(stream_chunk)?;
11878 }
11879 }
11880 MainStreamSource::Stream(mut stream) => {
11881 while let Some(chunk_result) = stream.next().await {
11882 let chunk = chunk_result.map_err(|e| AgentError::LLM(e.to_string()))?;
11883 accumulated.push_str(&chunk.delta);
11884 let stream_chunk = StreamChunk::content(chunk.delta);
11885 if routing_resolved.load(Ordering::SeqCst) {
11886 chunks.push(stream_chunk);
11887 } else {
11888 buffer.push(stream_chunk)?;
11889 }
11890 }
11891 }
11892 }
11893 chunks.splice(0..0, buffer.drain());
11894 let content = accumulated.trim().to_string();
11895 let draft = if let Some(calls) = self.parse_tool_calls(&content)? {
11896 MainResponseDraft::ToolCalls {
11897 raw_content: content,
11898 calls,
11899 thinking: None,
11900 }
11901 } else {
11902 MainResponseDraft::Text {
11903 raw_content: content,
11904 thinking: None,
11905 }
11906 };
11907 Ok(StreamingDraftResult::new(draft, chunks))
11908 }
11909
11910 async fn try_buffered_streaming_branches(
11911 &self,
11912 processed_input: &str,
11913 input_context: &HashMap<String, Value>,
11914 ) -> Result<Option<(AgentResponse, Vec<StreamChunk>)>> {
11915 let optimization = &self.runtime_config.optimization;
11916 if !optimization.enabled {
11917 return Ok(None);
11918 }
11919 if !matches!(
11925 self.get_effective_reasoning_config().mode,
11926 ReasoningMode::None
11927 ) {
11928 return Ok(None);
11929 }
11930 let transition_enabled =
11931 optimization.speculative_state_transitions && self.has_parallel_transition_candidates();
11932 if !transition_enabled {
11933 return Ok(None);
11934 }
11935 let mut branch_scheduler =
11936 TurnBranchScheduler::new(optimization.max_parallel_runtime_tasks)?;
11937 if !branch_scheduler.reserve_task() {
11938 return Ok(None);
11939 }
11940 if !self
11941 .reserve_active_speculative_llm_call(RuntimeOptimizationKind::BufferedStreamingRouting)
11942 {
11943 branch_scheduler.release_task();
11944 return Ok(None);
11945 }
11946 if !branch_scheduler.reserve_task() {
11947 branch_scheduler.release_task();
11948 return Ok(None);
11949 }
11950 let mut main_branch = RuntimeBranch::new(
11951 RuntimeTaskPurpose::MainResponse,
11952 RuntimeOptimizationKind::BufferedStreamingRouting,
11953 RuntimeTaskPriority::Normal,
11954 RuntimeCommitBehavior::FinalResponse,
11955 );
11956 let mut transition_branch = RuntimeBranch::new(
11957 RuntimeTaskPurpose::StateTransition,
11958 RuntimeOptimizationKind::ParallelStateTransition,
11959 RuntimeTaskPriority::Critical,
11960 RuntimeCommitBehavior::TransitionDecision,
11961 );
11962 let main_id = main_branch.branch_id();
11963 let transition_id = transition_branch.branch_id();
11964 let routing_resolved = Arc::new(AtomicBool::new(false));
11965 let mut main_future =
11966 Box::pin(crate::optimization::observability::with_branch_observation(
11967 &main_id,
11968 RuntimeOptimizationKind::BufferedStreamingRouting,
11969 RuntimeCommitBehavior::FinalResponse,
11970 self.generate_buffered_streaming_draft(
11971 processed_input,
11972 Arc::clone(&routing_resolved),
11973 ),
11974 ));
11975 let mut transition_future =
11976 Box::pin(crate::optimization::observability::with_branch_observation(
11977 &transition_id,
11978 RuntimeOptimizationKind::ParallelStateTransition,
11979 RuntimeCommitBehavior::TransitionDecision,
11980 self.select_parallel_transition_candidate(processed_input),
11981 ));
11982 let mut main_pending = true;
11983 let mut transition_pending = true;
11984 let mut main_result: Option<Result<StreamingDraftResult>> = None;
11985 let mut transition_finalized = false;
11986 let mut transition_candidate: Option<TransitionCandidate> = None;
11987 loop {
11988 if let Some(candidate) = transition_candidate.take() {
11989 if self
11990 .approve_transition_target(&candidate.from_state, candidate.target())
11991 .await?
11992 {
11993 drop(main_future);
11995 drop(transition_future);
11996 self.finalize_branch_loss(
11997 &main_id,
11998 RuntimeOptimizationKind::BufferedStreamingRouting,
11999 RuntimeCommitBehavior::FinalResponse,
12000 main_pending,
12001 main_result.as_ref().map(|result| result.is_err()),
12002 );
12003 if !self
12004 .apply_pre_response_transition_candidate(
12005 &candidate,
12006 &HashMap::new(),
12007 processed_input,
12008 )
12009 .await?
12010 {
12011 self.finalize_optional_branch(
12012 &transition_id,
12013 RuntimeOptimizationKind::ParallelStateTransition,
12014 RuntimeCommitBehavior::TransitionDecision,
12015 "discarded",
12016 false,
12017 );
12018 return Ok(None);
12019 }
12020 self.finalize_optional_branch(
12021 &transition_id,
12022 RuntimeOptimizationKind::ParallelStateTransition,
12023 RuntimeCommitBehavior::TransitionDecision,
12024 "committed",
12025 true,
12026 );
12027 let response = self.redispatch_current_state(processed_input).await?;
12028 return Ok(Some((
12029 response.clone(),
12030 vec![StreamChunk::content(response.content)],
12031 )));
12032 }
12033 self.finalize_optional_branch(
12034 &transition_id,
12035 RuntimeOptimizationKind::ParallelStateTransition,
12036 RuntimeCommitBehavior::TransitionDecision,
12037 "discarded",
12038 false,
12039 );
12040 transition_finalized = true;
12041 }
12042 if transition_finalized && !routing_resolved.load(Ordering::SeqCst) {
12048 match self
12049 .resolve_buffered_skill_after_transition(processed_input, &routing_resolved)
12050 .await
12051 {
12052 Ok(Some(candidate)) => {
12053 drop(main_future);
12055 drop(transition_future);
12056 self.finalize_branch_loss(
12057 &main_id,
12058 RuntimeOptimizationKind::BufferedStreamingRouting,
12059 RuntimeCommitBehavior::FinalResponse,
12060 main_pending,
12061 main_result.as_ref().map(|result| result.is_err()),
12062 );
12063 return match self
12064 .commit_winning_skill_candidate(
12065 candidate,
12066 processed_input,
12067 input_context,
12068 )
12069 .await?
12070 {
12071 Some(response) => Ok(Some((
12072 response.clone(),
12073 vec![StreamChunk::content(response.content)],
12074 ))),
12075 None => Ok(None),
12076 };
12077 }
12078 Ok(None) => {}
12079 Err(error) => {
12080 drop(main_future);
12081 drop(transition_future);
12082 self.finalize_branch_loss(
12083 &main_id,
12084 RuntimeOptimizationKind::BufferedStreamingRouting,
12085 RuntimeCommitBehavior::FinalResponse,
12086 main_pending,
12087 main_result.as_ref().map(|result| result.is_err()),
12088 );
12089 return Err(error);
12090 }
12091 }
12092 }
12093 if transition_finalized
12094 && routing_resolved.load(Ordering::SeqCst)
12095 && let Some(result) = main_result.take()
12096 {
12097 let stream_draft = match result {
12098 Ok(stream_draft) => stream_draft,
12099 Err(error) => {
12100 self.finalize_optional_branch(
12101 &main_id,
12102 RuntimeOptimizationKind::BufferedStreamingRouting,
12103 RuntimeCommitBehavior::FinalResponse,
12104 "failed",
12105 false,
12106 );
12107 return Err(error);
12108 }
12109 };
12110 let raw_draft_content = stream_draft.draft.raw_content().to_string();
12111 let buffered_chunks = stream_draft.chunks;
12112 self.finalize_optional_branch(
12113 &main_id,
12114 RuntimeOptimizationKind::BufferedStreamingRouting,
12115 RuntimeCommitBehavior::FinalResponse,
12116 "committed",
12117 true,
12118 );
12119 let response = self
12120 .commit_main_response_draft(
12121 processed_input,
12122 input_context,
12123 stream_draft.draft,
12124 ReasoningMode::None,
12125 false,
12126 )
12127 .await?;
12128 let chunks = if response.content == raw_draft_content {
12129 buffered_chunks
12130 } else {
12131 vec![StreamChunk::content(response.content.clone())]
12132 };
12133 return Ok(Some((response, chunks)));
12134 }
12135 tokio::select! {
12136 result = &mut main_future, if main_pending => {
12137 main_pending = false;
12138 main_branch.transition_to(RuntimeBranchStatus::Completed)?;
12139 main_result = Some(result);
12140 }
12141 result = &mut transition_future, if transition_pending => {
12142 transition_pending = false;
12143 transition_branch.transition_to(RuntimeBranchStatus::Completed)?;
12144 match result {
12145 Ok(ParallelTransitionSelection::Candidate(candidate)) => {
12146 transition_candidate = Some(candidate)
12147 }
12148 Ok(ParallelTransitionSelection::NoMatch) => {
12149 self.finalize_optional_branch(
12150 &transition_id,
12151 RuntimeOptimizationKind::ParallelStateTransition,
12152 RuntimeCommitBehavior::TransitionDecision,
12153 "discarded",
12154 false,
12155 );
12156 transition_finalized = true;
12157 }
12158 Ok(ParallelTransitionSelection::ReservationExhausted) => {
12159 self.finalize_optional_branch(
12160 &transition_id,
12161 RuntimeOptimizationKind::ParallelStateTransition,
12162 RuntimeCommitBehavior::TransitionDecision,
12163 "cancelled",
12164 false,
12165 );
12166 routing_resolved.store(true, Ordering::SeqCst);
12167 self.finalize_branch_loss(
12168 &main_id,
12169 RuntimeOptimizationKind::BufferedStreamingRouting,
12170 RuntimeCommitBehavior::FinalResponse,
12171 main_pending,
12172 main_result.as_ref().map(|result| result.is_err()),
12173 );
12174 return Ok(None);
12175 }
12176 Err(_) => {
12177 self.finalize_optional_branch(
12178 &transition_id,
12179 RuntimeOptimizationKind::ParallelStateTransition,
12180 RuntimeCommitBehavior::TransitionDecision,
12181 "failed",
12182 false,
12183 );
12184 transition_finalized = true;
12185 }
12186 }
12187 }
12188 }
12189 }
12190 }
12191
12192 async fn resolve_buffered_skill_after_transition(
12198 &self,
12199 processed_input: &str,
12200 routing_resolved: &AtomicBool,
12201 ) -> Result<Option<SkillCandidate>> {
12202 let candidate = if self.skill_router.is_some() {
12203 self.select_skill_candidate(processed_input).await?
12204 } else {
12205 None
12206 };
12207 if candidate.is_none() {
12208 routing_resolved.store(true, Ordering::SeqCst);
12209 }
12210 Ok(candidate)
12211 }
12212
12213 fn run_loop_internal_stream<'a>(
12217 &'a self,
12218 input: &'a str,
12219 terminal: RuntimeStreamTerminalSlot,
12220 ) -> Pin<Box<dyn Stream<Item = StreamChunk> + Send + 'a>> {
12221 let include_state_events = self.streaming.include_state_events;
12222
12223 Box::pin(async_stream::stream! {
12224 self.begin_root_turn();
12225 self.pre_turn_session_lifecycle().await;
12227
12228 let input_data = match self.process_input(input).await {
12229 Ok(data) => data,
12230 Err(e) => {
12231 yield StreamChunk::error(e.to_string());
12232 return;
12233 }
12234 };
12235 self.update_active_turn_context(&input_data.content, input_data.context.clone());
12236
12237 for (key, value) in &input_data.context {
12239 let _ = self.context_manager.set(key, value.clone());
12240 }
12241
12242 if input_data.metadata.rejected {
12243 let reason = input_data
12244 .metadata
12245 .rejection_reason
12246 .unwrap_or_else(|| "Input rejected".to_string());
12247 warn!(reason = %reason, "Input rejected (stream)");
12248 let response = AgentResponse::new(&reason);
12251 if let Err(e) = self.finish_turn_if_root(&response).await {
12252 yield StreamChunk::error(e.to_string());
12253 return;
12254 }
12255 yield StreamChunk::content(&reason);
12256 record_runtime_stream_final(&terminal, response);
12257 yield StreamChunk::Done {};
12258 return;
12259 }
12260
12261 let processed_input = &input_data.content;
12262
12263 let streaming_policy = self.runtime_config.optimization.streaming_policy;
12264
12265 if self.runtime_config.optimization.enabled
12272 && !matches!(
12273 streaming_policy,
12274 crate::optimization::StreamingOptimizationPolicy::Disabled
12275 )
12276 {
12277 match self.try_pre_response_transition(processed_input).await {
12278 Ok(Some(response)) => {
12279 yield StreamChunk::content(&response.content);
12280 record_runtime_stream_final(&terminal, response);
12281 yield StreamChunk::Done {};
12282 return;
12283 }
12284 Ok(None) => {}
12285 Err(e) => {
12286 yield StreamChunk::error(e.to_string());
12287 return;
12288 }
12289 }
12290 }
12291
12292 if self.runtime_config.optimization.enabled
12293 && matches!(
12294 streaming_policy,
12295 crate::optimization::StreamingOptimizationPolicy::BufferUntilRoutingDone
12296 )
12297 {
12298 match Box::pin(self.try_buffered_streaming_branches(processed_input, &input_data.context)).await {
12303 Ok(Some((response, chunks))) => {
12304 for chunk in chunks {
12305 yield chunk;
12306 }
12307 record_runtime_stream_final(&terminal, response);
12308 yield StreamChunk::Done {};
12309 return;
12310 }
12311 Ok(None) => {}
12312 Err(e) => {
12313 yield StreamChunk::error(e.to_string());
12314 return;
12315 }
12316 }
12317 }
12318
12319 if let Some(ref sm) = self.state_machine
12321 && let Some(def) = sm.current_definition()
12322 {
12323 let orchestration_result = if let Some(ref delegate_id) = def.delegate {
12324 Some(self.handle_delegated_state(processed_input, delegate_id, &def).await)
12325 } else if let Some(ref concurrent_config) = def.concurrent {
12326 Some(self.handle_concurrent_state(processed_input, concurrent_config).await)
12327 } else if let Some(ref group_chat_config) = def.group_chat {
12328 Some(self.handle_group_chat_state(processed_input, group_chat_config).await)
12329 } else if let Some(ref pipeline_config) = def.pipeline {
12330 Some(self.handle_pipeline_state(processed_input, pipeline_config).await)
12331 } else if let Some(ref handoff_config) = def.handoff {
12332 Some(self.handle_handoff_state(processed_input, handoff_config).await)
12333 } else {
12334 None
12335 };
12336
12337 if let Some(result) = orchestration_result {
12338 match result {
12339 Ok(response) => {
12340 yield StreamChunk::content(&response.content);
12341 record_runtime_stream_final(&terminal, response);
12342 yield StreamChunk::Done {};
12343 }
12344 Err(e) => {
12345 yield StreamChunk::error(e.to_string());
12346 }
12347 }
12348 return;
12349 }
12350 }
12351
12352 match self.try_skill_route(processed_input).await {
12354 Ok(SkillRouteResult::Response { skill_id, content }) => {
12355 if let Err(e) = self.commit_root_user_message(processed_input).await {
12356 yield StreamChunk::error(e.to_string());
12357 return;
12358 }
12359 match self.handle_skill_response(processed_input, &skill_id, content, &input_data.context).await {
12360 Ok(resp) => {
12361 yield StreamChunk::content(&resp.content);
12362 record_runtime_stream_final(&terminal, resp);
12363 yield StreamChunk::Done {};
12364 return;
12365 }
12366 Err(e) => {
12367 yield StreamChunk::error(e.to_string());
12368 return;
12369 }
12370 }
12371 }
12372 Ok(SkillRouteResult::NeedsClarification {
12373 response,
12374 ownership,
12375 }) => {
12376 let admission = match self
12377 .admit_optional_disambiguation_ownership(ownership)
12378 .await
12379 {
12380 Ok(admission) => admission,
12381 Err(e) => {
12382 yield StreamChunk::error(e.to_string());
12383 return;
12384 }
12385 };
12386 if let Err(e) = self.commit_root_user_message(processed_input).await {
12387 yield StreamChunk::error(e.to_string());
12388 return;
12389 }
12390 if Self::skill_clarification_needs_memory_record(&response)
12392 && let Err(e) = self.memory.add_message(ChatMessage::assistant(&response.content)).await
12393 {
12394 yield StreamChunk::error(e.to_string());
12395 return;
12396 }
12397 drop(admission);
12398 if let Err(e) = self.finish_turn_if_root(&response).await {
12399 yield StreamChunk::error(e.to_string());
12400 return;
12401 }
12402 yield StreamChunk::content(&response.content);
12403 record_runtime_stream_final(&terminal, response);
12404 yield StreamChunk::Done {};
12405 return;
12406 }
12407 Ok(SkillRouteResult::NoMatch) => {} Err(e) => {
12409 yield StreamChunk::error(e.to_string());
12410 return;
12411 }
12412 }
12413
12414 let effective_reasoning = self.get_effective_reasoning_config();
12416 let reasoning_mode = match self.determine_reasoning_mode(processed_input).await {
12417 Ok(mode) => mode,
12418 Err(e) => {
12419 yield StreamChunk::error(e.to_string());
12420 return;
12421 }
12422 };
12423 let auto_detected = matches!(effective_reasoning.mode, ReasoningMode::Auto);
12424
12425 info!(
12426 reasoning_mode = ?reasoning_mode,
12427 auto_detected = auto_detected,
12428 "Reasoning mode determined (stream)"
12429 );
12430
12431 if matches!(reasoning_mode, ReasoningMode::PlanAndExecute) {
12433 if let Err(e) = self.commit_root_user_message(processed_input).await {
12434 yield StreamChunk::error(e.to_string());
12435 return;
12436 }
12437 match self.handle_plan_and_execute(processed_input, &input_data.context, auto_detected).await {
12438 Ok(resp) => {
12439 yield StreamChunk::content(&resp.content);
12440 record_runtime_stream_final(&terminal, resp);
12441 yield StreamChunk::Done {};
12442 return;
12443 }
12444 Err(e) => {
12445 yield StreamChunk::error(e.to_string());
12446 return;
12447 }
12448 }
12449 }
12450
12451 if let Err(e) = self.commit_root_user_message(processed_input).await {
12452 yield StreamChunk::error(e.to_string());
12453 return;
12454 }
12455
12456 let llm = match self.get_state_llm() {
12457 Ok(llm) => llm,
12458 Err(e) => {
12459 yield StreamChunk::error(e.to_string());
12460 return;
12461 }
12462 };
12463
12464 let mut iterations = 0u32;
12465 let mut all_tool_calls: Vec<ToolCall> = Vec::new();
12466 let mut thinking_content: Option<String> = None;
12467
12468 loop {
12469 let effective_max = if reasoning_mode != ReasoningMode::None {
12471 let rc = self.get_effective_reasoning_config();
12472 self.max_iterations.min(rc.max_iterations)
12473 } else {
12474 self.max_iterations
12475 };
12476
12477 if iterations >= effective_max {
12478 let err_msg = format!("Max iterations ({}) exceeded", effective_max);
12479 let err = AgentError::Other(err_msg.clone());
12480 self.hooks.on_error(&err).await;
12481 error!(iterations = iterations, "Max iterations exceeded (stream)");
12482 yield StreamChunk::error(err_msg);
12483 return;
12484 }
12485 iterations += 1;
12486 *self.iteration_count.write() = iterations;
12487
12488 debug!(iteration = iterations, max = effective_max, "LLM call (stream)");
12489
12490 let protocol = match self.main_tool_protocol(llm.as_ref(), false).await {
12491 Ok(protocol) => protocol,
12492 Err(e) => {
12493 yield StreamChunk::error(e.to_string());
12494 return;
12495 }
12496 };
12497 let mut messages = match self
12498 .build_messages_internal(true, None, protocol.choice.is_none())
12499 .await
12500 {
12501 Ok(m) => m,
12502 Err(e) => {
12503 yield StreamChunk::error(e.to_string());
12504 return;
12505 }
12506 };
12507 self.inject_reasoning_prompt(&mut messages, &reasoning_mode, iterations == 1);
12508
12509 self.hooks.on_llm_start(&messages).await;
12510 let llm_start = Instant::now();
12511
12512 let buffered_decision = self.main_stream_must_buffer(&reasoning_mode, &protocol);
12513 let content = if buffered_decision {
12514 let response = match self
12518 .complete_main_llm_with_recovery(
12519 Arc::clone(&llm),
12520 &messages,
12521 &protocol,
12522 )
12523 .await
12524 {
12525 Ok(r) => r,
12526 Err(e) => {
12527 yield StreamChunk::error(e.to_string());
12528 return;
12529 }
12530 };
12531 let llm_duration_ms = llm_start.elapsed().as_millis() as u64;
12532 self.hooks.on_llm_complete(&response, llm_duration_ms).await;
12533 response.content.trim().to_string()
12534 } else {
12535 let source = match self
12537 .open_main_stream_with_recovery(Arc::clone(&llm), &messages, &protocol)
12538 .await
12539 {
12540 Ok(source) => source,
12541 Err(e) => {
12542 yield StreamChunk::error(e.to_string());
12543 return;
12544 }
12545 };
12546 let mut accumulated = String::new();
12547 match source {
12548 MainStreamSource::StaticResponse(text) => {
12549 accumulated.push_str(&text);
12550 yield StreamChunk::content(text);
12551 }
12552 MainStreamSource::Stream(mut stream_inner) => {
12553 while let Some(chunk_result) = stream_inner.next().await {
12554 match chunk_result {
12555 Ok(chunk) => {
12556 accumulated.push_str(&chunk.delta);
12557 yield StreamChunk::content(chunk.delta);
12558 }
12559 Err(e) => {
12560 yield StreamChunk::error(e.to_string());
12562 return;
12563 }
12564 }
12565 }
12566 }
12567 }
12568 let llm_duration_ms = llm_start.elapsed().as_millis() as u64;
12569 let llm_response = ai_agents_core::LLMResponse::new(
12571 accumulated.trim(),
12572 ai_agents_core::FinishReason::Stop,
12573 );
12574 self.hooks.on_llm_complete(&llm_response, llm_duration_ms).await;
12575 accumulated.trim().to_string()
12576 };
12577
12578 let parsed_tool_calls = match self.parse_main_tool_calls(&content, &protocol) {
12580 Ok(calls) => calls,
12581 Err(error) => {
12582 yield StreamChunk::error(error.to_string());
12583 return;
12584 }
12585 };
12586 if let Some(tool_calls) = parsed_tool_calls {
12587 let mut events = Vec::new();
12590 let outcome = self
12591 .handle_tool_calls(
12592 processed_input,
12593 &content,
12594 tool_calls,
12595 &mut all_tool_calls,
12596 Some(&mut events),
12597 )
12598 .await;
12599 for chunk in events.drain(..) {
12600 yield chunk;
12601 }
12602 match outcome {
12603 Ok(ToolCallOutcome::Continue) | Ok(ToolCallOutcome::TransitionFired) => continue,
12604 Ok(ToolCallOutcome::Rejected(response)) => {
12605 if let Err(finalize_error) = self.finish_turn_if_root(&response).await {
12606 yield StreamChunk::error(finalize_error.to_string());
12607 return;
12608 }
12609 let legacy_error = response.content.clone();
12610 record_runtime_stream_final(&terminal, response);
12611 yield StreamChunk::error(legacy_error);
12612 yield StreamChunk::Done {};
12613 return;
12614 }
12615 Err(e) => {
12616 yield StreamChunk::error(e.to_string());
12617 return;
12618 }
12619 }
12620 }
12621
12622 let (extracted_thinking, answer) = self.extract_thinking(&content);
12624 if extracted_thinking.is_some() {
12625 thinking_content = extracted_thinking;
12626 }
12627
12628 let output_data = match self.process_output(&answer, &input_data.context).await {
12629 Ok(d) => d,
12630 Err(e) => {
12631 yield StreamChunk::error(e.to_string());
12632 return;
12633 }
12634 };
12635
12636 let final_content = if output_data.metadata.rejected {
12637 output_data
12638 .metadata
12639 .rejection_reason
12640 .unwrap_or_else(|| answer.to_string())
12641 } else {
12642 output_data.content
12643 };
12644
12645 let (final_content, reflection_metadata) = match self
12647 .run_reflection(&*llm, processed_input, final_content)
12648 .await
12649 {
12650 Ok(r) => r,
12651 Err(e) => {
12652 yield StreamChunk::error(e.to_string());
12653 return;
12654 }
12655 };
12656
12657 let final_content = self.format_response_with_thinking(
12658 thinking_content.as_deref(),
12659 &final_content,
12660 );
12661
12662 if buffered_decision {
12664 yield StreamChunk::content(&final_content);
12665 }
12666
12667 let post_result = match self
12671 .post_loop_processing(processed_input, final_content)
12672 .await
12673 {
12674 Ok(r) => r,
12675 Err(e) => {
12676 yield StreamChunk::error(e.to_string());
12677 return;
12678 }
12679 };
12680
12681 let applied = match self.apply_post_loop_result(processed_input, post_result).await {
12682 Ok(applied) => applied,
12683 Err(e) => {
12684 yield StreamChunk::error(e.to_string());
12685 return;
12686 }
12687 };
12688
12689 if applied.transitioned {
12690 if include_state_events
12691 && let Some(state) = self.current_state()
12692 {
12693 yield StreamChunk::state_transition(None, state);
12694 }
12695 if applied.regenerated {
12701 yield StreamChunk::content(&applied.content);
12702 }
12703 }
12704 let final_content = applied.content;
12705
12706 let final_response = self.build_agent_response(AgentResponseParts {
12708 content: final_content,
12709 all_tool_calls,
12710 reasoning_mode,
12711 auto_detected,
12712 iterations,
12713 thinking: thinking_content,
12714 reflection_metadata,
12715 });
12716 if let Err(e) = self.finish_turn_if_root(&final_response).await {
12717 yield StreamChunk::error(e.to_string());
12718 return;
12719 }
12720
12721 record_runtime_stream_final(&terminal, final_response);
12722 yield StreamChunk::Done {};
12723 return;
12724 }
12725 })
12726 }
12727
12728 fn run_loop_stream<'a>(
12731 &'a self,
12732 input: &'a str,
12733 terminal: RuntimeStreamTerminalSlot,
12734 ) -> Pin<Box<dyn Stream<Item = StreamChunk> + Send + 'a>> {
12735 Box::pin(async_stream::stream! {
12736 self.begin_root_turn();
12737 let _root_cleanup = RootTurnCleanup::new(self);
12738 self.hooks.on_message_received(input).await;
12739
12740 if let Err(e) = self.prepare_turn_context().await {
12742 yield StreamChunk::error(e.to_string());
12743 return;
12744 }
12745
12746 self.clear_disambiguation_context();
12748
12749 let input_to_run = match self.resolve_disambiguation(input).await {
12753 Err(e) => {
12754 yield StreamChunk::error(e.to_string());
12755 return;
12756 }
12757 Ok(DisambiguationDispatch::Terminal(response)) => {
12758 yield StreamChunk::content(&response.content);
12759 record_runtime_stream_final(&terminal, response);
12760 yield StreamChunk::Done {};
12761 return;
12762 }
12763 Ok(DisambiguationDispatch::RecheckSkill {
12764 skill_id,
12765 enriched_input,
12766 disambiguation_epoch,
12767 state_generation,
12768 }) => {
12769 match self
12770 .recheck_skill_disambiguation(
12771 &skill_id,
12772 &enriched_input,
12773 disambiguation_epoch,
12774 state_generation,
12775 )
12776 .await
12777 {
12778 Ok(resp) => {
12779 yield StreamChunk::content(&resp.content);
12780 record_runtime_stream_final(&terminal, resp);
12781 yield StreamChunk::Done {};
12782 return;
12783 }
12784 Err(e) => {
12785 yield StreamChunk::error(e.to_string());
12786 return;
12787 }
12788 }
12789 }
12790 Ok(DisambiguationDispatch::Proceed(input)) => input,
12791 };
12792
12793 let mut inner = self.run_loop_internal_stream(&input_to_run, Arc::clone(&terminal));
12794 while let Some(chunk) = inner.next().await {
12795 yield chunk;
12796 }
12797 })
12798 }
12799
12800 pub fn info(&self) -> AgentInfo {
12801 self.info.clone()
12802 }
12803
12804 pub fn skills(&self) -> &[SkillDefinition] {
12805 &self.skills
12806 }
12807
12808 async fn reset_runtime_state(&self) -> Result<()> {
12810 let _admission = self.disambiguation_admission.write().await;
12811 if self.state_transition_reserved.load(Ordering::SeqCst) {
12812 return Err(AgentError::Other(
12813 "Cannot reset while a state transition is in progress".to_string(),
12814 ));
12815 }
12816 self.disambiguation_epoch.fetch_add(1, Ordering::SeqCst);
12817 *self.pending_skill_id.write() = None;
12818 if let Some(disambiguator) = self.disambiguation_manager.as_ref() {
12819 disambiguator.clear_pending().await;
12820 }
12821 self.memory.clear().await?;
12822 self.active_native_exchanges.write().clear();
12823 *self.iteration_count.write() = 0;
12824 self.tool_call_history.write().clear();
12825 if let Some(ref sm) = self.state_machine {
12826 sm.reset();
12827 }
12828 Ok(())
12829 }
12830
12831 pub async fn reset(&self) -> Result<()> {
12833 self.reset_runtime_state().await
12834 }
12835
12836 pub fn max_context_tokens(&self) -> u32 {
12837 self.max_context_tokens
12838 }
12839
12840 pub fn llm_registry(&self) -> &Arc<LLMRegistry> {
12841 &self.llm_registry
12842 }
12843
12844 pub fn state_machine(&self) -> Option<&Arc<StateMachine>> {
12845 self.state_machine.as_ref()
12846 }
12847
12848 pub fn context_manager(&self) -> &Arc<ContextManager> {
12849 &self.context_manager
12850 }
12851
12852 pub fn tool_call_history(&self) -> Vec<ToolCallRecord> {
12853 self.tool_call_history.read().clone()
12854 }
12855
12856 pub fn memory_token_budget(&self) -> Option<&MemoryTokenBudget> {
12857 self.memory_token_budget.as_ref()
12858 }
12859
12860 pub fn parallel_tools_config(&self) -> &ParallelToolsConfig {
12861 &self.parallel_tools
12862 }
12863
12864 pub fn streaming_config(&self) -> &StreamingConfig {
12865 &self.streaming
12866 }
12867
12868 pub fn hooks(&self) -> &Arc<dyn AgentHooks> {
12869 &self.hooks
12870 }
12871
12872 pub fn hitl_engine(&self) -> Option<&HITLEngine> {
12873 self.hitl_engine.as_ref()
12874 }
12875
12876 pub fn approval_handler(&self) -> &Arc<dyn ApprovalHandler> {
12877 &self.approval_handler
12878 }
12879
12880 fn build_hitl_language_context(&self) -> HashMap<String, Value> {
12882 let mut ctx = HashMap::new();
12883 for key in &["user.language", "input.detected.language", "language"] {
12884 if let Some(val) = self.context_manager.get(key) {
12885 ctx.insert(key.to_string(), val);
12886 }
12887 }
12888 ctx
12889 }
12890
12891 async fn request_hitl_approval(&self, check_result: HITLCheckResult) -> Result<ApprovalResult> {
12893 let Some(request) = check_result.into_request() else {
12894 return Ok(ApprovalResult::Approved);
12895 };
12896
12897 self.hooks.on_approval_requested(&request).await;
12898
12899 let timeout = request.timeout;
12900
12901 let raw_result = if let Some(duration) = timeout {
12902 match tokio::time::timeout(
12903 duration,
12904 self.approval_handler.request_approval(request.clone()),
12905 )
12906 .await
12907 {
12908 Ok(result) => result,
12909 Err(_) => ApprovalResult::timeout(),
12910 }
12911 } else {
12912 self.approval_handler
12913 .request_approval(request.clone())
12914 .await
12915 };
12916
12917 self.hooks
12918 .on_approval_result(&request.id, &raw_result)
12919 .await;
12920
12921 let (outcome, effective_result): (ApprovalResolvedOutcome, Result<ApprovalResult>) =
12922 match &raw_result {
12923 ApprovalResult::Approved => (
12924 ApprovalResolvedOutcome::Approved,
12925 Ok(ApprovalResult::Approved),
12926 ),
12927 ApprovalResult::Rejected { reason } => (
12928 ApprovalResolvedOutcome::Rejected {
12929 reason: reason.clone(),
12930 },
12931 Ok(ApprovalResult::Rejected {
12932 reason: reason.clone(),
12933 }),
12934 ),
12935 ApprovalResult::Modified { changes } => (
12936 ApprovalResolvedOutcome::Modified {
12937 changes: changes.clone(),
12938 },
12939 Ok(ApprovalResult::Modified {
12940 changes: changes.clone(),
12941 }),
12942 ),
12943 ApprovalResult::Timeout => {
12944 if let Some(ref engine) = self.hitl_engine {
12945 match engine.config().on_timeout {
12946 TimeoutAction::Approve => (
12947 ApprovalResolvedOutcome::Approved,
12948 Ok(ApprovalResult::Approved),
12949 ),
12950 TimeoutAction::Reject => {
12951 let reason = Some("Timeout".to_string());
12952 (
12953 ApprovalResolvedOutcome::Rejected {
12954 reason: reason.clone(),
12955 },
12956 Ok(ApprovalResult::Rejected { reason }),
12957 )
12958 }
12959 TimeoutAction::Error => {
12960 let message = "HITL approval timeout".to_string();
12961 (
12962 ApprovalResolvedOutcome::Error {
12963 message: message.clone(),
12964 },
12965 Err(AgentError::Other(message)),
12966 )
12967 }
12968 }
12969 } else {
12970 let reason = Some("Timeout (no engine)".to_string());
12971 (
12972 ApprovalResolvedOutcome::Rejected {
12973 reason: reason.clone(),
12974 },
12975 Ok(ApprovalResult::Rejected { reason }),
12976 )
12977 }
12978 }
12979 };
12980
12981 self.hooks
12982 .on_approval_resolved(&request, &raw_result, &outcome)
12983 .await;
12984
12985 effective_result
12986 }
12987
12988 pub async fn check_state_hitl(&self, from: Option<&str>, to: &str) -> Result<bool> {
12989 if let Some(ref hitl_engine) = self.hitl_engine {
12990 let hitl_lang_ctx = self.build_hitl_language_context();
12991 let check_result = self
12992 .observe_purpose(
12993 ObservationPurpose::HitlLocalization,
12994 hitl_engine.check_state_transition_with_localization(
12995 from,
12996 to,
12997 &hitl_lang_ctx,
12998 self.approval_handler.as_ref(),
12999 Some(&self.llm_registry),
13000 ),
13001 )
13002 .await?;
13003 if check_result.is_required() {
13004 let result = self.request_hitl_approval(check_result).await?;
13005 return Ok(matches!(
13006 result,
13007 ApprovalResult::Approved | ApprovalResult::Modified { .. }
13008 ));
13009 }
13010 }
13011 Ok(true)
13012 }
13013
13014 async fn execute_tools_parallel(
13016 &self,
13017 tool_calls: &[ToolCall],
13018 ) -> Vec<(String, Result<String>)> {
13019 let can_run_parallel = tool_calls.iter().all(|tc| {
13020 self.tools
13021 .resolve(&tc.name)
13022 .map(|resolved| resolved.tool.classify_call(&tc.arguments).concurrency_safe)
13023 .unwrap_or(false)
13024 });
13025
13026 if !self.parallel_tools.enabled || tool_calls.len() <= 1 || !can_run_parallel {
13027 let mut results = Vec::new();
13028 for tc in tool_calls {
13029 let result = self
13030 .observe_purpose(
13031 current_observation_context()
13032 .map(|context| context.purpose)
13033 .unwrap_or_default(),
13034 self.execute_tool_smart(tc),
13035 )
13036 .await;
13037 results.push((tc.id.clone(), result));
13038 }
13039 return results;
13040 }
13041
13042 let chunks: Vec<_> = tool_calls
13043 .chunks(self.parallel_tools.max_parallel)
13044 .collect();
13045
13046 let mut all_results = Vec::new();
13047
13048 for chunk in chunks {
13049 let futures: Vec<_> = chunk
13050 .iter()
13051 .map(|tc| {
13052 let tc = tc.clone();
13053 async move {
13054 let result = self.execute_tool_smart(&tc).await;
13055 (tc.id.clone(), result)
13056 }
13057 })
13058 .collect();
13059
13060 let results = futures::future::join_all(futures).await;
13061 all_results.extend(results);
13062 }
13063
13064 all_results
13065 }
13066
13067 pub async fn chat_stream<'a>(
13071 &'a self,
13072 input: &'a str,
13073 ) -> Result<Pin<Box<dyn Stream<Item = StreamChunk> + Send + 'a>>> {
13074 let RootTurnAdmission {
13075 guard: root_turn_guard,
13076 identity_stack,
13077 } = self.acquire_root_turn().await?;
13078 scope_runtime_gate_identity_stack(&identity_stack, self.init_storage()).await?;
13082 info!(input_len = input.len(), "Starting streaming chat");
13083 let terminal = new_runtime_stream_terminal_slot();
13084 let inner = self.run_loop_stream(input, terminal);
13085 let observation_context = self.build_observation_context(None);
13086 let stream: Pin<Box<dyn Stream<Item = StreamChunk> + Send + 'a>> =
13087 Box::pin(async_stream::stream! {
13088 let mut root_turn_guard = Some(root_turn_guard);
13089 let mut inner = inner;
13090 loop {
13091 let next = scope_runtime_gate_identity_stack(&identity_stack, async {
13092 if let Some(context) = observation_context.as_ref() {
13093 with_observation_context(context.clone(), inner.next()).await
13094 } else {
13095 inner.next().await
13096 }
13097 })
13098 .await;
13099 match next {
13100 Some(StreamChunk::Done {}) => {
13101 while scope_runtime_gate_identity_stack(&identity_stack, async {
13102 if let Some(context) = observation_context.as_ref() {
13103 with_observation_context(context.clone(), inner.next())
13104 .await
13105 .is_some()
13106 } else {
13107 inner.next().await.is_some()
13108 }
13109 })
13110 .await
13111 {}
13112 if observation_context.is_some() {
13113 scope_runtime_gate_identity_stack(
13114 &identity_stack,
13115 self.export_observability_if_configured(),
13116 )
13117 .await;
13118 }
13119 drop(root_turn_guard.take());
13120 yield StreamChunk::Done {};
13121 return;
13122 }
13123 Some(chunk) => yield chunk,
13124 None => {
13125 if observation_context.is_some() {
13126 scope_runtime_gate_identity_stack(
13127 &identity_stack,
13128 self.export_observability_if_configured(),
13129 )
13130 .await;
13131 }
13132 drop(root_turn_guard.take());
13133 return;
13134 }
13135 }
13136 }
13137 });
13138 Ok(stream)
13139 }
13140
13141 pub async fn chat_stream_events<'a>(
13145 &'a self,
13146 input: &'a str,
13147 ) -> Result<Pin<Box<dyn Stream<Item = AgentStreamEvent> + Send + 'a>>> {
13148 let RootTurnAdmission {
13149 guard,
13150 identity_stack,
13151 } = self.acquire_root_turn().await?;
13152 scope_runtime_gate_identity_stack(&identity_stack, self.init_storage()).await?;
13156 info!(input_len = input.len(), "Starting streaming chat events");
13157 let terminal = new_runtime_stream_terminal_slot();
13158 let inner = self.run_loop_stream(input, Arc::clone(&terminal));
13159 let observation_context = self.build_observation_context(None);
13160 Ok(self.drive_event_stream(
13161 inner,
13162 terminal,
13163 guard,
13164 identity_stack,
13165 observation_context,
13166 None,
13167 ))
13168 }
13169
13170 pub async fn chat_stream_events_with_actor_context<'a>(
13176 &'a self,
13177 input: &'a str,
13178 actor_context: crate::TurnActorContext,
13179 ) -> Result<Pin<Box<dyn Stream<Item = AgentStreamEvent> + Send + 'a>>> {
13180 let RootTurnAdmission {
13181 guard,
13182 identity_stack,
13183 } = self.acquire_root_turn().await?;
13184 scope_runtime_gate_identity_stack(&identity_stack, self.init_storage()).await?;
13185 info!(
13186 input_len = input.len(),
13187 "Starting streaming chat events with actor context"
13188 );
13189 let actor_id = actor_context.effective_actor_id().map(str::to_string);
13190 let terminal = new_runtime_stream_terminal_slot();
13191 let inner = self.run_loop_stream(input, Arc::clone(&terminal));
13192 let observation_context = self.build_observation_context(actor_id);
13193 Ok(self.drive_event_stream(
13194 inner,
13195 terminal,
13196 guard,
13197 identity_stack,
13198 observation_context,
13199 Some(actor_context),
13200 ))
13201 }
13202
13203 fn drive_event_stream<'a>(
13212 &'a self,
13213 mut inner: Pin<Box<dyn Stream<Item = StreamChunk> + Send + 'a>>,
13214 terminal: RuntimeStreamTerminalSlot,
13215 root_turn_guard: tokio::sync::OwnedMutexGuard<()>,
13216 identity_stack: RootTurnGateIdentityStack,
13217 observation_context: Option<SpanContext>,
13218 actor_context: Option<crate::TurnActorContext>,
13219 ) -> Pin<Box<dyn Stream<Item = AgentStreamEvent> + Send + 'a>> {
13220 Box::pin(async_stream::stream! {
13221 let mut root_turn_guard = Some(root_turn_guard);
13222 loop {
13223 let next = poll_scoped_chunk(
13224 &mut inner,
13225 &identity_stack,
13226 observation_context.as_ref(),
13227 actor_context.as_ref(),
13228 )
13229 .await;
13230 match next {
13231 Some(StreamChunk::Done {}) => {
13232 let terminal_event = { terminal.write().take() };
13233 if let Some(response) = terminal_event {
13234 while poll_scoped_chunk(
13235 &mut inner,
13236 &identity_stack,
13237 observation_context.as_ref(),
13238 actor_context.as_ref(),
13239 )
13240 .await
13241 .is_some()
13242 {}
13243 if observation_context.is_some() {
13244 scope_runtime_gate_identity_stack(
13245 &identity_stack,
13246 self.export_observability_if_configured(),
13247 )
13248 .await;
13249 }
13250 drop(root_turn_guard.take());
13251 yield AgentStreamEvent::Final(response);
13252 return;
13253 }
13254 }
13255 Some(StreamChunk::Error { message }) => {
13256 let finalized = { terminal.read().is_some() };
13257 if finalized {
13258 continue;
13259 }
13260 while poll_scoped_chunk(
13261 &mut inner,
13262 &identity_stack,
13263 observation_context.as_ref(),
13264 actor_context.as_ref(),
13265 )
13266 .await
13267 .is_some()
13268 {}
13269 if observation_context.is_some() {
13270 scope_runtime_gate_identity_stack(
13271 &identity_stack,
13272 self.export_observability_if_configured(),
13273 )
13274 .await;
13275 }
13276 drop(root_turn_guard.take());
13277 yield AgentStreamEvent::Chunk(StreamChunk::Error { message });
13278 return;
13279 }
13280 Some(chunk) => yield AgentStreamEvent::Chunk(chunk),
13281 None => {
13282 if observation_context.is_some() {
13283 scope_runtime_gate_identity_stack(
13284 &identity_stack,
13285 self.export_observability_if_configured(),
13286 )
13287 .await;
13288 }
13289 drop(root_turn_guard.take());
13290 return;
13291 }
13292 }
13293 }
13294 })
13295 }
13296}
13297
13298async fn poll_scoped_chunk<'a>(
13304 inner: &mut Pin<Box<dyn Stream<Item = StreamChunk> + Send + 'a>>,
13305 identity_stack: &RootTurnGateIdentityStack,
13306 observation_context: Option<&SpanContext>,
13307 actor_context: Option<&crate::TurnActorContext>,
13308) -> Option<StreamChunk> {
13309 scope_runtime_gate_identity_stack(identity_stack, async {
13310 let next = inner.next();
13311 match (observation_context, actor_context) {
13312 (Some(observation), Some(actor)) => {
13313 with_observation_context(
13314 observation.clone(),
13315 scope_actor_context(actor.clone(), next),
13316 )
13317 .await
13318 }
13319 (Some(observation), None) => with_observation_context(observation.clone(), next).await,
13320 (None, Some(actor)) => scope_actor_context(actor.clone(), next).await,
13321 (None, None) => next.await,
13322 }
13323 })
13324 .await
13325}
13326
13327#[async_trait]
13328impl ToolInvoker for RuntimeAgent {
13329 async fn invoke_tool(&self, request: ToolExecutionRequest) -> Result<ToolExecutionRecord> {
13330 self.execute_tool_record(request).await
13331 }
13332}
13333
13334#[async_trait]
13335impl Agent for RuntimeAgent {
13336 async fn chat(&self, input: &str) -> Result<AgentResponse> {
13338 let RootTurnAdmission {
13339 guard,
13340 identity_stack,
13341 } = self.acquire_root_turn().await?;
13342 let result = scope_runtime_gate_identity_stack(&identity_stack, async {
13343 let result = if let Some(context) = self.build_observation_context(None) {
13344 with_observation_context(context, self.run_loop(input)).await
13345 } else {
13346 self.run_loop(input).await
13347 };
13348 self.export_observability_if_configured().await;
13349 result
13350 })
13351 .await;
13352 drop(guard);
13353 result
13354 }
13355
13356 fn info(&self) -> AgentInfo {
13357 self.info.clone()
13358 }
13359
13360 async fn reset(&self) -> Result<()> {
13362 self.reset_runtime_state().await
13363 }
13364}
13365
13366fn background_maintenance_tags(
13376 label: &str,
13377 stage: &str,
13378 reason: Option<&str>,
13379 policy: Option<&crate::optimization::config::MaintenanceTaskPolicy>,
13380) -> HashMap<String, String> {
13381 let mut tags = HashMap::new();
13382 tags.insert("runtime.background".to_string(), "true".to_string());
13383 tags.insert("runtime.maintenance".to_string(), label.to_string());
13384 tags.insert("runtime.maintenance_stage".to_string(), stage.to_string());
13385 if let Some(policy) = policy {
13386 tags.insert(
13387 "runtime.await_before_next_turn".to_string(),
13388 await_before_next_turn_label(policy.await_before_next_turn).to_string(),
13389 );
13390 tags.insert(
13391 "runtime.maintenance_mode".to_string(),
13392 maintenance_mode_label(policy.mode).to_string(),
13393 );
13394 }
13395 if let Some(reason) = reason {
13396 tags.insert("runtime.reason".to_string(), reason.to_string());
13397 }
13398 tags
13399}
13400
13401fn await_before_next_turn_label(policy: AwaitBeforeNextTurn) -> &'static str {
13402 match policy {
13403 AwaitBeforeNextTurn::Never => "never",
13404 AwaitBeforeNextTurn::SameActor => "same_actor",
13405 AwaitBeforeNextTurn::Always => "always",
13406 }
13407}
13408
13409fn maintenance_mode_label(mode: MaintenanceMode) -> &'static str {
13410 match mode {
13411 MaintenanceMode::InlineSerial => "inline_serial",
13412 MaintenanceMode::InlineParallel => "inline_parallel",
13413 MaintenanceMode::Background => "background",
13414 }
13415}
13416
13417fn record_background_maintenance_event(
13419 manager: Option<&Arc<ObservabilityManager>>,
13420 label: &str,
13421 status: EventStatus,
13422 duration_ms: u64,
13423 stage: &str,
13424 reason: Option<String>,
13425 policy: Option<&crate::optimization::config::MaintenanceTaskPolicy>,
13426) {
13427 if let Some(manager) = manager {
13428 manager.record_lifecycle_event(
13429 EventType::MemoryOperation {
13430 operation: format!("{}_background_{}", label, stage),
13431 },
13432 ObservationPurpose::Other(format!("{}_maintenance", label)),
13433 status,
13434 duration_ms,
13435 background_maintenance_tags(label, stage, reason.as_deref(), policy),
13436 None,
13437 );
13438 }
13439}
13440
13441fn effective_maintenance_mode(mode: MaintenanceMode, force_parallel: bool) -> MaintenanceMode {
13442 if force_parallel && matches!(mode, MaintenanceMode::InlineSerial) {
13443 MaintenanceMode::InlineParallel
13444 } else {
13445 mode
13446 }
13447}
13448
13449fn observation_purpose_for_process(hint: ProcessPurposeHint) -> ObservationPurpose {
13450 match hint {
13451 ProcessPurposeHint::Detect => ObservationPurpose::ProcessDetect,
13452 ProcessPurposeHint::Extract => ObservationPurpose::ProcessExtract,
13453 ProcessPurposeHint::Validate => ObservationPurpose::ProcessValidate,
13454 ProcessPurposeHint::Transform | ProcessPurposeHint::Other => {
13455 ObservationPurpose::ProcessTransform
13456 }
13457 }
13458}
13459
13460fn new_tool_resource_locks() -> ToolResourceLocks {
13461 Arc::new(RwLock::new(HashMap::new()))
13462}
13463
13464fn tool_resource_lock_keys(
13469 _canonical_id: &str,
13470 args: &Value,
13471 bindings: &ai_agents_core::ToolPolicyBindings,
13472 classification: &ai_agents_core::ToolCallClassification,
13473) -> Vec<String> {
13474 if classification.concurrency_safe {
13475 return Vec::new();
13476 }
13477
13478 let mut keys = Vec::new();
13479 let mut has_path_resource = false;
13480 for binding in &bindings.path_fields {
13481 let value = value_at_argument_path(args, &binding.field)
13482 .cloned()
13483 .or_else(|| {
13484 binding
13485 .default_path
13486 .as_ref()
13487 .map(|path| Value::String(path.clone()))
13488 });
13489 if let Some(value) = value {
13490 collect_resource_strings(&value, |_| {
13491 has_path_resource = true;
13492 });
13493 }
13494 }
13495 for binding in &bindings.domain_fields {
13496 if let Some(value) = value_at_argument_path(args, &binding.field) {
13497 collect_resource_strings(value, |domain| {
13498 let normalized = if binding.is_url {
13499 normalized_url_resource_key(domain)
13500 } else {
13501 domain.trim().trim_end_matches('.').to_ascii_lowercase()
13502 };
13503 keys.push(format!("domain:{}", normalized));
13504 });
13505 }
13506 }
13507 for binding in &bindings.command_fields {
13508 if !matches!(binding.kind, ai_agents_core::CommandBindingKind::Cwd) {
13509 continue;
13510 }
13511 if let Some(value) = value_at_argument_path(args, &binding.field) {
13512 collect_resource_strings(value, |_| {
13513 has_path_resource = true;
13514 });
13515 }
13516 }
13517 if has_path_resource {
13518 keys.push("path-mutation:global".to_string());
13519 }
13520 if keys.is_empty() {
13521 keys.push("side-effect:unbound".to_string());
13522 }
13523 keys.sort();
13524 keys.dedup();
13525 keys
13526}
13527
13528fn value_at_argument_path<'a>(value: &'a Value, field: &str) -> Option<&'a Value> {
13529 let mut current = value;
13530 for segment in field.split('.') {
13531 if segment.is_empty() {
13532 return None;
13533 }
13534 current = current.get(segment)?;
13535 }
13536 Some(current)
13537}
13538
13539fn collect_resource_strings(value: &Value, mut collect: impl FnMut(&str)) {
13540 match value {
13541 Value::String(value) => collect(value),
13542 Value::Array(values) => {
13543 for value in values {
13544 if let Some(value) = value.as_str() {
13545 collect(value);
13546 }
13547 }
13548 }
13549 _ => {}
13550 }
13551}
13552
13553fn normalized_url_resource_key(value: &str) -> String {
13554 let value = value.trim();
13555 let Some((scheme, remainder)) = value.split_once("://") else {
13556 return value.to_ascii_lowercase();
13557 };
13558 let authority_end = remainder.find(['/', '?', '#']).unwrap_or(remainder.len());
13559 let (authority, suffix) = remainder.split_at(authority_end);
13560 format!(
13561 "{}://{}{}",
13562 scheme.to_ascii_lowercase(),
13563 authority.to_ascii_lowercase(),
13564 suffix
13565 )
13566}
13567
13568fn render_concurrent_template(
13569 template: &str,
13570 user_input: &str,
13571 context_values: &std::collections::HashMap<String, serde_json::Value>,
13572) -> Result<String> {
13573 let mut env = minijinja::Environment::new();
13574 env.add_template("concurrent", template)
13575 .map_err(|e| AgentError::Other(format!("Concurrent template parse error: {}", e)))?;
13576
13577 let mut ctx = std::collections::BTreeMap::new();
13578 ctx.insert("user_input".to_string(), minijinja::Value::from(user_input));
13579
13580 let context_obj = minijinja::Value::from_serialize(context_values);
13582 ctx.insert("context".to_string(), context_obj);
13583
13584 let tmpl = env
13585 .get_template("concurrent")
13586 .map_err(|e| AgentError::Other(format!("Concurrent template error: {}", e)))?;
13587
13588 tmpl.render(minijinja::Value::from_serialize(&ctx))
13589 .map_err(|e| AgentError::Other(format!("Concurrent template render error: {}", e)))
13590}
13591
13592#[cfg(test)]
13593mod tests {
13594 use super::*;
13595 use crate::AgentBuilder;
13596 use ai_agents_core::{LLMChunk, LLMConfig, LLMError, LLMFeature, Tool};
13597 use ai_agents_llm::mock::MockLLMProvider;
13598 use ai_agents_skills::{SkillDefinition, SkillStep};
13599 use ai_agents_tools::{
13600 CalculatorTool, CopyPathTool, DeletePathTool, FileWriteTool, MovePathTool, ToolAliases,
13601 ToolDescriptor, ToolProvider, ToolProviderError, ToolProviderType, WebFetchResolver,
13602 WebFetchTool, WebFetchTransport, WebFetchTransportRequest, WebFetchTransportResponse,
13603 };
13604
13605 fn mock_with_response(response: &str) -> MockLLMProvider {
13606 let mut mock = MockLLMProvider::new("test");
13607 mock.set_response(response);
13608 mock
13609 }
13610
13611 fn mock_with_responses(responses: Vec<&str>) -> MockLLMProvider {
13612 let mut mock = MockLLMProvider::new("test");
13613 mock.set_responses(responses.into_iter().map(String::from).collect(), true);
13614 mock
13615 }
13616
13617 async fn collect_stream_events(
13619 agent: &RuntimeAgent,
13620 input: &str,
13621 ) -> (String, Vec<StreamChunk>, Option<AgentResponse>) {
13622 use futures::StreamExt;
13623 let mut events = agent.chat_stream_events(input).await.expect("stream opens");
13624 let mut content = String::new();
13625 let mut chunks = Vec::new();
13626 let mut final_response = None;
13627 while let Some(event) = events.next().await {
13628 match event {
13629 AgentStreamEvent::Chunk(chunk) => {
13630 if let StreamChunk::Content { text } = &chunk {
13631 content.push_str(text);
13632 }
13633 chunks.push(chunk);
13634 }
13635 AgentStreamEvent::Final(response) => final_response = Some(response),
13636 }
13637 }
13638 (content, chunks, final_response)
13639 }
13640
13641 fn metadata_keys(response: &AgentResponse) -> std::collections::BTreeSet<String> {
13642 response
13643 .metadata
13644 .as_ref()
13645 .map(|m| m.keys().cloned().collect())
13646 .unwrap_or_default()
13647 }
13648
13649 async fn assert_blocking_streaming_parity<F>(
13652 build: F,
13653 input: &str,
13654 ) -> (AgentResponse, AgentResponse, Vec<StreamChunk>)
13655 where
13656 F: Fn() -> RuntimeAgent,
13657 {
13658 let blocking_agent = build();
13659 let streaming_agent = build();
13660
13661 let blocking = blocking_agent
13662 .chat(input)
13663 .await
13664 .expect("blocking chat succeeds");
13665 let (_, chunks, final_response) = collect_stream_events(&streaming_agent, input).await;
13666 let streamed = final_response.expect("streaming must emit Final when blocking succeeds");
13667
13668 assert_eq!(
13669 blocking.content, streamed.content,
13670 "committed content differs"
13671 );
13672 assert_eq!(
13673 metadata_keys(&blocking),
13674 metadata_keys(&streamed),
13675 "metadata key sets differ"
13676 );
13677 assert_eq!(
13678 blocking.tool_calls.as_ref().map(Vec::len),
13679 streamed.tool_calls.as_ref().map(Vec::len),
13680 "tool call counts differ"
13681 );
13682 assert_eq!(
13683 blocking_agent.current_state(),
13684 streaming_agent.current_state(),
13685 "final states differ"
13686 );
13687 (blocking, streamed, chunks)
13688 }
13689
13690 fn signed_calculator_response(
13691 exchange_id: &str,
13692 call_id: &str,
13693 expression: &str,
13694 ) -> LLMResponse {
13695 let call = ToolCall {
13696 id: call_id.to_string(),
13697 name: "calculator".to_string(),
13698 arguments: serde_json::json!({"expression": expression}),
13699 };
13700 let state = ai_agents_core::NativeProviderState::new(
13701 exchange_id,
13702 "fixture",
13703 "native-tools",
13704 ai_agents_core::NativeProviderTarget::new("https://fixture.invalid/", "fixture-model")
13705 .unwrap(),
13706 serde_json::json!({
13707 "role": "model",
13708 "parts": [{
13709 "functionCall": {"name": "calculator", "args": {"expression": expression}},
13710 "thoughtSignature": format!("signature-{exchange_id}")
13711 }]
13712 }),
13713 vec![ai_agents_core::NativeCallBinding::new(call_id, 0).unwrap()],
13714 )
13715 .unwrap();
13716 LLMResponse::new("", FinishReason::ToolCall)
13717 .with_provider_state(state)
13718 .unwrap()
13719 .with_tool_calls(vec![call])
13720 .unwrap()
13721 }
13722
13723 struct TerminalHistoryProvider {
13724 calls: Arc<std::sync::atomic::AtomicU32>,
13725 }
13726
13727 struct DroppingSignedAssistantMemory {
13728 messages: RwLock<Vec<ChatMessage>>,
13729 }
13730
13731 struct DroppingEarlierSequentialMemory {
13732 messages: RwLock<Vec<ChatMessage>>,
13733 signed_seen: std::sync::atomic::AtomicUsize,
13734 }
13735
13736 #[async_trait]
13737 impl ai_agents_core::Memory for DroppingSignedAssistantMemory {
13738 async fn add_message(&self, message: ChatMessage) -> Result<()> {
13739 let signed = message.role == ai_agents_core::Role::Assistant
13740 && ai_agents_core::decode_native_tool_call_markers(&message.content)
13741 .map_err(|error| AgentError::LLM(error.to_string()))?
13742 .is_some_and(|batch| batch.provider_state().is_some());
13743 if !signed {
13744 self.messages.write().push(message);
13745 }
13746 Ok(())
13747 }
13748
13749 async fn get_messages(&self, limit: Option<usize>) -> Result<Vec<ChatMessage>> {
13750 let messages = self.messages.read();
13751 let start = limit
13752 .map(|limit| messages.len().saturating_sub(limit))
13753 .unwrap_or(0);
13754 Ok(messages[start..].to_vec())
13755 }
13756
13757 async fn clear(&self) -> Result<()> {
13758 self.messages.write().clear();
13759 Ok(())
13760 }
13761
13762 fn len(&self) -> usize {
13763 self.messages.read().len()
13764 }
13765
13766 async fn restore(&self, snapshot: ai_agents_core::MemorySnapshot) -> Result<()> {
13767 *self.messages.write() = snapshot.messages;
13768 Ok(())
13769 }
13770 }
13771
13772 #[async_trait]
13773 impl ai_agents_memory::Memory for DroppingSignedAssistantMemory {}
13774
13775 #[async_trait]
13776 impl ai_agents_core::Memory for DroppingEarlierSequentialMemory {
13777 async fn add_message(&self, message: ChatMessage) -> Result<()> {
13778 let signed = message.role == ai_agents_core::Role::Assistant
13779 && ai_agents_core::decode_native_tool_call_markers(&message.content)
13780 .map_err(|error| AgentError::LLM(error.to_string()))?
13781 .is_some_and(|batch| batch.provider_state().is_some());
13782 let mut messages = self.messages.write();
13783 if signed && self.signed_seen.fetch_add(1, Ordering::SeqCst) == 1 {
13784 messages.retain(|stored| !stored.content.contains("seq-call-1"));
13785 }
13786 messages.push(message);
13787 Ok(())
13788 }
13789
13790 async fn get_messages(&self, limit: Option<usize>) -> Result<Vec<ChatMessage>> {
13791 let messages = self.messages.read();
13792 let start = limit
13793 .map(|limit| messages.len().saturating_sub(limit))
13794 .unwrap_or(0);
13795 Ok(messages[start..].to_vec())
13796 }
13797
13798 async fn clear(&self) -> Result<()> {
13799 self.messages.write().clear();
13800 self.signed_seen.store(0, Ordering::SeqCst);
13801 Ok(())
13802 }
13803
13804 fn len(&self) -> usize {
13805 self.messages.read().len()
13806 }
13807
13808 async fn restore(&self, snapshot: ai_agents_core::MemorySnapshot) -> Result<()> {
13809 *self.messages.write() = snapshot.messages;
13810 self.signed_seen.store(0, Ordering::SeqCst);
13811 Ok(())
13812 }
13813 }
13814
13815 #[async_trait]
13816 impl ai_agents_memory::Memory for DroppingEarlierSequentialMemory {}
13817
13818 #[async_trait]
13819 impl LLMProvider for TerminalHistoryProvider {
13820 async fn complete(
13821 &self,
13822 _messages: &[ChatMessage],
13823 _config: Option<&LLMConfig>,
13824 ) -> std::result::Result<LLMResponse, LLMError> {
13825 self.calls.fetch_add(1, Ordering::SeqCst);
13826 Err(LLMError::Serialization(
13827 "native history integrity failure".to_string(),
13828 ))
13829 }
13830
13831 async fn complete_stream(
13832 &self,
13833 _messages: &[ChatMessage],
13834 _config: Option<&LLMConfig>,
13835 ) -> std::result::Result<
13836 Box<dyn Stream<Item = std::result::Result<LLMChunk, LLMError>> + Unpin + Send>,
13837 LLMError,
13838 > {
13839 Err(LLMError::Serialization(
13840 "native history integrity failure".to_string(),
13841 ))
13842 }
13843
13844 fn provider_name(&self) -> &str {
13845 "terminal-history"
13846 }
13847
13848 fn supports(&self, _feature: LLMFeature) -> bool {
13849 false
13850 }
13851
13852 fn is_terminal_error(&self, error: &LLMError) -> bool {
13853 matches!(error, LLMError::Serialization(_))
13854 }
13855 }
13856
13857 fn disambiguation_state_machine(
13859 state_enabled: Option<bool>,
13860 require_confirmation: bool,
13861 ) -> Arc<StateMachine> {
13862 let definition = ai_agents_state::StateDefinition {
13863 prompt: Some("Handle the resolved request.".to_string()),
13864 disambiguation: Some(ai_agents_disambiguation::StateDisambiguationOverride {
13865 enabled: state_enabled,
13866 require_confirmation,
13867 ..Default::default()
13868 }),
13869 ..Default::default()
13870 };
13871 let review = ai_agents_state::StateDefinition {
13872 prompt: Some("Review a fresh request.".to_string()),
13873 ..Default::default()
13874 };
13875 Arc::new(
13876 StateMachine::new(ai_agents_state::StateConfig {
13877 initial: "active".to_string(),
13878 states: std::collections::HashMap::from([
13879 ("active".to_string(), definition),
13880 ("review".to_string(), review),
13881 ]),
13882 global_transitions: Vec::new(),
13883 fallback: None,
13884 max_no_transition: None,
13885 regenerate_on_transition: true,
13886 })
13887 .unwrap(),
13888 )
13889 }
13890
13891 fn state_disambiguation_agent(
13893 responses: Vec<&str>,
13894 manager_enabled: bool,
13895 state_enabled: Option<bool>,
13896 require_confirmation: bool,
13897 ) -> (RuntimeAgent, MockLLMProvider) {
13898 state_disambiguation_agent_with_skills(
13899 responses,
13900 manager_enabled,
13901 state_enabled,
13902 require_confirmation,
13903 Vec::new(),
13904 )
13905 }
13906
13907 fn state_disambiguation_agent_with_skills(
13909 responses: Vec<&str>,
13910 manager_enabled: bool,
13911 state_enabled: Option<bool>,
13912 require_confirmation: bool,
13913 skills: Vec<SkillDefinition>,
13914 ) -> (RuntimeAgent, MockLLMProvider) {
13915 let mut mock = MockLLMProvider::new("state-confirmation");
13916 mock.set_responses(responses.into_iter().map(String::from).collect(), false);
13917 let observed = mock.clone();
13918 let agent = AgentBuilder::new()
13919 .system_prompt("Handle requests.")
13920 .llm(Arc::new(mock.clone()))
13921 .llm_alias("router", Arc::new(mock))
13922 .state_machine(disambiguation_state_machine(
13923 state_enabled,
13924 require_confirmation,
13925 ))
13926 .skills(skills)
13927 .build()
13928 .unwrap()
13929 .with_disambiguation(DisambiguationConfig {
13930 enabled: manager_enabled,
13931 ..Default::default()
13932 });
13933 (agent, observed)
13934 }
13935
13936 fn confirmation_skill() -> SkillDefinition {
13938 SkillDefinition {
13939 id: "send_report".to_string(),
13940 description: "Send a report after clarification".to_string(),
13941 trigger: "When the user asks to send a report".to_string(),
13942 steps: vec![SkillStep::Prompt {
13943 prompt: "Execute confirmed report skill for: {{ input }}".to_string(),
13944 llm: None,
13945 }],
13946 reasoning: None,
13947 reflection: None,
13948 disambiguation: Some(ai_agents_disambiguation::SkillDisambiguationOverride {
13949 enabled: Some(true),
13950 ..Default::default()
13951 }),
13952 }
13953 }
13954
13955 fn confirmation_skill_call_count(observed: &MockLLMProvider) -> usize {
13957 observed
13958 .call_history()
13959 .iter()
13960 .filter(|call| {
13961 call.messages
13962 .iter()
13963 .any(|message| message.content.contains("Execute confirmed report skill"))
13964 })
13965 .count()
13966 }
13967
13968 struct BlockingRuntimeConfirmationObserver {
13969 entered: tokio::sync::Barrier,
13970 release: tokio::sync::Notify,
13971 }
13972
13973 impl BlockingRuntimeConfirmationObserver {
13974 fn new() -> Self {
13975 Self {
13976 entered: tokio::sync::Barrier::new(2),
13977 release: tokio::sync::Notify::new(),
13978 }
13979 }
13980 }
13981
13982 struct ResetOnTransitionHooks {
13983 agent: parking_lot::Mutex<Option<Weak<RuntimeAgent>>>,
13984 invoked: AtomicBool,
13985 }
13986
13987 #[async_trait]
13988 impl AgentHooks for ResetOnTransitionHooks {
13989 async fn on_state_transition(&self, _from: Option<&str>, _to: &str, _reason: &str) {
13990 if self.invoked.swap(true, Ordering::SeqCst) {
13991 return;
13992 }
13993 let agent = self.agent.lock().as_ref().and_then(Weak::upgrade);
13994 if let Some(agent) = agent {
13995 agent.reset().await.unwrap();
13996 }
13997 }
13998 }
13999
14000 impl ClarificationObserver for BlockingRuntimeConfirmationObserver {
14001 fn observe_question<'a>(
14002 &'a self,
14003 future: ClarificationQuestionFuture<'a>,
14004 ) -> ClarificationQuestionFuture<'a> {
14005 future
14006 }
14007
14008 fn observe_parse<'a>(
14009 &'a self,
14010 future: ClarificationParseFuture<'a>,
14011 ) -> ClarificationParseFuture<'a> {
14012 future
14013 }
14014
14015 fn observe_confirmation_parse<'a>(
14016 &'a self,
14017 future: ConfirmationParseFuture<'a>,
14018 ) -> ConfirmationParseFuture<'a> {
14019 Box::pin(async move {
14020 self.entered.wait().await;
14021 self.release.notified().await;
14022 future.await
14023 })
14024 }
14025 }
14026
14027 #[tokio::test]
14028 async fn state_confirmation_blocks_redispatch_until_explicit_agreement() {
14029 let (agent, observed) = state_disambiguation_agent(
14030 vec![
14031 r#"{"is_ambiguous":true,"confidence":0.2,"ambiguity_type":"missing_target","reasoning":"target missing","what_is_unclear":["target"],"detected_language":"en"}"#,
14032 r#"{"question":"What should I send?","options":null}"#,
14033 r#"{"status":"answered","selected_option":null,"enriched_input":"Send the report to Ada","resolved":{"intent":"send_report"}}"#,
14034 r#"{"question":"Should I send the report to Ada?"}"#,
14035 r#"{"status":"confirmed"}"#,
14036 "Request executed.",
14037 ],
14038 true,
14039 None,
14040 true,
14041 );
14042
14043 let clarification = agent.chat("Send it").await.unwrap();
14044 assert_eq!(clarification.content, "What should I send?");
14045 assert_eq!(observed.call_count(), 2);
14046
14047 let confirmation = agent.chat("The report to Ada").await.unwrap();
14048 assert_eq!(confirmation.content, "Should I send the report to Ada?");
14049 assert_eq!(
14050 confirmation
14051 .metadata
14052 .as_ref()
14053 .and_then(|metadata| metadata.get("disambiguation"))
14054 .and_then(|metadata| metadata.get("status"))
14055 .and_then(Value::as_str),
14056 Some("awaiting_confirmation")
14057 );
14058 assert_eq!(observed.call_count(), 4);
14059
14060 let completed = agent.chat("Yes").await.unwrap();
14061 assert_eq!(completed.content, "Request executed.");
14062 assert_eq!(observed.call_count(), 6);
14063 }
14064
14065 #[tokio::test]
14066 async fn streaming_state_confirmation_ends_the_turn_before_redispatch() {
14067 let (agent, observed) = state_disambiguation_agent(
14068 vec![
14069 r#"{"is_ambiguous":true,"confidence":0.2,"ambiguity_type":"missing_target","reasoning":"target missing","what_is_unclear":["target"],"detected_language":"en"}"#,
14070 r#"{"question":"What should I send?","options":null}"#,
14071 r#"{"status":"answered","selected_option":null,"enriched_input":"Send the report to Ada","resolved":{"intent":"send_report"}}"#,
14072 r#"{"question":"Should I send the report to Ada?"}"#,
14073 r#"{"status":"confirmed"}"#,
14074 "Request executed.",
14075 ],
14076 true,
14077 None,
14078 true,
14079 );
14080
14081 let mut clarification_stream = agent.chat_stream("Send it").await.unwrap();
14082 let mut clarification = String::new();
14083 while let Some(chunk) = clarification_stream.next().await {
14084 match chunk {
14085 StreamChunk::Content { text } => clarification.push_str(&text),
14086 StreamChunk::Done {} => break,
14087 StreamChunk::Error { message } => panic!("unexpected stream error: {message}"),
14088 _ => {}
14089 }
14090 }
14091 assert_eq!(clarification, "What should I send?");
14092 assert_eq!(observed.call_count(), 2);
14093
14094 let mut confirmation_stream = agent.chat_stream_events("The report to Ada").await.unwrap();
14095 let mut confirmation = None;
14096 while let Some(event) = confirmation_stream.next().await {
14097 match event {
14098 AgentStreamEvent::Final(response) => confirmation = Some(response),
14099 AgentStreamEvent::Chunk(StreamChunk::Error { message }) => {
14100 panic!("unexpected stream error: {message}")
14101 }
14102 AgentStreamEvent::Chunk(_) => {}
14103 }
14104 }
14105 let confirmation = confirmation.expect("confirmation must finalize");
14106 assert_eq!(confirmation.content, "Should I send the report to Ada?");
14107 assert_eq!(
14108 confirmation
14109 .metadata
14110 .as_ref()
14111 .and_then(|metadata| metadata.get("disambiguation"))
14112 .and_then(|metadata| metadata.get("status"))
14113 .and_then(Value::as_str),
14114 Some("awaiting_confirmation")
14115 );
14116 assert_eq!(observed.call_count(), 4);
14117
14118 let mut completed_stream = agent.chat_stream("Yes").await.unwrap();
14119 let mut completed = String::new();
14120 while let Some(chunk) = completed_stream.next().await {
14121 match chunk {
14122 StreamChunk::Content { text } => completed.push_str(&text),
14123 StreamChunk::Done {} => break,
14124 StreamChunk::Error { message } => panic!("unexpected stream error: {message}"),
14125 _ => {}
14126 }
14127 }
14128 assert_eq!(completed, "Request executed.");
14129 assert_eq!(observed.call_count(), 6);
14130 }
14131
14132 #[tokio::test]
14134 async fn root_turn_gate_serializes_blocking_and_streaming_entry_points() {
14135 let (complete_entered, mut complete_events) = tokio::sync::mpsc::unbounded_channel();
14136 let agent = Arc::new(
14137 AgentBuilder::new()
14138 .system_prompt("Serialize root turns.")
14139 .llm(Arc::new(RootTurnProbeProvider { complete_entered }))
14140 .build()
14141 .unwrap(),
14142 );
14143 let blocking_agent = Arc::clone(&agent);
14144
14145 let legacy_stream = agent.chat_stream("stream owner").await.unwrap();
14146 assert!(agent.root_turn_gate.try_lock().is_err());
14147 let blocking = tokio::spawn(async move { blocking_agent.chat("blocked").await.unwrap() });
14148 assert!(
14149 tokio::time::timeout(std::time::Duration::from_millis(50), complete_events.recv())
14150 .await
14151 .is_err(),
14152 "blocking turn reached the provider while the legacy stream owned the root gate"
14153 );
14154
14155 drop(legacy_stream);
14156 assert_eq!(
14157 tokio::time::timeout(std::time::Duration::from_secs(2), complete_events.recv())
14158 .await
14159 .expect("blocking turn did not enter after stream drop"),
14160 Some(())
14161 );
14162 let response = tokio::time::timeout(std::time::Duration::from_secs(2), blocking)
14163 .await
14164 .expect("blocking turn did not finish after stream drop")
14165 .unwrap();
14166 assert_eq!(response.content, "blocking complete");
14167
14168 let mut event_stream = agent.chat_stream_events("event terminal").await.unwrap();
14169 assert!(agent.root_turn_gate.try_lock().is_err());
14170 let mut saw_final = false;
14171 while let Some(event) = event_stream.next().await {
14172 if matches!(event, AgentStreamEvent::Final(_)) {
14173 saw_final = true;
14174 break;
14175 }
14176 }
14177 assert!(saw_final);
14178 assert!(
14179 agent.root_turn_gate.try_lock().is_ok(),
14180 "authoritative terminal event retained the root gate"
14181 );
14182 }
14183
14184 #[tokio::test]
14186 async fn response_hook_rejects_same_runtime_chat_reentry() {
14187 let hooks = Arc::new(ResponseChatHooks {
14188 target: parking_lot::Mutex::new(None),
14189 invoked: AtomicBool::new(false),
14190 nested_result: parking_lot::Mutex::new(None),
14191 });
14192 let agent = Arc::new(
14193 AgentBuilder::new()
14194 .system_prompt("Reject response hook reentry.")
14195 .llm(Arc::new(mock_with_response("outer response")))
14196 .hooks(hooks.clone())
14197 .build()
14198 .unwrap(),
14199 );
14200 *hooks.target.lock() = Some(Arc::downgrade(&agent));
14201
14202 let response = tokio::time::timeout(
14203 std::time::Duration::from_secs(2),
14204 agent.chat("outer request"),
14205 )
14206 .await
14207 .expect("same-runtime response hook reentry must fail without deadlocking")
14208 .unwrap();
14209
14210 assert_eq!(response.content, "outer response");
14211 let nested_result = hooks
14212 .nested_result
14213 .lock()
14214 .clone()
14215 .expect("response hook must record its nested call");
14216 let error = nested_result.expect_err("same-runtime nested chat must be rejected");
14217 assert!(error.contains("reentrant root turn ownership"));
14218 }
14219
14220 #[tokio::test]
14222 async fn root_turn_gate_allows_nested_runtime_and_rejects_cycles() {
14223 let agent_a = AgentBuilder::new()
14224 .system_prompt("Runtime A.")
14225 .llm(Arc::new(mock_with_response("response A")))
14226 .build()
14227 .unwrap();
14228 let agent_b = AgentBuilder::new()
14229 .system_prompt("Runtime B.")
14230 .llm(Arc::new(mock_with_response("response B")))
14231 .build()
14232 .unwrap();
14233 let RootTurnAdmission {
14234 guard: guard_a,
14235 identity_stack: stack_a,
14236 } = agent_a.acquire_root_turn().await.unwrap();
14237
14238 let cycle_error = scope_runtime_gate_identity_stack(&stack_a, async {
14239 let RootTurnAdmission {
14240 guard: guard_b,
14241 identity_stack: stack_b,
14242 } = agent_b
14243 .acquire_root_turn()
14244 .await
14245 .expect("runtime B must acquire a different gate");
14246 let result =
14247 scope_runtime_gate_identity_stack(&stack_b, agent_a.acquire_root_turn()).await;
14248 drop(guard_b);
14249 match result {
14250 Err(error) => error,
14251 Ok(_) => panic!("runtime A accepted a repeated gate identity"),
14252 }
14253 })
14254 .await;
14255 drop(guard_a);
14256
14257 assert!(
14258 cycle_error
14259 .to_string()
14260 .contains("reentrant root turn ownership")
14261 );
14262 }
14263
14264 #[tokio::test]
14266 async fn concurrent_orchestration_propagates_root_gate_ancestry() {
14267 let registry = Arc::new(crate::spawner::AgentRegistry::new());
14268 let hooks_a = Arc::new(ConcurrentResponseHooks {
14269 registry: Arc::downgrade(®istry),
14270 child_id: "runtime-b".to_string(),
14271 invoked: AtomicBool::new(false),
14272 nested_result: parking_lot::Mutex::new(None),
14273 });
14274 let hooks_b = Arc::new(ResponseChatHooks {
14275 target: parking_lot::Mutex::new(None),
14276 invoked: AtomicBool::new(false),
14277 nested_result: parking_lot::Mutex::new(None),
14278 });
14279 let agent_a = AgentBuilder::new()
14280 .system_prompt("Runtime A dispatches runtime B concurrently.")
14281 .llm(Arc::new(mock_with_response("response A")))
14282 .hooks(hooks_a.clone())
14283 .build()
14284 .unwrap();
14285 let agent_b = AgentBuilder::new()
14286 .system_prompt("Runtime B attempts to re-enter runtime A.")
14287 .llm(Arc::new(mock_with_response("response B")))
14288 .hooks(hooks_b.clone())
14289 .build()
14290 .unwrap();
14291 let spec_a = crate::spec::AgentSpec {
14292 name: "runtime-a".to_string(),
14293 system_prompt: "Runtime A dispatches runtime B concurrently.".to_string(),
14294 ..crate::spec::AgentSpec::default()
14295 };
14296 let spec_b = crate::spec::AgentSpec {
14297 name: "runtime-b".to_string(),
14298 system_prompt: "Runtime B attempts to re-enter runtime A.".to_string(),
14299 ..crate::spec::AgentSpec::default()
14300 };
14301 registry
14302 .register(crate::spawner::SpawnedAgent::from_runtime(
14303 "runtime-a".to_string(),
14304 agent_a,
14305 spec_a,
14306 ))
14307 .await
14308 .unwrap();
14309 registry
14310 .register(crate::spawner::SpawnedAgent::from_runtime(
14311 "runtime-b".to_string(),
14312 agent_b,
14313 spec_b,
14314 ))
14315 .await
14316 .unwrap();
14317 let runtime_a = registry.get("runtime-a").unwrap();
14318 *hooks_b.target.lock() = Some(Arc::downgrade(&runtime_a));
14319
14320 let response = tokio::time::timeout(
14321 std::time::Duration::from_secs(2),
14322 runtime_a.chat("outer concurrent request"),
14323 )
14324 .await
14325 .expect("concurrent orchestration cycle must fail without deadlocking")
14326 .unwrap();
14327
14328 assert_eq!(response.content, "response A");
14329 let child_result = hooks_a
14330 .nested_result
14331 .lock()
14332 .clone()
14333 .expect("runtime A hook must record runtime B completion");
14334 assert_eq!(child_result.unwrap(), "response B");
14335 let cycle_result = hooks_b
14336 .nested_result
14337 .lock()
14338 .clone()
14339 .expect("runtime B hook must record runtime A reentry");
14340 assert!(
14341 cycle_result
14342 .expect_err("runtime A accepted a repeated gate identity")
14343 .contains("reentrant root turn ownership")
14344 );
14345 }
14346
14347 fn skill_clarification_responses() -> Vec<&'static str> {
14350 vec![
14351 r#"{"is_ambiguous":false,"confidence":0.99,"ambiguity_type":null,"reasoning":"top-level clear","what_is_unclear":[],"detected_language":"en"}"#,
14352 "send_report",
14353 r#"{"is_ambiguous":true,"confidence":0.2,"ambiguity_type":"missing_target","reasoning":"target missing","what_is_unclear":["target"],"detected_language":"en"}"#,
14354 r#"{"question":"What should I send?","options":null}"#,
14355 ]
14356 }
14357
14358 #[tokio::test]
14364 async fn test_stream_skill_clarification_memory_matches_blocking() {
14365 let (blocking_agent, _) = state_disambiguation_agent_with_skills(
14366 skill_clarification_responses(),
14367 true,
14368 None,
14369 true,
14370 vec![confirmation_skill()],
14371 );
14372 let blocking = blocking_agent.chat("Send it").await.unwrap();
14373 let blocking_messages = blocking_agent.memory.get_messages(None).await.unwrap();
14374
14375 let (streaming_agent, _) = state_disambiguation_agent_with_skills(
14376 skill_clarification_responses(),
14377 true,
14378 None,
14379 true,
14380 vec![confirmation_skill()],
14381 );
14382 let (content, chunks, streamed) = collect_stream_events(&streaming_agent, "Send it").await;
14383 assert!(!chunks.iter().any(StreamChunk::is_error), "{chunks:?}");
14384 let streamed = streamed.expect("skill clarification must finalize as Final");
14385 let streaming_messages = streaming_agent.memory.get_messages(None).await.unwrap();
14386
14387 assert_eq!(blocking.content, "What should I send?");
14388 assert_eq!(streamed.content, blocking.content);
14389 assert_eq!(content, streamed.content);
14390 assert_eq!(
14391 blocking
14392 .metadata
14393 .as_ref()
14394 .and_then(|m| m.get("disambiguation")),
14395 streamed
14396 .metadata
14397 .as_ref()
14398 .and_then(|m| m.get("disambiguation")),
14399 );
14400 assert_eq!(
14401 streamed
14402 .metadata
14403 .as_ref()
14404 .and_then(|m| m.get("disambiguation"))
14405 .and_then(|d| d.get("status"))
14406 .and_then(Value::as_str),
14407 Some("awaiting_clarification"),
14408 );
14409 let shape = |messages: &[ChatMessage]| {
14410 messages
14411 .iter()
14412 .map(|m| (format!("{:?}", m.role), m.content.clone()))
14413 .collect::<Vec<_>>()
14414 };
14415 assert_eq!(shape(&blocking_messages), shape(&streaming_messages));
14416 assert_eq!(
14417 shape(&streaming_messages),
14418 vec![
14419 ("User".to_string(), "Send it".to_string()),
14420 ("Assistant".to_string(), "What should I send?".to_string()),
14421 ],
14422 );
14423 assert_eq!(
14424 *streaming_agent.pending_skill_id.read(),
14425 Some("send_report".to_string()),
14426 );
14427 }
14428
14429 #[tokio::test]
14431 async fn test_stream_skill_clarification_memory_failure_surfaces_as_error() {
14432 let build = || {
14434 let mut mock = MockLLMProvider::new("skill-clarification");
14435 mock.set_responses(
14436 skill_clarification_responses()
14437 .into_iter()
14438 .map(String::from)
14439 .collect(),
14440 false,
14441 );
14442 AgentBuilder::new()
14443 .system_prompt("Handle requests.")
14444 .llm(Arc::new(mock.clone()))
14445 .llm_alias("router", Arc::new(mock))
14446 .state_machine(disambiguation_state_machine(None, true))
14447 .skills(vec![confirmation_skill()])
14448 .memory(Arc::new(FailingMemory {
14449 messages: parking_lot::RwLock::new(Vec::new()),
14450 fail_on_add: 2,
14451 adds: std::sync::atomic::AtomicUsize::new(0),
14452 }))
14453 .build()
14454 .unwrap()
14455 .with_disambiguation(DisambiguationConfig {
14456 enabled: true,
14457 ..Default::default()
14458 })
14459 };
14460
14461 let blocking = build().chat("Send it").await;
14462 assert!(
14463 blocking.is_err(),
14464 "blocking must surface the failed clarification write: {blocking:?}"
14465 );
14466
14467 let (_, chunks, streamed) = collect_stream_events(&build(), "Send it").await;
14468 assert!(
14469 streamed.is_none(),
14470 "a failed write must not finalize the turn"
14471 );
14472 assert!(
14473 chunks.iter().any(|chunk| matches!(
14474 chunk,
14475 StreamChunk::Error { message } if message.contains("simulated memory failure")
14476 )),
14477 "streaming must surface the failed clarification write: {chunks:?}"
14478 );
14479 }
14480
14481 #[tokio::test]
14483 async fn confirmed_skill_route_executes_exactly_once() {
14484 let (agent, observed) = state_disambiguation_agent_with_skills(
14485 vec![
14486 r#"{"is_ambiguous":false,"confidence":0.99,"ambiguity_type":null,"reasoning":"top-level clear","what_is_unclear":[],"detected_language":"en"}"#,
14487 "send_report",
14488 r#"{"is_ambiguous":true,"confidence":0.2,"ambiguity_type":"missing_target","reasoning":"target missing","what_is_unclear":["target"],"detected_language":"en"}"#,
14489 r#"{"question":"What should I send?","options":null}"#,
14490 r#"{"status":"answered","selected_option":null,"enriched_input":"Send the report to Ada","resolved":{"intent":"send_report"}}"#,
14491 r#"{"question":"Should I send the report to Ada?"}"#,
14492 r#"{"status":"confirmed"}"#,
14493 r#"{"is_ambiguous":false,"confidence":0.99,"ambiguity_type":null,"reasoning":"resolved","what_is_unclear":[],"detected_language":"en"}"#,
14494 "Report skill executed.",
14495 ],
14496 true,
14497 None,
14498 true,
14499 vec![confirmation_skill()],
14500 );
14501
14502 let clarification = agent.chat("Send it").await.unwrap();
14503 assert_eq!(clarification.content, "What should I send?");
14504 assert_eq!(confirmation_skill_call_count(&observed), 0);
14505
14506 let confirmation = agent.chat("The report to Ada").await.unwrap();
14507 assert_eq!(confirmation.content, "Should I send the report to Ada?");
14508 assert_eq!(
14509 confirmation
14510 .metadata
14511 .as_ref()
14512 .and_then(|metadata| metadata.get("disambiguation"))
14513 .and_then(|metadata| metadata.get("status"))
14514 .and_then(Value::as_str),
14515 Some("awaiting_confirmation")
14516 );
14517 assert_eq!(confirmation_skill_call_count(&observed), 0);
14518
14519 let completed = agent.chat("Yes").await.unwrap();
14520 assert_eq!(completed.content, "Report skill executed.");
14521 assert_eq!(confirmation_skill_call_count(&observed), 1);
14522 assert!(agent.pending_skill_id.read().is_none());
14523 let messages = agent.memory.get_messages(None).await.unwrap();
14524 assert!(!messages.iter().any(|message| message.content == "Yes"));
14525 }
14526
14527 #[tokio::test]
14529 async fn confirmed_skill_recheck_preserves_new_clarification_metadata() {
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 r#"{"status":"confirmed"}"#,
14539 r#"{"is_ambiguous":true,"confidence":0.3,"ambiguity_type":"missing_parameters","reasoning":"timing missing","what_is_unclear":["timing"],"detected_language":"en"}"#,
14540 r#"{"question":"When should I send it?","options":null}"#,
14541 ],
14542 true,
14543 None,
14544 true,
14545 vec![confirmation_skill()],
14546 );
14547
14548 agent.chat("Send it").await.unwrap();
14549 agent.chat("The report to Ada").await.unwrap();
14550 let follow_up = agent.chat("Yes").await.unwrap();
14551
14552 assert_eq!(follow_up.content, "When should I send it?");
14553 let metadata = follow_up
14554 .metadata
14555 .as_ref()
14556 .and_then(|metadata| metadata.get("disambiguation"))
14557 .unwrap();
14558 assert_eq!(
14559 metadata.get("status").and_then(Value::as_str),
14560 Some("awaiting_clarification")
14561 );
14562 assert_eq!(
14563 metadata.get("skill_id").and_then(Value::as_str),
14564 Some("send_report")
14565 );
14566 assert!(metadata.get("detection").is_some());
14567 assert_eq!(confirmation_skill_call_count(&observed), 0);
14568 }
14569
14570 #[tokio::test]
14572 async fn rejected_skill_confirmation_never_executes() {
14573 let (agent, observed) = state_disambiguation_agent_with_skills(
14574 vec![
14575 r#"{"is_ambiguous":false,"confidence":0.99,"ambiguity_type":null,"reasoning":"top-level clear","what_is_unclear":[],"detected_language":"en"}"#,
14576 "send_report",
14577 r#"{"is_ambiguous":true,"confidence":0.2,"ambiguity_type":"missing_target","reasoning":"target missing","what_is_unclear":["target"],"detected_language":"en"}"#,
14578 r#"{"question":"What should I send?","options":null}"#,
14579 r#"{"status":"answered","selected_option":null,"enriched_input":"Send the report to Ada","resolved":{"intent":"send_report"}}"#,
14580 r#"{"question":"Should I send the report to Ada?"}"#,
14581 r#"{"status":"rejected"}"#,
14582 "Confirmation rejected.",
14583 ],
14584 true,
14585 None,
14586 true,
14587 vec![confirmation_skill()],
14588 );
14589
14590 agent.chat("Send it").await.unwrap();
14591 agent.chat("The report to Ada").await.unwrap();
14592 let rejected = agent.chat("No").await.unwrap();
14593
14594 assert_eq!(rejected.content, "Confirmation rejected.");
14595 assert_eq!(confirmation_skill_call_count(&observed), 0);
14596 assert!(agent.pending_skill_id.read().is_none());
14597 }
14598
14599 #[tokio::test]
14601 async fn reset_invalidates_pending_skill_confirmation_before_streaming_input() {
14602 let (agent, observed) = state_disambiguation_agent_with_skills(
14603 vec![
14604 r#"{"is_ambiguous":false,"confidence":0.99,"ambiguity_type":null,"reasoning":"top-level clear","what_is_unclear":[],"detected_language":"en"}"#,
14605 "send_report",
14606 r#"{"is_ambiguous":true,"confidence":0.2,"ambiguity_type":"missing_target","reasoning":"target missing","what_is_unclear":["target"],"detected_language":"en"}"#,
14607 r#"{"question":"What should I send?","options":null}"#,
14608 r#"{"status":"answered","selected_option":null,"enriched_input":"Send the report to Ada","resolved":{"intent":"send_report"}}"#,
14609 r#"{"question":"Should I send the report to Ada?"}"#,
14610 r#"{"is_ambiguous":false,"confidence":0.99,"ambiguity_type":null,"reasoning":"fresh input","what_is_unclear":[],"detected_language":"en"}"#,
14611 "none",
14612 "Fresh response.",
14613 ],
14614 true,
14615 None,
14616 true,
14617 vec![confirmation_skill()],
14618 );
14619
14620 agent.chat("Send it").await.unwrap();
14621 agent.chat("The report to Ada").await.unwrap();
14622 agent.reset().await.unwrap();
14623 assert!(agent.pending_skill_id.read().is_none());
14624 assert!(
14625 !agent
14626 .disambiguation_manager()
14627 .unwrap()
14628 .has_pending_clarification()
14629 .await
14630 );
14631
14632 let mut stream = agent.chat_stream("Yes").await.unwrap();
14633 let mut content = String::new();
14634 while let Some(chunk) = stream.next().await {
14635 match chunk {
14636 StreamChunk::Content { text } => content.push_str(&text),
14637 StreamChunk::Done {} => break,
14638 StreamChunk::Error { message } => panic!("unexpected stream error: {message}"),
14639 _ => {}
14640 }
14641 }
14642
14643 assert_eq!(content, "Fresh response.");
14644 assert_eq!(confirmation_skill_call_count(&observed), 0);
14645 }
14646
14647 #[tokio::test]
14649 async fn trait_reset_clears_pending_skill_confirmation() {
14650 let (agent, _) = state_disambiguation_agent_with_skills(
14651 vec![
14652 r#"{"is_ambiguous":false,"confidence":0.99,"ambiguity_type":null,"reasoning":"top-level clear","what_is_unclear":[],"detected_language":"en"}"#,
14653 "send_report",
14654 r#"{"is_ambiguous":true,"confidence":0.2,"ambiguity_type":"missing_target","reasoning":"target missing","what_is_unclear":["target"],"detected_language":"en"}"#,
14655 r#"{"question":"What should I send?","options":null}"#,
14656 r#"{"status":"answered","selected_option":null,"enriched_input":"Send the report to Ada","resolved":{"intent":"send_report"}}"#,
14657 r#"{"question":"Should I send the report to Ada?"}"#,
14658 ],
14659 true,
14660 None,
14661 true,
14662 vec![confirmation_skill()],
14663 );
14664
14665 agent.chat("Send it").await.unwrap();
14666 agent.chat("The report to Ada").await.unwrap();
14667 <RuntimeAgent as Agent>::reset(&agent).await.unwrap();
14668
14669 assert!(agent.pending_skill_id.read().is_none());
14670 assert!(
14671 !agent
14672 .disambiguation_manager()
14673 .unwrap()
14674 .has_pending_clarification()
14675 .await
14676 );
14677 }
14678
14679 #[tokio::test]
14681 async fn state_change_invalidates_pending_skill_confirmation() {
14682 let (agent, observed) = state_disambiguation_agent_with_skills(
14683 vec![
14684 r#"{"is_ambiguous":false,"confidence":0.99,"ambiguity_type":null,"reasoning":"top-level clear","what_is_unclear":[],"detected_language":"en"}"#,
14685 "send_report",
14686 r#"{"is_ambiguous":true,"confidence":0.2,"ambiguity_type":"missing_target","reasoning":"target missing","what_is_unclear":["target"],"detected_language":"en"}"#,
14687 r#"{"question":"What should I send?","options":null}"#,
14688 r#"{"status":"answered","selected_option":null,"enriched_input":"Send the report to Ada","resolved":{"intent":"send_report"}}"#,
14689 r#"{"question":"Should I send the report to Ada?"}"#,
14690 r#"{"is_ambiguous":false,"confidence":0.99,"ambiguity_type":null,"reasoning":"fresh input","what_is_unclear":[],"detected_language":"en"}"#,
14691 "none",
14692 "Fresh response.",
14693 ],
14694 true,
14695 None,
14696 true,
14697 vec![confirmation_skill()],
14698 );
14699
14700 agent.chat("Send it").await.unwrap();
14701 agent.chat("The report to Ada").await.unwrap();
14702 agent.transition_to("review").await.unwrap();
14703 let cancelled = agent.chat("Yes").await.unwrap();
14704
14705 assert_eq!(cancelled.content, "Fresh response.");
14706 assert_eq!(confirmation_skill_call_count(&observed), 0);
14707 assert!(agent.pending_skill_id.read().is_none());
14708 }
14709
14710 #[tokio::test]
14712 async fn in_flight_confirmation_cannot_redispatch_after_reset() {
14713 let (mut agent, observed) = state_disambiguation_agent_with_skills(
14714 vec![
14715 r#"{"is_ambiguous":false,"confidence":0.99,"ambiguity_type":null,"reasoning":"top-level clear","what_is_unclear":[],"detected_language":"en"}"#,
14716 "send_report",
14717 r#"{"is_ambiguous":true,"confidence":0.2,"ambiguity_type":"missing_target","reasoning":"target missing","what_is_unclear":["target"],"detected_language":"en"}"#,
14718 r#"{"question":"What should I send?","options":null}"#,
14719 r#"{"status":"answered","selected_option":null,"enriched_input":"Send the report to Ada","resolved":{"intent":"send_report"}}"#,
14720 r#"{"question":"Should I send the report to Ada?"}"#,
14721 r#"{"status":"confirmed"}"#,
14722 "Confirmation cancelled.",
14723 ],
14724 true,
14725 None,
14726 true,
14727 vec![confirmation_skill()],
14728 );
14729 let observer = Arc::new(BlockingRuntimeConfirmationObserver::new());
14730 let manager = agent
14731 .disambiguation_manager
14732 .take()
14733 .unwrap()
14734 .with_clarification_observer(observer.clone());
14735 agent.disambiguation_manager = Some(manager);
14736 let agent = Arc::new(agent);
14737
14738 agent.chat("Send it").await.unwrap();
14739 agent.chat("The report to Ada").await.unwrap();
14740
14741 let confirming_agent = Arc::clone(&agent);
14742 let confirmation = tokio::spawn(async move { confirming_agent.chat("Yes").await });
14743 observer.entered.wait().await;
14744 agent.reset().await.unwrap();
14745 observer.release.notify_one();
14746
14747 let response = confirmation.await.unwrap().unwrap();
14748 assert_eq!(response.content, "Confirmation cancelled.");
14749 assert_eq!(confirmation_skill_call_count(&observed), 0);
14750 assert!(agent.pending_skill_id.read().is_none());
14751 }
14752
14753 #[tokio::test]
14755 async fn queued_reset_prevents_stale_confirmation_question_publication() {
14756 let (agent, observed) = state_disambiguation_agent(
14757 vec![
14758 r#"{"is_ambiguous":true,"confidence":0.2,"ambiguity_type":"missing_target","reasoning":"target missing","what_is_unclear":["target"],"detected_language":"en"}"#,
14759 r#"{"question":"What should I send?","options":null}"#,
14760 r#"{"status":"answered","selected_option":null,"enriched_input":"Send the report to Ada","resolved":{"intent":"send_report"}}"#,
14761 r#"{"question":"Should I send the report to Ada?"}"#,
14762 ],
14763 true,
14764 None,
14765 true,
14766 );
14767 let agent = Arc::new(agent);
14768 agent.chat("Send it").await.unwrap();
14769
14770 let admission = agent.disambiguation_admission.write().await;
14771 let (started_tx, started_rx) = tokio::sync::oneshot::channel();
14772 let resetting_agent = Arc::clone(&agent);
14773 let reset = tokio::spawn(async move {
14774 let _ = started_tx.send(());
14775 resetting_agent.reset().await
14776 });
14777 started_rx.await.unwrap();
14778 tokio::task::yield_now().await;
14779
14780 let responding_agent = Arc::clone(&agent);
14781 let response =
14782 tokio::spawn(async move { responding_agent.chat("The report to Ada").await });
14783 tokio::time::timeout(std::time::Duration::from_secs(2), async {
14784 while observed.call_count() < 4 {
14785 tokio::task::yield_now().await;
14786 }
14787 })
14788 .await
14789 .expect("clarification processing must reach terminal publication");
14790 drop(admission);
14791
14792 reset.await.unwrap().unwrap();
14793 let error = response.await.unwrap().unwrap_err();
14794 assert!(error.to_string().contains("ownership changed"));
14795 assert!(
14796 !agent
14797 .disambiguation_manager()
14798 .unwrap()
14799 .has_pending_clarification()
14800 .await
14801 );
14802 assert!(agent.memory.get_messages(None).await.unwrap().is_empty());
14803 }
14804
14805 #[tokio::test]
14807 async fn queued_reset_prevents_stale_skill_clarification_publication() {
14808 let (agent, observed) = state_disambiguation_agent_with_skills(
14809 vec![
14810 r#"{"is_ambiguous":false,"confidence":0.99,"ambiguity_type":null,"reasoning":"top-level clear","what_is_unclear":[],"detected_language":"en"}"#,
14811 "send_report",
14812 r#"{"is_ambiguous":true,"confidence":0.2,"ambiguity_type":"missing_target","reasoning":"target missing","what_is_unclear":["target"],"detected_language":"en"}"#,
14813 r#"{"question":"What should I send?","options":null}"#,
14814 ],
14815 true,
14816 None,
14817 true,
14818 vec![confirmation_skill()],
14819 );
14820 let agent = Arc::new(agent);
14821 let admission = agent.disambiguation_admission.write().await;
14822 let (started_tx, started_rx) = tokio::sync::oneshot::channel();
14823 let resetting_agent = Arc::clone(&agent);
14824 let reset = tokio::spawn(async move {
14825 let _ = started_tx.send(());
14826 resetting_agent.reset().await
14827 });
14828 started_rx.await.unwrap();
14829 tokio::task::yield_now().await;
14830
14831 let responding_agent = Arc::clone(&agent);
14832 let response = tokio::spawn(async move { responding_agent.chat("Send it").await });
14833 tokio::time::timeout(std::time::Duration::from_secs(2), async {
14834 while observed.call_count() < 4 {
14835 tokio::task::yield_now().await;
14836 }
14837 })
14838 .await
14839 .expect("skill clarification must reach terminal publication");
14840 drop(admission);
14841
14842 reset.await.unwrap().unwrap();
14843 let error = response.await.unwrap().unwrap_err();
14844 assert!(error.to_string().contains("ownership changed"));
14845 assert_eq!(confirmation_skill_call_count(&observed), 0);
14846 assert!(agent.pending_skill_id.read().is_none());
14847 assert!(agent.memory.get_messages(None).await.unwrap().is_empty());
14848 }
14849
14850 #[tokio::test]
14852 async fn transition_hook_can_reset_without_admission_deadlock() {
14853 let hooks = Arc::new(ResetOnTransitionHooks {
14854 agent: parking_lot::Mutex::new(None),
14855 invoked: AtomicBool::new(false),
14856 });
14857 let agent = Arc::new(
14858 AgentBuilder::new()
14859 .system_prompt("Test transition hook reentrancy.")
14860 .llm(Arc::new(mock_with_response("done")))
14861 .state_machine(disambiguation_state_machine(None, false))
14862 .build()
14863 .unwrap()
14864 .with_hooks(hooks.clone()),
14865 );
14866 *hooks.agent.lock() = Some(Arc::downgrade(&agent));
14867
14868 let transitioned = tokio::time::timeout(
14869 std::time::Duration::from_secs(2),
14870 agent.apply_transition_target("active", "review", "test transition", None),
14871 )
14872 .await
14873 .expect("transition hook reset must not deadlock")
14874 .unwrap();
14875
14876 assert!(transitioned);
14877 assert!(hooks.invoked.load(Ordering::SeqCst));
14878 assert_eq!(agent.current_state().as_deref(), Some("active"));
14879 }
14880
14881 #[tokio::test]
14883 async fn concurrent_transition_cannot_duplicate_exit_actions() {
14884 let gate = PathMutationGate::new();
14885 let active = ai_agents_state::StateDefinition {
14886 on_exit: vec![StateAction::Tool {
14887 tool: "transition_exit".to_string(),
14888 args: Some(serde_json::json!({"path": "./transition-exit.txt"})),
14889 }],
14890 ..Default::default()
14891 };
14892 let state_machine = Arc::new(
14893 StateMachine::new(ai_agents_state::StateConfig {
14894 initial: "active".to_string(),
14895 states: HashMap::from([
14896 ("active".to_string(), active),
14897 (
14898 "review".to_string(),
14899 ai_agents_state::StateDefinition::default(),
14900 ),
14901 ]),
14902 global_transitions: Vec::new(),
14903 fallback: None,
14904 max_no_transition: None,
14905 regenerate_on_transition: true,
14906 })
14907 .unwrap(),
14908 );
14909 let agent = Arc::new(
14910 AgentBuilder::new()
14911 .system_prompt("Test transition reservation.")
14912 .llm(Arc::new(mock_with_response("done")))
14913 .tool(Arc::new(BlockingPathMutationTool {
14914 id: "transition_exit",
14915 path_fields: vec![ai_agents_core::PathPolicyBinding::write("path")],
14916 gate: gate.clone(),
14917 }))
14918 .state_machine(state_machine)
14919 .build()
14920 .unwrap(),
14921 );
14922
14923 let first_agent = Arc::clone(&agent);
14924 let first = tokio::spawn(async move { first_agent.transition_to("review").await });
14925 tokio::time::timeout(std::time::Duration::from_secs(2), gate.wait_until_entered())
14926 .await
14927 .expect("reserved transition must enter its exit action");
14928
14929 let second = tokio::time::timeout(
14930 std::time::Duration::from_secs(2),
14931 agent.transition_to("review"),
14932 )
14933 .await
14934 .expect("competing transition must fail without waiting for the exit action")
14935 .unwrap_err();
14936 assert!(second.to_string().contains("already in progress"));
14937
14938 gate.release();
14939 first.await.unwrap().unwrap();
14940 assert_eq!(agent.current_state().as_deref(), Some("review"));
14941 }
14942
14943 #[tokio::test]
14945 async fn concurrent_transition_cannot_overtake_enter_actions() {
14946 let gate = PathMutationGate::new();
14947 let review = ai_agents_state::StateDefinition {
14948 on_enter: vec![StateAction::Tool {
14949 tool: "transition_enter".to_string(),
14950 args: Some(serde_json::json!({"path": "./transition-enter.txt"})),
14951 }],
14952 ..Default::default()
14953 };
14954 let state_machine = Arc::new(
14955 StateMachine::new(ai_agents_state::StateConfig {
14956 initial: "active".to_string(),
14957 states: HashMap::from([
14958 (
14959 "active".to_string(),
14960 ai_agents_state::StateDefinition::default(),
14961 ),
14962 ("review".to_string(), review),
14963 ]),
14964 global_transitions: Vec::new(),
14965 fallback: None,
14966 max_no_transition: None,
14967 regenerate_on_transition: true,
14968 })
14969 .unwrap(),
14970 );
14971 let agent = Arc::new(
14972 AgentBuilder::new()
14973 .system_prompt("Test transition lifecycle reservation.")
14974 .llm(Arc::new(mock_with_response("done")))
14975 .tool(Arc::new(BlockingPathMutationTool {
14976 id: "transition_enter",
14977 path_fields: vec![ai_agents_core::PathPolicyBinding::write("path")],
14978 gate: gate.clone(),
14979 }))
14980 .state_machine(state_machine)
14981 .build()
14982 .unwrap(),
14983 );
14984
14985 let first_agent = Arc::clone(&agent);
14986 let first = tokio::spawn(async move { first_agent.transition_to("review").await });
14987 tokio::time::timeout(std::time::Duration::from_secs(2), gate.wait_until_entered())
14988 .await
14989 .expect("committed transition must enter its destination action");
14990
14991 let second = agent.transition_to("active").await.unwrap_err();
14992 assert!(second.to_string().contains("already in progress"));
14993 assert!(agent.reset().await.is_err());
14994
14995 gate.release();
14996 first.await.unwrap().unwrap();
14997 assert_eq!(agent.current_state().as_deref(), Some("review"));
14998 }
14999
15000 #[tokio::test]
15002 async fn same_state_restore_invalidates_pending_skill_confirmation() {
15003 let (agent, observed) = state_disambiguation_agent_with_skills(
15004 vec![
15005 r#"{"is_ambiguous":false,"confidence":0.99,"ambiguity_type":null,"reasoning":"top-level clear","what_is_unclear":[],"detected_language":"en"}"#,
15006 "send_report",
15007 r#"{"is_ambiguous":true,"confidence":0.2,"ambiguity_type":"missing_target","reasoning":"target missing","what_is_unclear":["target"],"detected_language":"en"}"#,
15008 r#"{"question":"What should I send?","options":null}"#,
15009 r#"{"status":"answered","selected_option":null,"enriched_input":"Send the report to Ada","resolved":{"intent":"send_report"}}"#,
15010 r#"{"question":"Should I send the report to Ada?"}"#,
15011 ],
15012 true,
15013 None,
15014 true,
15015 vec![confirmation_skill()],
15016 );
15017
15018 agent.chat("Send it").await.unwrap();
15019 agent.chat("The report to Ada").await.unwrap();
15020 let snapshot = agent.save_state().await.unwrap();
15021 assert_eq!(agent.current_state().as_deref(), Some("active"));
15022
15023 agent.restore_state(snapshot).await.unwrap();
15024
15025 assert_eq!(agent.current_state().as_deref(), Some("active"));
15026 assert!(agent.pending_skill_id.read().is_none());
15027 assert!(
15028 !agent
15029 .disambiguation_manager()
15030 .unwrap()
15031 .has_pending_clarification()
15032 .await
15033 );
15034 assert_eq!(confirmation_skill_call_count(&observed), 0);
15035 }
15036
15037 #[tokio::test]
15039 async fn direct_state_generation_change_invalidates_confirmation() {
15040 let (agent, observed) = state_disambiguation_agent_with_skills(
15041 vec![
15042 r#"{"is_ambiguous":false,"confidence":0.99,"ambiguity_type":null,"reasoning":"top-level clear","what_is_unclear":[],"detected_language":"en"}"#,
15043 "send_report",
15044 r#"{"is_ambiguous":true,"confidence":0.2,"ambiguity_type":"missing_target","reasoning":"target missing","what_is_unclear":["target"],"detected_language":"en"}"#,
15045 r#"{"question":"What should I send?","options":null}"#,
15046 r#"{"status":"answered","selected_option":null,"enriched_input":"Send the report to Ada","resolved":{"intent":"send_report"}}"#,
15047 r#"{"question":"Should I send the report to Ada?"}"#,
15048 "Confirmation cancelled.",
15049 ],
15050 true,
15051 None,
15052 true,
15053 vec![confirmation_skill()],
15054 );
15055
15056 agent.chat("Send it").await.unwrap();
15057 agent.chat("The report to Ada").await.unwrap();
15058 let state_machine = agent.state_machine().unwrap();
15059 state_machine
15060 .transition_to("review", "external test")
15061 .unwrap();
15062 state_machine
15063 .transition_to("active", "external test")
15064 .unwrap();
15065
15066 let response = agent.chat("Yes").await.unwrap();
15067
15068 assert_eq!(response.content, "Confirmation cancelled.");
15069 assert_eq!(confirmation_skill_call_count(&observed), 0);
15070 assert!(agent.pending_skill_id.read().is_none());
15071 }
15072
15073 #[tokio::test]
15074 async fn state_confirmation_does_not_add_a_question_for_clear_input() {
15075 let (agent, observed) = state_disambiguation_agent(
15076 vec![
15077 r#"{"is_ambiguous":false,"confidence":0.99,"ambiguity_type":null,"reasoning":"clear","what_is_unclear":[],"detected_language":"en"}"#,
15078 "Request executed.",
15079 ],
15080 true,
15081 None,
15082 true,
15083 );
15084
15085 let response = agent.chat("Send the report to Ada").await.unwrap();
15086
15087 assert_eq!(response.content, "Request executed.");
15088 assert_eq!(observed.call_count(), 2);
15089 }
15090
15091 #[tokio::test]
15092 async fn state_override_cannot_activate_a_disabled_top_level_manager() {
15093 let (agent, observed) =
15094 state_disambiguation_agent(vec!["Request executed."], false, Some(true), true);
15095
15096 assert!(!agent.has_disambiguation());
15097 let response = agent.chat("Send it").await.unwrap();
15098
15099 assert_eq!(response.content, "Request executed.");
15100 assert_eq!(observed.call_count(), 1);
15101 }
15102
15103 #[tokio::test]
15104 async fn native_required_choice_executes_through_the_shared_tool_path() {
15105 let mut mock = MockLLMProvider::new("native-required");
15106 mock.set_tool_choice(Some(ToolChoice::Required));
15107 let native_call = ToolCall {
15108 id: "provider-call-1".to_string(),
15109 name: "calculator".to_string(),
15110 arguments: serde_json::json!({"expression": "2 + 2"}),
15111 };
15112 let provider_state = ai_agents_core::NativeProviderState::new(
15113 "fixture-exchange-1",
15114 "fixture",
15115 "native-tools",
15116 ai_agents_core::NativeProviderTarget::new("https://fixture.invalid/", "fixture-model")
15117 .unwrap(),
15118 serde_json::json!({
15119 "role": "model",
15120 "parts": [{
15121 "functionCall": {"name": "calculator", "args": {"expression": "2 + 2"}},
15122 "thoughtSignature": "fixture-signature"
15123 }]
15124 }),
15125 vec![ai_agents_core::NativeCallBinding::new("provider-call-1", 0).unwrap()],
15126 )
15127 .unwrap();
15128 mock.add_response(
15129 LLMResponse::new("", FinishReason::ToolCall)
15130 .with_provider_state(provider_state)
15131 .unwrap()
15132 .with_tool_calls(vec![native_call])
15133 .unwrap(),
15134 );
15135 mock.add_response(LLMResponse::new("The answer is 4.", FinishReason::Stop));
15136 let observed = mock.clone();
15137 let agent = AgentBuilder::new()
15138 .system_prompt("Use the calculator when needed.")
15139 .llm(Arc::new(mock))
15140 .tool(Arc::new(CalculatorTool::new()))
15141 .build()
15142 .unwrap();
15143
15144 let response = agent.chat("What is 2 + 2?").await.unwrap();
15145
15146 assert_eq!(response.content, "The answer is 4.");
15147 assert_eq!(
15148 response.tool_calls.as_ref().unwrap()[0].id,
15149 "provider-call-1"
15150 );
15151 let calls = observed.call_history();
15152 assert_eq!(calls.len(), 2);
15153 assert!(matches!(
15154 calls[0].request.as_ref().map(|request| &request.choice),
15155 Some(ToolChoice::Required)
15156 ));
15157 assert!(matches!(
15158 calls[1].request.as_ref().map(|request| &request.choice),
15159 Some(ToolChoice::Auto)
15160 ));
15161 let replay_batch = calls[1]
15162 .messages
15163 .iter()
15164 .find_map(|message| {
15165 ai_agents_core::decode_native_tool_call_markers(&message.content).unwrap()
15166 })
15167 .expect("signed native call marker must be replayed");
15168 assert_eq!(
15169 replay_batch.provider_state().unwrap().exchange_id(),
15170 "fixture-exchange-1"
15171 );
15172 assert!(calls[1].messages.iter().any(|message| {
15173 ai_agents_core::decode_native_tool_result_markers(&message.content)
15174 .is_ok_and(|results| results.is_some())
15175 }));
15176 }
15177
15178 #[tokio::test]
15179 async fn custom_memory_loss_stops_before_signed_tool_execution() {
15180 let mut mock = MockLLMProvider::new("native-custom-memory");
15181 mock.set_tool_choice(Some(ToolChoice::Required));
15182 let call = ToolCall {
15183 id: "provider-call-drop".to_string(),
15184 name: "calculator".to_string(),
15185 arguments: serde_json::json!({"expression": "3 + 4"}),
15186 };
15187 let state = ai_agents_core::NativeProviderState::new(
15188 "fixture-exchange-drop",
15189 "fixture",
15190 "native-tools",
15191 ai_agents_core::NativeProviderTarget::new("https://fixture.invalid/", "fixture-model")
15192 .unwrap(),
15193 serde_json::json!({
15194 "role": "model",
15195 "parts": [{
15196 "functionCall": {"name": "calculator", "args": {"expression": "3 + 4"}},
15197 "thoughtSignature": "fixture-signature-drop"
15198 }]
15199 }),
15200 vec![ai_agents_core::NativeCallBinding::new("provider-call-drop", 0).unwrap()],
15201 )
15202 .unwrap();
15203 mock.add_response(
15204 LLMResponse::new("", FinishReason::ToolCall)
15205 .with_provider_state(state)
15206 .unwrap()
15207 .with_tool_calls(vec![call])
15208 .unwrap(),
15209 );
15210 let agent = AgentBuilder::new()
15211 .system_prompt("Use the calculator.")
15212 .llm(Arc::new(mock))
15213 .memory(Arc::new(DroppingSignedAssistantMemory {
15214 messages: RwLock::new(Vec::new()),
15215 }))
15216 .tool(Arc::new(CalculatorTool::new()))
15217 .build()
15218 .unwrap();
15219
15220 let error = agent.chat("What is 3 + 4?").await.unwrap_err();
15221
15222 assert!(
15223 error
15224 .to_string()
15225 .contains("removed before provider continuation")
15226 );
15227 assert!(agent.tool_call_history.read().is_empty());
15228 }
15229
15230 #[tokio::test]
15231 async fn sequential_signed_history_validates_every_prior_exchange() {
15232 let mut mock = MockLLMProvider::new("native-sequential-memory");
15233 mock.set_tool_choice(Some(ToolChoice::Required));
15234 mock.add_response(signed_calculator_response(
15235 "seq-exchange-1",
15236 "seq-call-1",
15237 "1 + 1",
15238 ));
15239 mock.add_response(signed_calculator_response(
15240 "seq-exchange-2",
15241 "seq-call-2",
15242 "2 + 2",
15243 ));
15244 let agent = AgentBuilder::new()
15245 .system_prompt("Use the calculator sequentially.")
15246 .llm(Arc::new(mock))
15247 .memory(Arc::new(DroppingEarlierSequentialMemory {
15248 messages: RwLock::new(Vec::new()),
15249 signed_seen: std::sync::atomic::AtomicUsize::new(0),
15250 }))
15251 .tool(Arc::new(CalculatorTool::new()))
15252 .build()
15253 .unwrap();
15254
15255 let error = agent.chat("Calculate twice.").await.unwrap_err();
15256
15257 assert!(error.to_string().contains("seq-exchange-1"));
15258 assert_eq!(agent.tool_call_history.read().len(), 1);
15259 }
15260
15261 #[tokio::test]
15262 async fn post_transition_signed_hitl_rejection_stops_before_continuation() {
15263 let mut native = MockLLMProvider::new("post-transition-native");
15264 native.set_tool_choice(Some(ToolChoice::Auto));
15265 let call = ToolCall {
15266 id: "post-transition-call".to_string(),
15267 name: "echo".to_string(),
15268 arguments: serde_json::json!({"message": "hello"}),
15269 };
15270 let state = ai_agents_core::NativeProviderState::new(
15271 "post-transition-exchange",
15272 "fixture",
15273 "native-tools",
15274 ai_agents_core::NativeProviderTarget::new("https://fixture.invalid/", "fixture-model")
15275 .unwrap(),
15276 serde_json::json!({
15277 "role": "model",
15278 "parts": [{
15279 "functionCall": {"name": "echo", "args": {"message": "hello"}},
15280 "thoughtSignature": "post-transition-signature"
15281 }]
15282 }),
15283 vec![ai_agents_core::NativeCallBinding::new("post-transition-call", 0).unwrap()],
15284 )
15285 .unwrap();
15286 native.add_response(
15287 LLMResponse::new("", FinishReason::ToolCall)
15288 .with_provider_state(state)
15289 .unwrap()
15290 .with_tool_calls(vec![call])
15291 .unwrap(),
15292 );
15293 let observed_native = native.clone();
15294 let yaml = r#"
15295name: PostTransitionNativeReject
15296system_prompt: test
15297tools: [echo]
15298hitl:
15299 tools:
15300 echo:
15301 require_approval: true
15302states:
15303 initial: intake
15304 states:
15305 intake:
15306 prompt: intake
15307 transitions:
15308 - to: active
15309 guard:
15310 context:
15311 route:
15312 eq: active
15313 active:
15314 prompt: active
15315 llm: native
15316"#;
15317 let agent = AgentBuilder::from_yaml(yaml)
15318 .unwrap()
15319 .llm(Arc::new(mock_with_response("stale intake response")))
15320 .llm_alias("native", Arc::new(native))
15321 .auto_configure_features()
15322 .unwrap()
15323 .build()
15324 .unwrap();
15325 agent
15326 .set_context("route", serde_json::json!("active"))
15327 .unwrap();
15328
15329 let error = agent.chat("move to active").await.unwrap_err();
15330
15331 assert!(matches!(error, AgentError::HITLRejected(_)));
15332 assert_eq!(observed_native.call_count(), 1);
15333 }
15334
15335 #[test]
15336 fn runtime_overflow_removes_a_past_signed_user_turn_as_one_prefix() {
15337 let call = ToolCall {
15338 id: "overflow-call".to_string(),
15339 name: "calculator".to_string(),
15340 arguments: serde_json::json!({"expression": "1 + 1"}),
15341 };
15342 let state = ai_agents_core::NativeProviderState::new(
15343 "overflow-exchange",
15344 "google",
15345 "generateContent",
15346 ai_agents_core::NativeProviderTarget::new("https://example.invalid/", "gemini-3")
15347 .unwrap(),
15348 serde_json::json!({
15349 "role": "model",
15350 "parts": [{
15351 "functionCall": {"name": "calculator", "args": {"expression": "1 + 1"}},
15352 "thoughtSignature": "overflow-signature"
15353 }]
15354 }),
15355 vec![ai_agents_core::NativeCallBinding::new("overflow-call", 0).unwrap()],
15356 )
15357 .unwrap();
15358 let call_marker = ai_agents_core::encode_native_tool_call_markers(
15359 std::slice::from_ref(&call),
15360 Some(&state),
15361 )
15362 .unwrap();
15363 let result_marker = ai_agents_core::encode_native_tool_result_marker(
15364 &call,
15365 serde_json::json!({"result": 2}),
15366 )
15367 .unwrap();
15368 let history = vec![
15369 ChatMessage::user("old question"),
15370 ChatMessage::assistant(call_marker),
15371 ChatMessage::function("calculator", result_marker),
15372 ChatMessage::assistant("old answer"),
15373 ChatMessage::user("new question"),
15374 ];
15375
15376 let removable = RuntimeAgent::native_safe_prefix_at_least(&history, 1).unwrap();
15377
15378 assert_eq!(removable, 4);
15379 }
15380
15381 #[test]
15382 fn auxiliary_projection_does_not_interpret_user_marker_text() {
15383 let user_text = serde_json::json!({
15384 "_ai_agents_native_tool_call": true,
15385 "id": "",
15386 "tool": "user-data",
15387 "arguments": {}
15388 })
15389 .to_string();
15390
15391 let projected =
15392 RuntimeAgent::readable_native_messages(vec![ChatMessage::user(&user_text)]).unwrap();
15393
15394 assert_eq!(projected[0].content, user_text);
15395 }
15396
15397 #[tokio::test]
15398 async fn terminal_provider_history_error_skips_retry_and_static_fallback() {
15399 let calls = Arc::new(std::sync::atomic::AtomicU32::new(0));
15400 let recovery = RecoveryManager::new(ai_agents_recovery::ErrorRecoveryConfig {
15401 default: ai_agents_recovery::RetryConfig {
15402 max_retries: 3,
15403 ..Default::default()
15404 },
15405 llm: ai_agents_recovery::LLMRecoveryConfig {
15406 on_failure: LLMFailureAction::FallbackResponse {
15407 message: "must not be returned".to_string(),
15408 },
15409 ..Default::default()
15410 },
15411 ..Default::default()
15412 });
15413 let agent = AgentBuilder::new()
15414 .system_prompt("Reject corrupted native history.")
15415 .llm(Arc::new(TerminalHistoryProvider {
15416 calls: Arc::clone(&calls),
15417 }))
15418 .recovery_manager(recovery)
15419 .build()
15420 .unwrap();
15421
15422 let error = agent.chat("continue").await.unwrap_err();
15423
15424 assert!(
15425 error
15426 .to_string()
15427 .contains("native history integrity failure")
15428 );
15429 assert_eq!(calls.load(Ordering::SeqCst), 1);
15430 }
15431
15432 #[tokio::test]
15433 async fn prompt_fallback_uses_one_corrective_retry() {
15434 let mut mock = MockLLMProvider::new("prompt-required");
15435 mock.set_tool_choice(Some(ToolChoice::Required));
15436 mock.set_native_tool_support(false);
15437 mock.set_responses(
15438 vec![
15439 "I can calculate that.".to_string(),
15440 r#"{"tool":"calculator","arguments":{"expression":"2 + 2"}}"#.to_string(),
15441 "The answer is 4.".to_string(),
15442 ],
15443 false,
15444 );
15445 let observed = mock.clone();
15446 let agent = AgentBuilder::new()
15447 .system_prompt("Use tools.")
15448 .llm(Arc::new(mock))
15449 .tool(Arc::new(CalculatorTool::new()))
15450 .build()
15451 .unwrap();
15452
15453 let response = agent.chat("What is 2 + 2?").await.unwrap();
15454
15455 assert_eq!(response.content, "The answer is 4.");
15456 assert_eq!(observed.call_count(), 3);
15457 let corrective = &observed.call_history()[1].messages;
15458 assert!(
15459 corrective
15460 .last()
15461 .unwrap()
15462 .content
15463 .contains("previous response")
15464 );
15465 }
15466
15467 #[tokio::test]
15468 async fn prompt_fallback_fails_after_one_noncompliant_retry() {
15469 let mut mock = MockLLMProvider::new("prompt-required-failure");
15470 mock.set_tool_choice(Some(ToolChoice::Required));
15471 mock.set_native_tool_support(false);
15472 mock.set_responses(
15473 vec!["No tool.".to_string(), "Still no tool.".to_string()],
15474 false,
15475 );
15476 let observed = mock.clone();
15477 let agent = AgentBuilder::new()
15478 .system_prompt("Use tools.")
15479 .llm(Arc::new(mock))
15480 .tool(Arc::new(CalculatorTool::new()))
15481 .build()
15482 .unwrap();
15483
15484 let error = agent.chat("What is 2 + 2?").await.unwrap_err();
15485
15486 assert!(error.to_string().contains("one corrective retry"));
15487 assert_eq!(observed.call_count(), 2);
15488 }
15489
15490 #[tokio::test]
15491 async fn specific_choice_cannot_widen_the_effective_grant() {
15492 let mut mock = MockLLMProvider::new("specific-outside-grant");
15493 mock.set_tool_choice(Some(ToolChoice::Specific("random".to_string())));
15494 let observed = mock.clone();
15495 let agent = AgentBuilder::new()
15496 .system_prompt("Use tools.")
15497 .llm(Arc::new(mock))
15498 .tool(Arc::new(CalculatorTool::new()))
15499 .build()
15500 .unwrap();
15501
15502 let error = agent.chat("Generate a value.").await.unwrap_err();
15503
15504 assert!(error.to_string().contains("is not registered"));
15505 assert_eq!(observed.call_count(), 0);
15506 }
15507
15508 #[tokio::test]
15509 async fn none_choice_exposes_no_tool_protocol() {
15510 let mut mock = MockLLMProvider::new("no-tools");
15511 mock.set_tool_choice(Some(ToolChoice::None));
15512 mock.set_response(r#"{"tool":"calculator","arguments":{"expression":"2 + 2"}}"#);
15513 let observed = mock.clone();
15514 let agent = AgentBuilder::new()
15515 .system_prompt("Answer directly.")
15516 .llm(Arc::new(mock))
15517 .tool(Arc::new(CalculatorTool::new()))
15518 .build()
15519 .unwrap();
15520
15521 let response = agent.chat("Hello").await.unwrap();
15522
15523 assert!(response.tool_calls.is_none());
15524 assert_eq!(observed.call_count(), 1);
15525 let call = observed.last_call().unwrap();
15526 assert!(call.request.is_none());
15527 assert!(
15528 call.messages
15529 .iter()
15530 .all(|message| !message.content.contains("Available tools:"))
15531 );
15532 }
15533
15534 struct RuntimeStorage {
15535 capabilities: Box<[StorageCapability]>,
15536 snapshots: RwLock<HashMap<String, AgentSnapshot>>,
15537 metadata: RwLock<HashMap<String, ai_agents_core::SessionMetadata>>,
15538 metadata_save_calls: AtomicU64,
15539 metadata_load_calls: AtomicU64,
15540 fail_metadata_save: AtomicBool,
15541 fail_metadata_load: AtomicBool,
15542 }
15543
15544 impl RuntimeStorage {
15545 fn new(capabilities: impl IntoIterator<Item = StorageCapability>) -> Self {
15546 Self {
15547 capabilities: capabilities.into_iter().collect(),
15548 snapshots: RwLock::new(HashMap::new()),
15549 metadata: RwLock::new(HashMap::new()),
15550 metadata_save_calls: AtomicU64::new(0),
15551 metadata_load_calls: AtomicU64::new(0),
15552 fail_metadata_save: AtomicBool::new(false),
15553 fail_metadata_load: AtomicBool::new(false),
15554 }
15555 }
15556 }
15557
15558 #[async_trait]
15559 impl AgentStorage for RuntimeStorage {
15560 fn supports(&self, capability: StorageCapability) -> bool {
15561 self.capabilities.contains(&capability)
15562 }
15563
15564 async fn save(&self, session_id: &str, snapshot: &AgentSnapshot) -> Result<()> {
15565 self.snapshots
15566 .write()
15567 .insert(session_id.to_string(), snapshot.clone());
15568 Ok(())
15569 }
15570
15571 async fn load(&self, session_id: &str) -> Result<Option<AgentSnapshot>> {
15572 Ok(self.snapshots.read().get(session_id).cloned())
15573 }
15574
15575 async fn delete(&self, session_id: &str) -> Result<()> {
15576 self.snapshots.write().remove(session_id);
15577 Ok(())
15578 }
15579
15580 async fn list_sessions(&self) -> Result<Vec<String>> {
15581 Ok(self.snapshots.read().keys().cloned().collect())
15582 }
15583
15584 async fn save_snapshot_with_metadata(
15585 &self,
15586 session_id: &str,
15587 snapshot: &AgentSnapshot,
15588 metadata: &ai_agents_core::SessionMetadata,
15589 ) -> Result<()> {
15590 self.metadata_save_calls.fetch_add(1, Ordering::SeqCst);
15591 if self.fail_metadata_save.load(Ordering::SeqCst) {
15592 return Err(AgentError::Persistence("metadata save failed".into()));
15593 }
15594 self.snapshots
15595 .write()
15596 .insert(session_id.to_string(), snapshot.clone());
15597 self.metadata
15598 .write()
15599 .insert(session_id.to_string(), metadata.clone());
15600 Ok(())
15601 }
15602
15603 async fn save_metadata(
15604 &self,
15605 session_id: &str,
15606 metadata: &ai_agents_core::SessionMetadata,
15607 ) -> Result<()> {
15608 self.metadata_save_calls.fetch_add(1, Ordering::SeqCst);
15609 if self.fail_metadata_save.load(Ordering::SeqCst) {
15610 return Err(AgentError::Persistence("metadata save failed".into()));
15611 }
15612 self.metadata
15613 .write()
15614 .insert(session_id.to_string(), metadata.clone());
15615 Ok(())
15616 }
15617
15618 async fn load_metadata(
15619 &self,
15620 session_id: &str,
15621 ) -> Result<Option<ai_agents_core::SessionMetadata>> {
15622 self.metadata_load_calls.fetch_add(1, Ordering::SeqCst);
15623 if self.fail_metadata_load.load(Ordering::SeqCst) {
15624 return Err(AgentError::Persistence("metadata load failed".into()));
15625 }
15626 Ok(self.metadata.read().get(session_id).cloned())
15627 }
15628 }
15629
15630 fn runtime_storage_agent() -> RuntimeAgent {
15631 AgentBuilder::new()
15632 .system_prompt("Test runtime storage integration.")
15633 .llm(Arc::new(mock_with_response("done")))
15634 .build()
15635 .unwrap()
15636 }
15637
15638 fn restore_spec(id: &str) -> crate::spec::AgentSpec {
15639 crate::spec::AgentSpec {
15640 name: id.to_string(),
15641 system_prompt: format!("Restore child {id}."),
15642 ..crate::spec::AgentSpec::default()
15643 }
15644 }
15645
15646 fn restore_entry(id: &str) -> ai_agents_core::SpawnedAgentEntry {
15647 ai_agents_core::SpawnedAgentEntry {
15648 id: id.to_string(),
15649 name: id.to_string(),
15650 spec_yaml: serde_yaml::to_string(&restore_spec(id)).unwrap(),
15651 }
15652 }
15653
15654 fn restore_spawner(
15655 storage: Arc<RuntimeStorage>,
15656 max_agents: usize,
15657 ) -> (
15658 Arc<crate::spawner::AgentSpawner>,
15659 Arc<crate::spawner::AgentRegistry>,
15660 ) {
15661 let mut llms = LLMRegistry::new();
15662 llms.register("default", Arc::new(mock_with_response("done")));
15663 (
15664 Arc::new(
15665 crate::spawner::AgentSpawner::new()
15666 .with_shared_llms(llms)
15667 .with_shared_storage(storage)
15668 .with_max_agents(max_agents),
15669 ),
15670 Arc::new(crate::spawner::AgentRegistry::new()),
15671 )
15672 }
15673
15674 async fn save_restore_target(
15675 parent: &RuntimeAgent,
15676 storage: &RuntimeStorage,
15677 session_id: &str,
15678 entries: Vec<ai_agents_core::SpawnedAgentEntry>,
15679 ) {
15680 let mut snapshot = parent.save_state().await.unwrap();
15681 snapshot.spawned_agents = Some(entries);
15682 storage.save(session_id, &snapshot).await.unwrap();
15683 storage
15684 .save_metadata(session_id, &ai_agents_core::SessionMetadata::default())
15685 .await
15686 .unwrap();
15687 }
15688
15689 #[tokio::test]
15690 async fn storage_init_requires_storage_for_actor_facts() {
15691 let facts = ai_agents_facts::FactsConfig {
15692 enabled: true,
15693 ..Default::default()
15694 };
15695 let agent = runtime_storage_agent().with_facts_config(None, Some(facts));
15696
15697 let error = agent.init_storage().await.unwrap_err();
15698 assert!(matches!(
15699 error,
15700 AgentError::Config(message)
15701 if message.contains("actor facts or actor memory")
15702 && message.contains("none is configured or injected")
15703 ));
15704 }
15705
15706 #[tokio::test]
15707 async fn storage_init_validates_actor_facts_capability() {
15708 let storage = Arc::new(RuntimeStorage::new([StorageCapability::Snapshot]));
15709 let actor_memory = ai_agents_facts::ActorMemoryConfig {
15710 enabled: true,
15711 ..Default::default()
15712 };
15713 let agent = runtime_storage_agent()
15714 .with_storage(storage)
15715 .with_facts_config(Some(actor_memory), None);
15716
15717 assert!(matches!(
15718 agent.init_storage().await,
15719 Err(AgentError::UnsupportedStorageCapability(
15720 StorageCapability::ActorFacts
15721 ))
15722 ));
15723 }
15724
15725 #[tokio::test]
15726 async fn blocking_chat_rejects_unsupported_required_storage() {
15727 let storage = Arc::new(RuntimeStorage::new([StorageCapability::Snapshot]));
15728 let facts = ai_agents_facts::FactsConfig {
15729 enabled: true,
15730 ..Default::default()
15731 };
15732 let agent = runtime_storage_agent()
15733 .with_storage(storage)
15734 .with_facts_config(None, Some(facts));
15735
15736 assert!(matches!(
15737 agent.chat("hello").await,
15738 Err(AgentError::UnsupportedStorageCapability(
15739 StorageCapability::ActorFacts
15740 ))
15741 ));
15742 }
15743
15744 #[tokio::test]
15745 async fn streaming_chat_rejects_unsupported_required_storage_before_stream_creation() {
15746 let storage = Arc::new(RuntimeStorage::new([StorageCapability::Snapshot]));
15747 let config = ai_agents_relationships::RelationshipConfig {
15748 enabled: true,
15749 ..Default::default()
15750 };
15751 let manager = Arc::new(RelationshipManager::from_config(config).unwrap());
15752 let agent = runtime_storage_agent()
15753 .with_storage(storage)
15754 .with_relationships(manager);
15755
15756 assert!(matches!(
15757 agent.chat_stream("hello").await,
15758 Err(AgentError::UnsupportedStorageCapability(
15759 StorageCapability::ActorRelationships
15760 ))
15761 ));
15762 }
15763
15764 #[tokio::test]
15765 async fn storage_init_completes_facts_for_injected_storage() {
15766 let storage = Arc::new(RuntimeStorage::new([
15767 StorageCapability::Snapshot,
15768 StorageCapability::ActorFacts,
15769 ]));
15770 let facts = ai_agents_facts::FactsConfig {
15771 enabled: true,
15772 ..Default::default()
15773 };
15774 let agent = runtime_storage_agent()
15775 .with_storage(storage)
15776 .with_facts_config(None, Some(facts));
15777
15778 agent.init_storage().await.unwrap();
15779 assert!(agent.fact_store().is_some());
15780 }
15781
15782 #[tokio::test]
15783 async fn storage_init_requires_storage_for_persistent_relationships() {
15784 let config = ai_agents_relationships::RelationshipConfig {
15785 enabled: true,
15786 ..Default::default()
15787 };
15788 let manager = Arc::new(RelationshipManager::from_config(config).unwrap());
15789 let agent = runtime_storage_agent().with_relationships(manager);
15790
15791 let error = agent.init_storage().await.unwrap_err();
15792 assert!(matches!(
15793 error,
15794 AgentError::Config(message)
15795 if message.contains("persistent relationships")
15796 && message.contains("none is configured or injected")
15797 ));
15798 }
15799
15800 #[tokio::test]
15801 async fn storage_init_validates_persistent_relationships_capability() {
15802 let storage = Arc::new(RuntimeStorage::new([StorageCapability::Snapshot]));
15803 let config = ai_agents_relationships::RelationshipConfig {
15804 enabled: true,
15805 ..Default::default()
15806 };
15807 let manager = Arc::new(RelationshipManager::from_config(config).unwrap());
15808 let agent = runtime_storage_agent()
15809 .with_storage(storage)
15810 .with_relationships(manager);
15811
15812 assert!(matches!(
15813 agent.init_storage().await,
15814 Err(AgentError::UnsupportedStorageCapability(
15815 StorageCapability::ActorRelationships
15816 ))
15817 ));
15818 }
15819
15820 #[tokio::test]
15821 async fn session_restore_updates_identity_and_clears_stale_actor_binding() {
15822 let storage = Arc::new(RuntimeStorage::new([
15823 StorageCapability::Snapshot,
15824 StorageCapability::SessionMetadata,
15825 ]));
15826 let agent = runtime_storage_agent().with_storage(storage.clone());
15827 agent.set_actor_id("old-actor").unwrap();
15828 agent.save_session("old").await.unwrap();
15829 storage
15830 .save("target", &agent.save_state().await.unwrap())
15831 .await
15832 .unwrap();
15833 storage
15834 .save_metadata("target", &ai_agents_core::SessionMetadata::default())
15835 .await
15836 .unwrap();
15837
15838 assert!(agent.load_session("target").await.unwrap());
15839
15840 assert_eq!(agent.current_session_id.read().as_deref(), Some("target"));
15841 assert_eq!(agent.actor_id(), None);
15842 }
15843
15844 #[tokio::test]
15845 async fn complete_restore_reconciles_growth_shrink_and_empty_topologies() {
15846 let storage = Arc::new(RuntimeStorage::new([
15847 StorageCapability::Snapshot,
15848 StorageCapability::SessionMetadata,
15849 ]));
15850 let (spawner, registry) = restore_spawner(storage.clone(), 3);
15851 let parent = runtime_storage_agent()
15852 .with_storage(storage.clone())
15853 .with_spawner_handles(Arc::clone(&spawner), Arc::clone(®istry));
15854
15855 for id in ["a", "b"] {
15856 let spawned = spawner
15857 .spawn_with_id(id.to_string(), restore_spec(id))
15858 .await
15859 .unwrap();
15860 spawned.agent.save_session("grow").await.unwrap();
15861 registry.register(spawned).await.unwrap();
15862 }
15863 let staged_c = crate::spawner::storage::NamespacedStorage::new(storage.clone(), "c");
15864 staged_c
15865 .save("grow", &AgentSnapshot::new("c".into()))
15866 .await
15867 .unwrap();
15868 staged_c
15869 .save_metadata("grow", &ai_agents_core::SessionMetadata::default())
15870 .await
15871 .unwrap();
15872 save_restore_target(
15873 &parent,
15874 storage.as_ref(),
15875 "grow",
15876 vec![restore_entry("a"), restore_entry("b"), restore_entry("c")],
15877 )
15878 .await;
15879
15880 assert_eq!(parent.restore_session_full("grow").await.unwrap(), 3);
15881 assert_eq!(registry.count(), 3);
15882 assert!(registry.contains("c"));
15883 assert_eq!(spawner.spawned_count(), 3);
15884
15885 for id in ["a", "b"] {
15886 registry
15887 .get(id)
15888 .unwrap()
15889 .save_session("shrink")
15890 .await
15891 .unwrap();
15892 }
15893 save_restore_target(
15894 &parent,
15895 storage.as_ref(),
15896 "shrink",
15897 vec![restore_entry("a"), restore_entry("b")],
15898 )
15899 .await;
15900
15901 assert_eq!(parent.restore_session_full("shrink").await.unwrap(), 2);
15902 assert_eq!(registry.count(), 2);
15903 assert!(!registry.contains("c"));
15904 assert_eq!(spawner.spawned_count(), 2);
15905
15906 save_restore_target(&parent, storage.as_ref(), "empty", Vec::new()).await;
15907
15908 assert_eq!(parent.restore_session_full("empty").await.unwrap(), 0);
15909 assert_eq!(registry.count(), 0);
15910 assert_eq!(spawner.spawned_count(), 0);
15911 assert_eq!(parent.current_session_id.read().as_deref(), Some("empty"));
15912 }
15913
15914 #[tokio::test]
15915 async fn storage_session_metadata_is_called_only_when_advertised() {
15916 let storage = Arc::new(RuntimeStorage::new([StorageCapability::Snapshot]));
15917 storage.fail_metadata_save.store(true, Ordering::SeqCst);
15918 storage.fail_metadata_load.store(true, Ordering::SeqCst);
15919 let agent = runtime_storage_agent().with_storage(storage.clone());
15920
15921 agent.save_session("session").await.unwrap();
15922 assert!(agent.load_session("session").await.unwrap());
15923 assert_eq!(storage.metadata_save_calls.load(Ordering::SeqCst), 0);
15924 assert_eq!(storage.metadata_load_calls.load(Ordering::SeqCst), 0);
15925 }
15926
15927 #[cfg(feature = "sqlite")]
15928 #[tokio::test]
15929 async fn sqlite_runtime_save_filter_reopen_and_reload_stay_consistent() {
15930 let directory =
15931 std::env::temp_dir().join(format!("ai-agents-runtime-sqlite-{}", uuid::Uuid::new_v4()));
15932 let path = directory.join("sessions.sqlite");
15933 let path_string = path.to_string_lossy().into_owned();
15934 let storage = Arc::new(
15935 ai_agents_storage::SqliteStorage::new(&path_string)
15936 .await
15937 .unwrap(),
15938 );
15939 let agent = runtime_storage_agent().with_storage(storage.clone());
15940 agent.set_session_metadata(ai_agents_core::SessionMetadata {
15941 tags: vec!["initial".into()],
15942 ..Default::default()
15943 });
15944 agent.chat("persist this turn").await.unwrap();
15945 agent.save_session("session").await.unwrap();
15946
15947 agent.set_session_metadata(ai_agents_core::SessionMetadata {
15948 tags: vec!["updated".into()],
15949 ..Default::default()
15950 });
15951 agent.save_session("session").await.unwrap();
15952 assert!(
15953 agent
15954 .list_sessions_filtered(&ai_agents_core::SessionFilter {
15955 tags: Some(vec!["initial".into()]),
15956 ..Default::default()
15957 })
15958 .await
15959 .unwrap()
15960 .is_empty()
15961 );
15962 assert_eq!(
15963 agent
15964 .list_sessions_filtered(&ai_agents_core::SessionFilter {
15965 tags: Some(vec!["updated".into()]),
15966 ..Default::default()
15967 })
15968 .await
15969 .unwrap()
15970 .len(),
15971 1
15972 );
15973 drop(agent);
15974 storage.close().await;
15975 drop(storage);
15976
15977 let reopened_storage = Arc::new(
15978 ai_agents_storage::SqliteStorage::new(&path_string)
15979 .await
15980 .unwrap(),
15981 );
15982 let restored = runtime_storage_agent().with_storage(reopened_storage.clone());
15983 assert!(restored.load_session("session").await.unwrap());
15984 assert_eq!(restored.session_metadata().tags, vec!["updated"]);
15985 assert_eq!(
15986 restored.current_session_id.read().as_deref(),
15987 Some("session")
15988 );
15989 assert!(restored.save_state().await.unwrap().memory.messages.len() >= 2);
15990 assert_eq!(
15991 restored
15992 .list_sessions_filtered(&ai_agents_core::SessionFilter {
15993 tags: Some(vec!["updated".into()]),
15994 ..Default::default()
15995 })
15996 .await
15997 .unwrap()
15998 .len(),
15999 1
16000 );
16001
16002 drop(restored);
16003 reopened_storage.close().await;
16004 drop(reopened_storage);
16005 crate::remove_sqlite_test_directory(&directory)
16006 .await
16007 .unwrap();
16008 }
16009
16010 #[tokio::test]
16011 async fn storage_session_metadata_backend_failures_propagate() {
16012 let storage = Arc::new(RuntimeStorage::new([
16013 StorageCapability::Snapshot,
16014 StorageCapability::SessionMetadata,
16015 ]));
16016 let agent = runtime_storage_agent().with_storage(storage.clone());
16017
16018 agent.save_session("session").await.unwrap();
16019 storage
16020 .save("target", &agent.save_state().await.unwrap())
16021 .await
16022 .unwrap();
16023 storage.fail_metadata_load.store(true, Ordering::SeqCst);
16024 assert!(matches!(
16025 agent.load_session("target").await,
16026 Err(AgentError::Persistence(message)) if message == "metadata load failed"
16027 ));
16028 assert_eq!(agent.current_session_id.read().as_deref(), Some("session"));
16029
16030 storage.fail_metadata_save.store(true, Ordering::SeqCst);
16031 assert!(matches!(
16032 agent.save_session("session").await,
16033 Err(AgentError::Persistence(message)) if message == "metadata save failed"
16034 ));
16035 }
16036
16037 struct ProviderFutureDropSignal {
16038 dropped: Arc<AtomicBool>,
16039 }
16040
16041 impl Drop for ProviderFutureDropSignal {
16042 fn drop(&mut self) {
16043 self.dropped.store(true, Ordering::SeqCst);
16044 }
16045 }
16046
16047 struct BufferedLockingProvider {
16048 lock: Arc<tokio::sync::Mutex<()>>,
16049 stream_started: Arc<tokio::sync::Notify>,
16050 stream_dropped: Arc<AtomicBool>,
16051 committed_after_drop: Arc<AtomicBool>,
16052 }
16053
16054 #[async_trait]
16055 impl LLMProvider for BufferedLockingProvider {
16056 async fn complete(
16057 &self,
16058 _messages: &[ChatMessage],
16059 _config: Option<&LLMConfig>,
16060 ) -> std::result::Result<LLMResponse, LLMError> {
16061 let _guard = self.lock.lock().await;
16062 self.committed_after_drop
16063 .store(self.stream_dropped.load(Ordering::SeqCst), Ordering::SeqCst);
16064 Ok(LLMResponse::new(
16065 "Committed technical response.",
16066 FinishReason::Stop,
16067 ))
16068 }
16069
16070 async fn complete_stream(
16071 &self,
16072 _messages: &[ChatMessage],
16073 _config: Option<&LLMConfig>,
16074 ) -> std::result::Result<
16075 Box<dyn Stream<Item = std::result::Result<LLMChunk, LLMError>> + Unpin + Send>,
16076 LLMError,
16077 > {
16078 let _guard = self.lock.lock().await;
16079 let _drop_signal = ProviderFutureDropSignal {
16080 dropped: Arc::clone(&self.stream_dropped),
16081 };
16082 self.stream_started.notify_one();
16083 std::future::pending().await
16084 }
16085
16086 fn provider_name(&self) -> &str {
16087 "buffered-locking"
16088 }
16089
16090 fn supports(&self, _feature: LLMFeature) -> bool {
16091 false
16092 }
16093 }
16094
16095 struct PendingDropStream {
16096 dropped: Arc<AtomicBool>,
16097 dropped_notify: Arc<tokio::sync::Notify>,
16098 }
16099
16100 impl Stream for PendingDropStream {
16101 type Item = std::result::Result<LLMChunk, LLMError>;
16102
16103 fn poll_next(
16104 self: Pin<&mut Self>,
16105 _cx: &mut std::task::Context<'_>,
16106 ) -> std::task::Poll<Option<Self::Item>> {
16107 std::task::Poll::Pending
16108 }
16109 }
16110
16111 impl Drop for PendingDropStream {
16112 fn drop(&mut self) {
16113 self.dropped.store(true, Ordering::SeqCst);
16114 self.dropped_notify.notify_one();
16115 }
16116 }
16117
16118 struct EstablishedStreamProvider {
16119 stream_started: Arc<tokio::sync::Notify>,
16120 stream_dropped: Arc<AtomicBool>,
16121 stream_dropped_notify: Arc<tokio::sync::Notify>,
16122 committed_after_drop: Arc<AtomicBool>,
16123 }
16124
16125 #[async_trait]
16126 impl LLMProvider for EstablishedStreamProvider {
16127 async fn complete(
16128 &self,
16129 _messages: &[ChatMessage],
16130 _config: Option<&LLMConfig>,
16131 ) -> std::result::Result<LLMResponse, LLMError> {
16132 if !self.stream_dropped.load(Ordering::SeqCst) {
16133 self.stream_dropped_notify.notified().await;
16134 }
16135 self.committed_after_drop
16136 .store(self.stream_dropped.load(Ordering::SeqCst), Ordering::SeqCst);
16137 Ok(LLMResponse::new(
16138 "Committed technical response.",
16139 FinishReason::Stop,
16140 ))
16141 }
16142
16143 async fn complete_stream(
16144 &self,
16145 _messages: &[ChatMessage],
16146 _config: Option<&LLMConfig>,
16147 ) -> std::result::Result<
16148 Box<dyn Stream<Item = std::result::Result<LLMChunk, LLMError>> + Unpin + Send>,
16149 LLMError,
16150 > {
16151 self.stream_started.notify_one();
16152 Ok(Box::new(PendingDropStream {
16153 dropped: Arc::clone(&self.stream_dropped),
16154 dropped_notify: Arc::clone(&self.stream_dropped_notify),
16155 }))
16156 }
16157
16158 fn provider_name(&self) -> &str {
16159 "established-stream"
16160 }
16161
16162 fn supports(&self, _feature: LLMFeature) -> bool {
16163 false
16164 }
16165 }
16166
16167 struct FirstCallLockingProvider {
16168 lock: Arc<tokio::sync::Mutex<()>>,
16169 first_started: Arc<tokio::sync::Notify>,
16170 first_dropped: Arc<AtomicBool>,
16171 committed_after_drop: Arc<AtomicBool>,
16172 calls: AtomicU64,
16173 }
16174
16175 #[async_trait]
16176 impl LLMProvider for FirstCallLockingProvider {
16177 async fn complete(
16178 &self,
16179 _messages: &[ChatMessage],
16180 _config: Option<&LLMConfig>,
16181 ) -> std::result::Result<LLMResponse, LLMError> {
16182 let _guard = self.lock.lock().await;
16183 let call = self.calls.fetch_add(1, Ordering::SeqCst);
16184 if call == 0 {
16185 let _drop_signal = ProviderFutureDropSignal {
16186 dropped: Arc::clone(&self.first_dropped),
16187 };
16188 self.first_started.notify_one();
16189 return std::future::pending().await;
16190 }
16191 self.committed_after_drop
16192 .store(self.first_dropped.load(Ordering::SeqCst), Ordering::SeqCst);
16193 Ok(LLMResponse::new(
16194 "Committed technical response.",
16195 FinishReason::Stop,
16196 ))
16197 }
16198
16199 async fn complete_stream(
16200 &self,
16201 _messages: &[ChatMessage],
16202 _config: Option<&LLMConfig>,
16203 ) -> std::result::Result<
16204 Box<dyn Stream<Item = std::result::Result<LLMChunk, LLMError>> + Unpin + Send>,
16205 LLMError,
16206 > {
16207 Err(LLMError::Other(
16208 "streaming is not used in this test".to_string(),
16209 ))
16210 }
16211
16212 fn provider_name(&self) -> &str {
16213 "first-call-locking"
16214 }
16215
16216 fn supports(&self, _feature: LLMFeature) -> bool {
16217 false
16218 }
16219 }
16220
16221 struct RoutingAfterProviderStart {
16222 provider_started: Arc<tokio::sync::Notify>,
16223 }
16224
16225 #[async_trait]
16226 impl LLMProvider for RoutingAfterProviderStart {
16227 async fn complete(
16228 &self,
16229 _messages: &[ChatMessage],
16230 _config: Option<&LLMConfig>,
16231 ) -> std::result::Result<LLMResponse, LLMError> {
16232 self.provider_started.notified().await;
16233 Ok(LLMResponse::new("1", FinishReason::Stop))
16234 }
16235
16236 async fn complete_stream(
16237 &self,
16238 _messages: &[ChatMessage],
16239 _config: Option<&LLMConfig>,
16240 ) -> std::result::Result<
16241 Box<dyn Stream<Item = std::result::Result<LLMChunk, LLMError>> + Unpin + Send>,
16242 LLMError,
16243 > {
16244 Err(LLMError::Other(
16245 "streaming is not used in this test".to_string(),
16246 ))
16247 }
16248
16249 fn provider_name(&self) -> &str {
16250 "routing-after-start"
16251 }
16252
16253 fn supports(&self, _feature: LLMFeature) -> bool {
16254 false
16255 }
16256 }
16257
16258 struct ResponseCountingHooks {
16260 responses: Arc<std::sync::atomic::AtomicUsize>,
16261 }
16262
16263 struct RootTurnProbeProvider {
16265 complete_entered: tokio::sync::mpsc::UnboundedSender<()>,
16266 }
16267
16268 struct ResponseChatHooks {
16270 target: parking_lot::Mutex<Option<Weak<RuntimeAgent>>>,
16271 invoked: AtomicBool,
16272 nested_result: parking_lot::Mutex<Option<std::result::Result<String, String>>>,
16273 }
16274
16275 struct ConcurrentResponseHooks {
16277 registry: Weak<crate::spawner::AgentRegistry>,
16278 child_id: String,
16279 invoked: AtomicBool,
16280 nested_result: parking_lot::Mutex<Option<std::result::Result<String, String>>>,
16281 }
16282
16283 struct RetryDeadlineTool {
16285 calls: Arc<std::sync::atomic::AtomicUsize>,
16286 deadlines: Arc<parking_lot::Mutex<Vec<chrono::DateTime<chrono::Utc>>>>,
16287 remaining_ms: Arc<parking_lot::Mutex<Vec<i64>>>,
16288 }
16289
16290 struct ToolLifecycleRecordingHooks {
16292 events: parking_lot::Mutex<Vec<String>>,
16293 records: parking_lot::Mutex<Vec<ToolExecutionRecord>>,
16294 }
16295
16296 impl ToolLifecycleRecordingHooks {
16297 fn new() -> Self {
16299 Self {
16300 events: parking_lot::Mutex::new(Vec::new()),
16301 records: parking_lot::Mutex::new(Vec::new()),
16302 }
16303 }
16304
16305 fn events(&self) -> Vec<String> {
16307 self.events.lock().clone()
16308 }
16309
16310 fn records(&self) -> Vec<ToolExecutionRecord> {
16312 self.records.lock().clone()
16313 }
16314 }
16315
16316 struct ContextEchoTool;
16318
16319 #[async_trait]
16320 impl LLMProvider for RootTurnProbeProvider {
16321 async fn complete(
16322 &self,
16323 _messages: &[ChatMessage],
16324 _config: Option<&LLMConfig>,
16325 ) -> std::result::Result<LLMResponse, LLMError> {
16326 let _ = self.complete_entered.send(());
16327 Ok(LLMResponse::new("blocking complete", FinishReason::Stop))
16328 }
16329
16330 async fn complete_stream(
16331 &self,
16332 _messages: &[ChatMessage],
16333 _config: Option<&LLMConfig>,
16334 ) -> std::result::Result<
16335 Box<dyn Stream<Item = std::result::Result<LLMChunk, LLMError>> + Unpin + Send>,
16336 LLMError,
16337 > {
16338 Ok(Box::new(futures::stream::iter(vec![Ok(
16339 LLMChunk::final_chunk("stream complete", FinishReason::Stop, None),
16340 )])))
16341 }
16342
16343 fn provider_name(&self) -> &str {
16344 "root-turn-probe"
16345 }
16346
16347 fn supports(&self, feature: LLMFeature) -> bool {
16348 matches!(feature, LLMFeature::Streaming)
16349 }
16350 }
16351
16352 #[async_trait]
16353 impl ai_agents_core::Tool for ContextEchoTool {
16354 fn id(&self) -> &str {
16355 "context_echo"
16356 }
16357
16358 fn name(&self) -> &str {
16359 "Context Echo"
16360 }
16361
16362 fn description(&self) -> &str {
16363 "Returns selected execution context fields."
16364 }
16365
16366 fn input_schema(&self) -> Value {
16367 serde_json::json!({"type": "object"})
16368 }
16369
16370 fn policy_bindings(&self) -> ai_agents_core::ToolPolicyBindings {
16371 ai_agents_core::ToolPolicyBindings {
16372 path_fields: vec![ai_agents_core::PathPolicyBinding::read("path")],
16373 result_limit_fields: vec![ai_agents_core::ResultLimitBinding::new(
16374 "max_results",
16375 ai_agents_core::ResultLimitKind::MaxResults,
16376 )],
16377 ..Default::default()
16378 }
16379 }
16380
16381 async fn execute(
16382 &self,
16383 _args: Value,
16384 ctx: ai_agents_core::ToolExecutionContext,
16385 ) -> ToolResult {
16386 ToolResult::ok(
16387 serde_json::json!({
16388 "requested_name": ctx.requested_name,
16389 "canonical_id": ctx.canonical_id,
16390 "display_name": ctx.display_name,
16391 "max_results": ctx.limits.max_results,
16392 "custom_config": ctx.custom_config,
16393 })
16394 .to_string(),
16395 )
16396 }
16397 }
16398
16399 #[async_trait]
16400 impl ai_agents_core::Tool for RetryDeadlineTool {
16401 fn id(&self) -> &str {
16402 "retry_deadline"
16403 }
16404
16405 fn name(&self) -> &str {
16406 "Retry Deadline"
16407 }
16408
16409 fn description(&self) -> &str {
16410 "Records one deadline per retry invocation."
16411 }
16412
16413 fn input_schema(&self) -> Value {
16414 serde_json::json!({"type": "object"})
16415 }
16416
16417 fn safety_metadata(&self) -> ai_agents_core::ToolSafetyMetadata {
16418 ai_agents_core::ToolSafetyMetadata::compute()
16419 }
16420
16421 fn classify_call(&self, _args: &Value) -> ai_agents_core::ToolCallClassification {
16422 let mut classification =
16423 ai_agents_core::ToolCallClassification::from_metadata(&self.safety_metadata());
16424 classification.timeout_ms = Some(1_000);
16425 classification.safely_retryable = true;
16426 classification
16427 }
16428
16429 async fn execute(
16431 &self,
16432 _args: Value,
16433 ctx: ai_agents_core::ToolExecutionContext,
16434 ) -> ToolResult {
16435 let deadline = ctx
16436 .deadline
16437 .expect("each invocation must receive a deadline");
16438 self.remaining_ms.lock().push(
16439 deadline
16440 .signed_duration_since(chrono::Utc::now())
16441 .num_milliseconds(),
16442 );
16443 self.deadlines.lock().push(deadline);
16444 let call = self.calls.fetch_add(1, Ordering::SeqCst);
16445 if call == 0 {
16446 tokio::time::sleep(std::time::Duration::from_millis(20)).await;
16447 ToolResult::error("retry")
16448 } else {
16449 ToolResult::ok("done")
16450 }
16451 }
16452 }
16453
16454 struct ClassifiedTimeoutTool {
16456 id: &'static str,
16457 calls: Arc<std::sync::atomic::AtomicUsize>,
16458 timeout_ms: u64,
16459 sleep_ms: u64,
16460 requires_approval: bool,
16461 remaining_ms: Arc<parking_lot::Mutex<Vec<i64>>>,
16462 }
16463
16464 struct ApprovalModifiedTimeoutTool {
16466 calls: Arc<std::sync::atomic::AtomicUsize>,
16467 }
16468
16469 struct SlowTool;
16471
16472 struct FlakyWriteTool {
16474 calls: Arc<std::sync::atomic::AtomicUsize>,
16475 }
16476
16477 struct LockedWriteTool {
16479 active: Arc<std::sync::atomic::AtomicUsize>,
16480 max_active: Arc<std::sync::atomic::AtomicUsize>,
16481 }
16482
16483 struct MultiResourceWriteTool {
16484 active: Arc<std::sync::atomic::AtomicUsize>,
16485 max_active: Arc<std::sync::atomic::AtomicUsize>,
16486 }
16487
16488 #[derive(Clone)]
16489 struct PathMutationGate {
16490 entered: Arc<AtomicBool>,
16491 entered_notify: Arc<tokio::sync::Notify>,
16492 release: Arc<tokio::sync::Notify>,
16493 }
16494
16495 impl PathMutationGate {
16496 fn new() -> Self {
16497 Self {
16498 entered: Arc::new(AtomicBool::new(false)),
16499 entered_notify: Arc::new(tokio::sync::Notify::new()),
16500 release: Arc::new(tokio::sync::Notify::new()),
16501 }
16502 }
16503
16504 async fn wait_until_entered(&self) {
16505 if !self.entered.load(Ordering::SeqCst) {
16506 self.entered_notify.notified().await;
16507 }
16508 }
16509
16510 fn release(&self) {
16511 self.release.notify_one();
16512 }
16513 }
16514
16515 struct BlockingPathMutationTool {
16516 id: &'static str,
16517 path_fields: Vec<ai_agents_core::PathPolicyBinding>,
16518 gate: PathMutationGate,
16519 }
16520
16521 struct NoBindingWriteTool {
16522 active: Arc<std::sync::atomic::AtomicUsize>,
16523 max_active: Arc<std::sync::atomic::AtomicUsize>,
16524 }
16525
16526 struct RecoveryTestTool {
16527 id: String,
16528 succeeds: bool,
16529 calls: Arc<std::sync::atomic::AtomicUsize>,
16530 max_output_chars: Option<usize>,
16531 }
16532
16533 struct BlockingApprovalHandler {
16534 entered: Arc<tokio::sync::Barrier>,
16535 release: Arc<tokio::sync::Notify>,
16536 result: ApprovalResult,
16537 }
16538
16539 struct CountingApprovalHandler {
16540 calls: Arc<std::sync::atomic::AtomicUsize>,
16541 }
16542
16543 struct DriftingFallbackProvider {
16545 refreshed: AtomicBool,
16546 primary_calls: Arc<std::sync::atomic::AtomicUsize>,
16547 secondary_calls: Arc<std::sync::atomic::AtomicUsize>,
16548 }
16549
16550 struct RefreshFallbackProviderHooks {
16552 agent: parking_lot::Mutex<Option<Weak<RuntimeAgent>>>,
16553 lifecycle: Arc<ToolLifecycleRecordingHooks>,
16554 }
16555
16556 struct RuntimeWebFetchTransport {
16557 calls: Arc<std::sync::atomic::AtomicUsize>,
16558 }
16559
16560 struct RuntimeWebFetchResolver;
16561
16562 struct ReentrantToolHooks {
16563 agent: parking_lot::Mutex<Option<Weak<RuntimeAgent>>>,
16564 invoked: AtomicBool,
16565 nested_success: AtomicBool,
16566 }
16567
16568 #[async_trait]
16569 impl ai_agents_core::Tool for ClassifiedTimeoutTool {
16570 fn id(&self) -> &str {
16572 self.id
16573 }
16574
16575 fn name(&self) -> &str {
16577 "Classified Timeout"
16578 }
16579
16580 fn description(&self) -> &str {
16582 "Records and waits under one call-level timeout."
16583 }
16584
16585 fn input_schema(&self) -> Value {
16587 serde_json::json!({"type": "object"})
16588 }
16589
16590 fn classify_call(&self, _args: &Value) -> ai_agents_core::ToolCallClassification {
16592 let mut classification =
16593 ai_agents_core::ToolCallClassification::from_metadata(&self.safety_metadata());
16594 classification.timeout_ms = Some(self.timeout_ms);
16595 classification.requires_approval = self.requires_approval;
16596 classification
16597 }
16598
16599 async fn execute(
16601 &self,
16602 _args: Value,
16603 ctx: ai_agents_core::ToolExecutionContext,
16604 ) -> ToolResult {
16605 self.calls.fetch_add(1, Ordering::SeqCst);
16606 let deadline = ctx
16607 .deadline
16608 .expect("each invocation must receive a deadline");
16609 self.remaining_ms.lock().push(
16610 deadline
16611 .signed_duration_since(chrono::Utc::now())
16612 .num_milliseconds(),
16613 );
16614 tokio::time::sleep(Duration::from_millis(self.sleep_ms)).await;
16615 ToolResult::ok("done")
16616 }
16617 }
16618
16619 #[async_trait]
16620 impl ai_agents_core::Tool for ApprovalModifiedTimeoutTool {
16621 fn id(&self) -> &str {
16623 "approval_modified_timeout"
16624 }
16625
16626 fn name(&self) -> &str {
16628 "Approval Modified Timeout"
16629 }
16630
16631 fn description(&self) -> &str {
16633 "Becomes invalid only after approval modifies its arguments."
16634 }
16635
16636 fn input_schema(&self) -> Value {
16638 serde_json::json!({"type": "object"})
16639 }
16640
16641 fn policy_bindings(&self) -> ai_agents_core::ToolPolicyBindings {
16643 ai_agents_core::ToolPolicyBindings {
16644 path_fields: vec![ai_agents_core::PathPolicyBinding::write("path")],
16645 ..Default::default()
16646 }
16647 }
16648
16649 fn safety_metadata(&self) -> ai_agents_core::ToolSafetyMetadata {
16651 ai_agents_core::ToolSafetyMetadata {
16652 read_only: false,
16653 concurrency_safe: false,
16654 operation: ai_agents_core::ToolOperationKind::Write,
16655 side_effect_level: ai_agents_core::ToolSideEffectLevel::LocalWrite,
16656 requires_network: false,
16657 destructive: false,
16658 open_world: false,
16659 host_dependent: false,
16660 requires_user_interaction: false,
16661 supports_cancellation: true,
16662 default_requires_approval: true,
16663 should_defer_schema: false,
16664 max_output_chars: Some(1024),
16665 max_result_size_chars: Some(1024),
16666 }
16667 }
16668
16669 fn classify_call(&self, args: &Value) -> ai_agents_core::ToolCallClassification {
16671 let mut classification =
16672 ai_agents_core::ToolCallClassification::from_metadata(&self.safety_metadata());
16673 classification.timeout_ms = Some(if args["invalid_timeout"].as_bool() == Some(true) {
16674 u64::MAX
16675 } else {
16676 1_000
16677 });
16678 classification
16679 }
16680
16681 async fn execute(
16683 &self,
16684 _args: Value,
16685 _ctx: ai_agents_core::ToolExecutionContext,
16686 ) -> ToolResult {
16687 self.calls.fetch_add(1, Ordering::SeqCst);
16688 ToolResult::ok("unexpected")
16689 }
16690 }
16691
16692 #[async_trait]
16693 impl ai_agents_core::Tool for SlowTool {
16694 fn id(&self) -> &str {
16695 "slow"
16696 }
16697
16698 fn name(&self) -> &str {
16699 "Slow"
16700 }
16701
16702 fn description(&self) -> &str {
16703 "Waits until cancelled or timed out."
16704 }
16705
16706 fn input_schema(&self) -> Value {
16707 serde_json::json!({"type": "object"})
16708 }
16709
16710 async fn execute(
16711 &self,
16712 _args: Value,
16713 _ctx: ai_agents_core::ToolExecutionContext,
16714 ) -> ToolResult {
16715 tokio::time::sleep(std::time::Duration::from_secs(5)).await;
16716 ToolResult::ok("done")
16717 }
16718 }
16719
16720 #[async_trait]
16721 impl ai_agents_core::Tool for FlakyWriteTool {
16722 fn id(&self) -> &str {
16723 "flaky_write"
16724 }
16725
16726 fn name(&self) -> &str {
16727 "Flaky Write"
16728 }
16729
16730 fn description(&self) -> &str {
16731 "Fails on the first write attempt."
16732 }
16733
16734 fn input_schema(&self) -> Value {
16735 serde_json::json!({"type": "object", "properties": {"path": {"type": "string"}}})
16736 }
16737
16738 fn policy_bindings(&self) -> ai_agents_core::ToolPolicyBindings {
16739 ai_agents_core::ToolPolicyBindings {
16740 path_fields: vec![ai_agents_core::PathPolicyBinding::write("path")],
16741 ..Default::default()
16742 }
16743 }
16744
16745 fn safety_metadata(&self) -> ai_agents_core::ToolSafetyMetadata {
16746 ai_agents_core::ToolSafetyMetadata {
16747 read_only: false,
16748 concurrency_safe: false,
16749 operation: ai_agents_core::ToolOperationKind::Write,
16750 side_effect_level: ai_agents_core::ToolSideEffectLevel::LocalWrite,
16751 requires_network: false,
16752 destructive: false,
16753 open_world: false,
16754 host_dependent: false,
16755 requires_user_interaction: false,
16756 supports_cancellation: true,
16757 default_requires_approval: false,
16758 should_defer_schema: false,
16759 max_output_chars: Some(1024),
16760 max_result_size_chars: Some(1024),
16761 }
16762 }
16763
16764 fn classify_call(&self, _args: &Value) -> ai_agents_core::ToolCallClassification {
16765 let mut classification =
16766 ai_agents_core::ToolCallClassification::from_metadata(&self.safety_metadata());
16767 classification.safely_retryable = false;
16768 classification
16769 }
16770
16771 async fn execute(
16772 &self,
16773 _args: Value,
16774 _ctx: ai_agents_core::ToolExecutionContext,
16775 ) -> ToolResult {
16776 let call = self.calls.fetch_add(1, Ordering::SeqCst);
16777 if call == 0 {
16778 ToolResult::error("first failure")
16779 } else {
16780 ToolResult::ok("second success")
16781 }
16782 }
16783 }
16784
16785 #[async_trait]
16786 impl ai_agents_core::Tool for LockedWriteTool {
16787 fn id(&self) -> &str {
16788 "locked_write"
16789 }
16790
16791 fn name(&self) -> &str {
16792 "Locked Write"
16793 }
16794
16795 fn description(&self) -> &str {
16796 "Tracks concurrent execution on one resource."
16797 }
16798
16799 fn input_schema(&self) -> Value {
16800 serde_json::json!({"type": "object", "properties": {"path": {"type": "string"}}})
16801 }
16802
16803 fn policy_bindings(&self) -> ai_agents_core::ToolPolicyBindings {
16804 ai_agents_core::ToolPolicyBindings {
16805 path_fields: vec![ai_agents_core::PathPolicyBinding::write("path")],
16806 ..Default::default()
16807 }
16808 }
16809
16810 fn safety_metadata(&self) -> ai_agents_core::ToolSafetyMetadata {
16811 ai_agents_core::ToolSafetyMetadata {
16812 read_only: false,
16813 concurrency_safe: false,
16814 operation: ai_agents_core::ToolOperationKind::Write,
16815 side_effect_level: ai_agents_core::ToolSideEffectLevel::LocalWrite,
16816 requires_network: false,
16817 destructive: false,
16818 open_world: false,
16819 host_dependent: false,
16820 requires_user_interaction: false,
16821 supports_cancellation: true,
16822 default_requires_approval: false,
16823 should_defer_schema: false,
16824 max_output_chars: Some(1024),
16825 max_result_size_chars: Some(1024),
16826 }
16827 }
16828
16829 async fn execute(
16830 &self,
16831 _args: Value,
16832 _ctx: ai_agents_core::ToolExecutionContext,
16833 ) -> ToolResult {
16834 let active = self.active.fetch_add(1, Ordering::SeqCst) + 1;
16835 loop {
16836 let current_max = self.max_active.load(Ordering::SeqCst);
16837 if active <= current_max {
16838 break;
16839 }
16840 if self
16841 .max_active
16842 .compare_exchange(current_max, active, Ordering::SeqCst, Ordering::SeqCst)
16843 .is_ok()
16844 {
16845 break;
16846 }
16847 }
16848 tokio::time::sleep(std::time::Duration::from_millis(100)).await;
16849 self.active.fetch_sub(1, Ordering::SeqCst);
16850 ToolResult::ok("done")
16851 }
16852 }
16853
16854 #[async_trait]
16855 impl ai_agents_core::Tool for MultiResourceWriteTool {
16856 fn id(&self) -> &str {
16857 "multi_resource_write"
16858 }
16859
16860 fn name(&self) -> &str {
16861 "Multi Resource Write"
16862 }
16863
16864 fn description(&self) -> &str {
16865 "Tracks concurrent execution across source and destination resources."
16866 }
16867
16868 fn input_schema(&self) -> Value {
16869 serde_json::json!({"type": "object"})
16870 }
16871
16872 fn policy_bindings(&self) -> ai_agents_core::ToolPolicyBindings {
16873 ai_agents_core::ToolPolicyBindings {
16874 path_fields: vec![
16875 ai_agents_core::PathPolicyBinding::read_write("source_path"),
16876 ai_agents_core::PathPolicyBinding::write("destination_path"),
16877 ],
16878 ..Default::default()
16879 }
16880 }
16881
16882 fn safety_metadata(&self) -> ai_agents_core::ToolSafetyMetadata {
16883 LockedWriteTool {
16884 active: Arc::clone(&self.active),
16885 max_active: Arc::clone(&self.max_active),
16886 }
16887 .safety_metadata()
16888 }
16889
16890 async fn execute(
16891 &self,
16892 _args: Value,
16893 _ctx: ai_agents_core::ToolExecutionContext,
16894 ) -> ToolResult {
16895 let active = self.active.fetch_add(1, Ordering::SeqCst) + 1;
16896 self.max_active.fetch_max(active, Ordering::SeqCst);
16897 tokio::time::sleep(std::time::Duration::from_millis(75)).await;
16898 self.active.fetch_sub(1, Ordering::SeqCst);
16899 ToolResult::ok("done")
16900 }
16901 }
16902
16903 #[async_trait]
16904 impl ai_agents_core::Tool for BlockingPathMutationTool {
16905 fn id(&self) -> &str {
16906 self.id
16907 }
16908
16909 fn name(&self) -> &str {
16910 self.id
16911 }
16912
16913 fn description(&self) -> &str {
16914 "Blocks a path mutation until the test releases it."
16915 }
16916
16917 fn input_schema(&self) -> Value {
16918 serde_json::json!({"type": "object"})
16919 }
16920
16921 fn policy_bindings(&self) -> ai_agents_core::ToolPolicyBindings {
16922 ai_agents_core::ToolPolicyBindings {
16923 path_fields: self.path_fields.clone(),
16924 ..Default::default()
16925 }
16926 }
16927
16928 fn safety_metadata(&self) -> ai_agents_core::ToolSafetyMetadata {
16929 ai_agents_core::ToolSafetyMetadata {
16930 read_only: false,
16931 concurrency_safe: false,
16932 operation: ai_agents_core::ToolOperationKind::Write,
16933 side_effect_level: ai_agents_core::ToolSideEffectLevel::LocalWrite,
16934 requires_network: false,
16935 destructive: false,
16936 open_world: false,
16937 host_dependent: false,
16938 requires_user_interaction: false,
16939 supports_cancellation: true,
16940 default_requires_approval: false,
16941 should_defer_schema: false,
16942 max_output_chars: Some(1024),
16943 max_result_size_chars: Some(1024),
16944 }
16945 }
16946
16947 async fn execute(
16948 &self,
16949 _args: Value,
16950 _ctx: ai_agents_core::ToolExecutionContext,
16951 ) -> ToolResult {
16952 self.gate.entered.store(true, Ordering::SeqCst);
16953 self.gate.entered_notify.notify_one();
16954 self.gate.release.notified().await;
16955 ToolResult::ok("done")
16956 }
16957 }
16958
16959 #[async_trait]
16960 impl ai_agents_core::Tool for NoBindingWriteTool {
16961 fn id(&self) -> &str {
16962 "no_binding_write"
16963 }
16964
16965 fn name(&self) -> &str {
16966 "No Binding Write"
16967 }
16968
16969 fn description(&self) -> &str {
16970 "Tracks concurrent execution without resource bindings."
16971 }
16972
16973 fn input_schema(&self) -> Value {
16974 serde_json::json!({"type": "object"})
16975 }
16976
16977 fn safety_metadata(&self) -> ai_agents_core::ToolSafetyMetadata {
16978 LockedWriteTool {
16979 active: Arc::clone(&self.active),
16980 max_active: Arc::clone(&self.max_active),
16981 }
16982 .safety_metadata()
16983 }
16984
16985 async fn execute(
16986 &self,
16987 _args: Value,
16988 _ctx: ai_agents_core::ToolExecutionContext,
16989 ) -> ToolResult {
16990 let active = self.active.fetch_add(1, Ordering::SeqCst) + 1;
16991 self.max_active.fetch_max(active, Ordering::SeqCst);
16992 tokio::time::sleep(std::time::Duration::from_millis(75)).await;
16993 self.active.fetch_sub(1, Ordering::SeqCst);
16994 ToolResult::ok("done")
16995 }
16996 }
16997
16998 #[async_trait]
16999 impl ai_agents_core::Tool for RecoveryTestTool {
17000 fn id(&self) -> &str {
17001 &self.id
17002 }
17003
17004 fn name(&self) -> &str {
17005 &self.id
17006 }
17007
17008 fn description(&self) -> &str {
17009 "Records recovery execution and returns a configured result."
17010 }
17011
17012 fn input_schema(&self) -> Value {
17013 serde_json::json!({"type": "object"})
17014 }
17015
17016 fn policy_bindings(&self) -> ai_agents_core::ToolPolicyBindings {
17017 ai_agents_core::ToolPolicyBindings {
17018 path_fields: vec![ai_agents_core::PathPolicyBinding::write("path")],
17019 ..Default::default()
17020 }
17021 }
17022
17023 fn safety_metadata(&self) -> ai_agents_core::ToolSafetyMetadata {
17025 ai_agents_core::ToolSafetyMetadata {
17026 read_only: false,
17027 concurrency_safe: false,
17028 operation: ai_agents_core::ToolOperationKind::Write,
17029 side_effect_level: ai_agents_core::ToolSideEffectLevel::LocalWrite,
17030 requires_network: false,
17031 destructive: false,
17032 open_world: false,
17033 host_dependent: false,
17034 requires_user_interaction: false,
17035 supports_cancellation: true,
17036 default_requires_approval: false,
17037 should_defer_schema: false,
17038 max_output_chars: Some(self.max_output_chars.unwrap_or(1024)),
17039 max_result_size_chars: Some(1024),
17040 }
17041 }
17042
17043 async fn execute(
17045 &self,
17046 _args: Value,
17047 _ctx: ai_agents_core::ToolExecutionContext,
17048 ) -> ToolResult {
17049 self.calls.fetch_add(1, Ordering::SeqCst);
17050 let mut result = if self.succeeds {
17051 ToolResult::ok(format!("{} succeeded", self.id))
17052 } else {
17053 ToolResult::error(format!("{} failed", self.id))
17054 };
17055 result.metadata = Some(HashMap::from([(
17056 "recovery_test_tool".to_string(),
17057 Value::String(self.id.clone()),
17058 )]));
17059 result
17060 }
17061 }
17062
17063 #[async_trait]
17064 impl WebFetchTransport for RuntimeWebFetchTransport {
17065 async fn send(
17067 &self,
17068 _request: WebFetchTransportRequest,
17069 ) -> std::result::Result<WebFetchTransportResponse, String> {
17070 Err("validated addresses are required".to_string())
17071 }
17072
17073 async fn send_validated(
17075 &self,
17076 _request: WebFetchTransportRequest,
17077 _addresses: &[std::net::SocketAddr],
17078 ) -> std::result::Result<WebFetchTransportResponse, String> {
17079 self.calls.fetch_add(1, Ordering::SeqCst);
17080 Ok(WebFetchTransportResponse {
17081 status: 200,
17082 content_type: Some("text/plain".to_string()),
17083 location: None,
17084 body: b"approved".to_vec(),
17085 })
17086 }
17087 }
17088
17089 #[async_trait]
17090 impl WebFetchResolver for RuntimeWebFetchResolver {
17091 async fn resolve(
17093 &self,
17094 _host: &str,
17095 _port: u16,
17096 ) -> std::result::Result<Vec<std::net::IpAddr>, String> {
17097 Ok(vec![std::net::IpAddr::V4(std::net::Ipv4Addr::new(
17098 93, 184, 216, 34,
17099 ))])
17100 }
17101 }
17102
17103 #[async_trait]
17104 impl ToolProvider for DriftingFallbackProvider {
17105 fn id(&self) -> &str {
17107 "drifting_fallback"
17108 }
17109
17110 fn name(&self) -> &str {
17112 "Drifting Fallback"
17113 }
17114
17115 fn provider_type(&self) -> ToolProviderType {
17117 ToolProviderType::Custom
17118 }
17119
17120 async fn list_tools(&self) -> Vec<ToolDescriptor> {
17122 let alias = ToolAliases::new().with_name("en", "fallback alias");
17123 let mut primary = ToolDescriptor::new(
17124 "primary",
17125 "Primary",
17126 "Fails before fallback.",
17127 serde_json::json!({"type": "object"}),
17128 );
17129 let mut secondary = ToolDescriptor::new(
17130 "secondary",
17131 "Secondary",
17132 "Must not execute after final canonical drift.",
17133 serde_json::json!({"type": "object"}),
17134 );
17135 if self.refreshed.load(Ordering::SeqCst) {
17136 primary = primary.with_aliases(alias);
17137 } else {
17138 secondary = secondary.with_aliases(alias);
17139 }
17140 vec![primary, secondary]
17141 }
17142
17143 async fn get_tool(&self, tool_id: &str) -> Option<Arc<dyn Tool>> {
17145 let calls = match tool_id {
17146 "primary" => Arc::clone(&self.primary_calls),
17147 "secondary" => Arc::clone(&self.secondary_calls),
17148 _ => return None,
17149 };
17150 Some(Arc::new(RecoveryTestTool {
17151 id: tool_id.to_string(),
17152 succeeds: false,
17153 calls,
17154 max_output_chars: None,
17155 }))
17156 }
17157
17158 fn supports_refresh(&self) -> bool {
17160 true
17161 }
17162
17163 async fn refresh(&self) -> std::result::Result<(), ToolProviderError> {
17165 self.refreshed.store(true, Ordering::SeqCst);
17166 Ok(())
17167 }
17168 }
17169
17170 #[async_trait]
17171 impl AgentHooks for RefreshFallbackProviderHooks {
17172 async fn on_tool_start(&self, tool: &str, args: &Value) {
17174 self.lifecycle.on_tool_start(tool, args).await;
17175 if tool != "secondary" {
17176 return;
17177 }
17178 let agent = self.agent.lock().as_ref().and_then(Weak::upgrade);
17179 if let Some(agent) = agent {
17180 agent
17181 .tools
17182 .refresh_provider("drifting_fallback")
17183 .await
17184 .unwrap();
17185 }
17186 }
17187
17188 async fn on_tool_complete(&self, tool: &str, result: &ToolResult, duration_ms: u64) {
17189 self.lifecycle
17190 .on_tool_complete(tool, result, duration_ms)
17191 .await;
17192 }
17193
17194 async fn on_tool_execution_record(&self, record: &ToolExecutionRecord) {
17195 self.lifecycle.on_tool_execution_record(record).await;
17196 }
17197
17198 async fn on_error(&self, error: &AgentError) {
17199 self.lifecycle.on_error(error).await;
17200 }
17201 }
17202
17203 #[async_trait]
17204 impl ApprovalHandler for BlockingApprovalHandler {
17205 async fn request_approval(
17206 &self,
17207 _request: ai_agents_hitl::ApprovalRequest,
17208 ) -> ApprovalResult {
17209 self.entered.wait().await;
17210 self.release.notified().await;
17211 self.result.clone()
17212 }
17213 }
17214
17215 #[async_trait]
17216 impl ApprovalHandler for CountingApprovalHandler {
17217 async fn request_approval(
17218 &self,
17219 _request: ai_agents_hitl::ApprovalRequest,
17220 ) -> ApprovalResult {
17221 self.calls.fetch_add(1, Ordering::SeqCst);
17222 ApprovalResult::Approved
17223 }
17224 }
17225
17226 #[async_trait]
17227 impl AgentHooks for ReentrantToolHooks {
17228 async fn on_tool_complete(&self, tool: &str, _result: &ToolResult, _duration_ms: u64) {
17229 if tool != "reentrant_write" || self.invoked.swap(true, Ordering::SeqCst) {
17230 return;
17231 }
17232 let agent = self.agent.lock().as_ref().and_then(Weak::upgrade);
17233 if let Some(agent) = agent {
17234 let result = agent
17235 .invoke_tool(ToolExecutionRequest::new(
17236 "nested-hook-call",
17237 "reentrant_write",
17238 serde_json::json!({"path": "./hook.txt"}),
17239 ToolCallSource::Manual,
17240 ))
17241 .await;
17242 self.nested_success
17243 .store(result.is_ok_and(|record| record.success), Ordering::SeqCst);
17244 }
17245 }
17246 }
17247
17248 #[async_trait]
17249 impl AgentHooks for ResponseCountingHooks {
17250 async fn on_response(&self, _response: &AgentResponse) {
17251 self.responses.fetch_add(1, Ordering::SeqCst);
17252 }
17253 }
17254
17255 #[async_trait]
17256 impl AgentHooks for ResponseChatHooks {
17257 async fn on_response(&self, _response: &AgentResponse) {
17259 if self.invoked.swap(true, Ordering::SeqCst) {
17260 return;
17261 }
17262 let target = self.target.lock().as_ref().and_then(Weak::upgrade);
17263 let result = if let Some(target) = target {
17264 target
17265 .chat("nested response hook call")
17266 .await
17267 .map(|response| response.content)
17268 .map_err(|error| error.to_string())
17269 } else {
17270 Err("response hook target is unavailable".to_string())
17271 };
17272 *self.nested_result.lock() = Some(result);
17273 }
17274 }
17275
17276 #[async_trait]
17277 impl AgentHooks for ConcurrentResponseHooks {
17278 async fn on_response(&self, _response: &AgentResponse) {
17280 if self.invoked.swap(true, Ordering::SeqCst) {
17281 return;
17282 }
17283 let Some(registry) = self.registry.upgrade() else {
17284 *self.nested_result.lock() =
17285 Some(Err("concurrent registry is unavailable".to_string()));
17286 return;
17287 };
17288 let agents = [ai_agents_state::ConcurrentAgentRef::Id(
17289 self.child_id.clone(),
17290 )];
17291 let aggregation = ai_agents_state::AggregationConfig {
17292 strategy: ai_agents_state::AggregationStrategy::FirstWins,
17293 synthesizer_llm: None,
17294 synthesizer_prompt: None,
17295 vote: None,
17296 };
17297 let result = crate::orchestration::concurrent(
17298 ®istry,
17299 "nested concurrent response hook call",
17300 &agents,
17301 &aggregation,
17302 None,
17303 Some(1),
17304 None,
17305 ai_agents_state::PartialFailureAction::Abort,
17306 None,
17307 )
17308 .await
17309 .map(|result| result.response.content)
17310 .map_err(|error| error.to_string());
17311 *self.nested_result.lock() = Some(result);
17312 }
17313 }
17314
17315 #[async_trait]
17316 impl AgentHooks for ToolLifecycleRecordingHooks {
17317 async fn on_tool_start(&self, tool: &str, _args: &Value) {
17318 self.events.lock().push(format!("start:{tool}"));
17319 }
17320
17321 async fn on_tool_complete(&self, tool: &str, result: &ToolResult, _duration_ms: u64) {
17322 self.events
17323 .lock()
17324 .push(format!("complete:{tool}:{}", result.success));
17325 }
17326
17327 async fn on_tool_execution_record(&self, record: &ToolExecutionRecord) {
17328 self.events.lock().push(format!(
17329 "record:{}:{}",
17330 record.canonical_id, record.executed
17331 ));
17332 self.records.lock().push(record.clone());
17333 }
17334
17335 async fn on_error(&self, _error: &AgentError) {
17337 self.events.lock().push("error".to_string());
17338 }
17339 }
17340
17341 struct ApprovalRecordingHooks {
17342 events: parking_lot::Mutex<Vec<String>>,
17343 }
17344
17345 impl ApprovalRecordingHooks {
17346 fn new() -> Self {
17347 Self {
17348 events: parking_lot::Mutex::new(Vec::new()),
17349 }
17350 }
17351
17352 fn events(&self) -> Vec<String> {
17353 self.events.lock().clone()
17354 }
17355 }
17356
17357 #[async_trait]
17358 impl AgentHooks for ApprovalRecordingHooks {
17359 async fn on_approval_result(&self, request_id: &str, result: &ApprovalResult) {
17360 self.events.lock().push(format!(
17361 "raw:{}:{}",
17362 request_id,
17363 approval_result_name(result)
17364 ));
17365 }
17366
17367 async fn on_approval_resolved(
17368 &self,
17369 request: &ai_agents_hitl::ApprovalRequest,
17370 raw_result: &ApprovalResult,
17371 outcome: &ApprovalResolvedOutcome,
17372 ) {
17373 self.events.lock().push(format!(
17374 "resolved:{}:{}:{}",
17375 request.id,
17376 approval_result_name(raw_result),
17377 approval_outcome_name(outcome)
17378 ));
17379 }
17380 }
17381
17382 fn approval_result_name(result: &ApprovalResult) -> &'static str {
17383 match result {
17384 ApprovalResult::Approved => "approved",
17385 ApprovalResult::Rejected { .. } => "rejected",
17386 ApprovalResult::Modified { .. } => "modified",
17387 ApprovalResult::Timeout => "timeout",
17388 }
17389 }
17390
17391 fn approval_outcome_name(outcome: &ApprovalResolvedOutcome) -> &'static str {
17392 match outcome {
17393 ApprovalResolvedOutcome::Approved => "approved",
17394 ApprovalResolvedOutcome::Rejected { .. } => "rejected",
17395 ApprovalResolvedOutcome::Modified { .. } => "modified",
17396 ApprovalResolvedOutcome::Error { .. } => "error",
17397 }
17398 }
17399
17400 fn assert_correlated_approval_events(
17401 events: &[String],
17402 raw_status: &str,
17403 outcome_status: &str,
17404 ) {
17405 assert_eq!(events.len(), 2);
17406 let raw: Vec<_> = events[0].split(':').collect();
17407 let resolved: Vec<_> = events[1].split(':').collect();
17408 assert_eq!(raw[0], "raw");
17409 assert_eq!(resolved[0], "resolved");
17410 assert_eq!(raw[1], resolved[1]);
17411 assert_eq!(raw[2], raw_status);
17412 assert_eq!(resolved[2], raw_status);
17413 assert_eq!(resolved[3], outcome_status);
17414 }
17415
17416 fn approval_security_config(policy_enabled: bool) -> ToolSecurityConfig {
17417 let mut security = ToolSecurityConfig {
17418 enabled: true,
17419 fail_closed: true,
17420 ..Default::default()
17421 };
17422 let policy = ai_agents_tools::ToolPolicyConfig {
17423 enabled: policy_enabled,
17424 write_paths: vec![".".to_string()],
17425 require_confirmation: true,
17426 ..Default::default()
17427 };
17428 security.tools.insert("locked_write".to_string(), policy);
17429 security
17430 }
17431
17432 struct MutationTestWorkspace {
17433 root: std::path::PathBuf,
17434 }
17435
17436 impl MutationTestWorkspace {
17437 fn new() -> Self {
17438 let root = std::env::temp_dir().join(format!(
17439 "ai-agents-runtime-mutation-{}",
17440 uuid::Uuid::new_v4()
17441 ));
17442 std::fs::create_dir_all(&root).unwrap();
17443 Self { root }
17444 }
17445 }
17446
17447 impl Drop for MutationTestWorkspace {
17448 fn drop(&mut self) {
17449 let _ = std::fs::remove_dir_all(&self.root);
17450 }
17451 }
17452
17453 async fn wait_for_resource_lock_strong_count(locks: &ToolResourceLocks, minimum: usize) {
17454 tokio::time::timeout(std::time::Duration::from_secs(2), async {
17455 loop {
17456 let strong_count = locks
17457 .read()
17458 .get("path-mutation:global")
17459 .map_or(0, |lock| lock.strong_count());
17460 if strong_count >= minimum {
17461 break;
17462 }
17463 tokio::task::yield_now().await;
17464 }
17465 })
17466 .await
17467 .expect("path mutation call did not reach the shared lock");
17468 }
17469
17470 async fn assert_path_mutation_pair_serialized(
17471 first_id: &'static str,
17472 first_fields: Vec<ai_agents_core::PathPolicyBinding>,
17473 first_args: Value,
17474 second_id: &'static str,
17475 second_fields: Vec<ai_agents_core::PathPolicyBinding>,
17476 second_args: Value,
17477 ) {
17478 let locks = new_tool_resource_locks();
17479 let first_gate = PathMutationGate::new();
17480 let second_gate = PathMutationGate::new();
17481 second_gate.release();
17482 let agent = Arc::new(
17483 AgentBuilder::new()
17484 .system_prompt("Test global path mutation locking.")
17485 .llm(Arc::new(mock_with_response("done")))
17486 .tool(Arc::new(BlockingPathMutationTool {
17487 id: first_id,
17488 path_fields: first_fields,
17489 gate: first_gate.clone(),
17490 }))
17491 .tool(Arc::new(BlockingPathMutationTool {
17492 id: second_id,
17493 path_fields: second_fields,
17494 gate: second_gate.clone(),
17495 }))
17496 .build()
17497 .unwrap()
17498 .with_shared_resource_locks(Arc::clone(&locks)),
17499 );
17500
17501 let first = {
17502 let agent = Arc::clone(&agent);
17503 tokio::spawn(async move {
17504 agent
17505 .invoke_tool(ToolExecutionRequest::new(
17506 format!("{}-first", first_id),
17507 first_id,
17508 first_args,
17509 ToolCallSource::Manual,
17510 ))
17511 .await
17512 .unwrap()
17513 })
17514 };
17515 first_gate.wait_until_entered().await;
17516
17517 let second = {
17518 let agent = Arc::clone(&agent);
17519 tokio::spawn(async move {
17520 agent
17521 .invoke_tool(ToolExecutionRequest::new(
17522 format!("{}-second", second_id),
17523 second_id,
17524 second_args,
17525 ToolCallSource::Manual,
17526 ))
17527 .await
17528 .unwrap()
17529 })
17530 };
17531 wait_for_resource_lock_strong_count(&locks, 2).await;
17532 assert!(!second_gate.entered.load(Ordering::SeqCst));
17533 assert!(!second.is_finished());
17534
17535 first_gate.release();
17536 let (first, second) = tokio::time::timeout(std::time::Duration::from_secs(2), async {
17537 tokio::join!(first, second)
17538 })
17539 .await
17540 .expect("serialized path mutation calls did not finish");
17541 assert!(first.unwrap().success);
17542 assert!(second.unwrap().success);
17543 assert!(second_gate.entered.load(Ordering::SeqCst));
17544 assert!(locks.read().is_empty());
17545 }
17546
17547 #[derive(Clone, Copy)]
17548 enum MutationDenial {
17549 Policy,
17550 Approval,
17551 }
17552
17553 fn mutation_denial_security_config(
17554 tool_id: &str,
17555 workspace: &std::path::Path,
17556 denial: MutationDenial,
17557 ) -> ToolSecurityConfig {
17558 let workspace = workspace.to_string_lossy().into_owned();
17559 let mut policy = ai_agents_tools::ToolPolicyConfig {
17560 read_paths: vec![workspace.clone()],
17561 write_paths: vec![workspace.clone()],
17562 ..Default::default()
17563 };
17564 match denial {
17565 MutationDenial::Policy => policy.blocked_paths = vec![workspace],
17566 MutationDenial::Approval => policy.require_confirmation = true,
17567 }
17568
17569 let mut security = ToolSecurityConfig {
17570 enabled: true,
17571 fail_closed: true,
17572 ..Default::default()
17573 };
17574 security.tools.insert(tool_id.to_string(), policy);
17575 security
17576 }
17577
17578 async fn assert_path_mutation_denied(tool: Arc<dyn Tool>, denial: MutationDenial) {
17579 let workspace = MutationTestWorkspace::new();
17580 let tool_id = tool.id().to_string();
17581 let preserved = workspace.root.join(format!("{}-preserved.txt", tool_id));
17582 let destination = workspace.root.join(format!("{}-destination.txt", tool_id));
17583 std::fs::write(&preserved, "preserved").unwrap();
17584 let arguments = match tool_id.as_str() {
17585 "copy_path" | "move_path" => serde_json::json!({
17586 "source_path": preserved.to_string_lossy(),
17587 "destination_path": destination.to_string_lossy(),
17588 "dry_run": false
17589 }),
17590 "delete_path" => serde_json::json!({
17591 "path": preserved.to_string_lossy(),
17592 "recursive": false,
17593 "dry_run": false
17594 }),
17595 _ => panic!("unsupported mutation tool: {}", tool_id),
17596 };
17597 let security = mutation_denial_security_config(&tool_id, &workspace.root, denial);
17598 let builder = AgentBuilder::new()
17599 .system_prompt("Test mutation denial.")
17600 .llm(Arc::new(mock_with_response("done")))
17601 .tool(tool)
17602 .tool_security(ToolSecurityEngine::new(security));
17603 let builder = match denial {
17604 MutationDenial::Policy => builder,
17605 MutationDenial::Approval => builder
17606 .hitl_engine(HITLEngine::new(ai_agents_hitl::HITLConfig::default()))
17607 .approval_handler(Arc::new(RejectAllHandler::new())),
17608 };
17609 let agent = builder.build().unwrap();
17610
17611 let record = agent
17612 .invoke_tool(ToolExecutionRequest::new(
17613 format!("{}-denied", tool_id),
17614 tool_id.clone(),
17615 arguments,
17616 ToolCallSource::Manual,
17617 ))
17618 .await
17619 .unwrap();
17620
17621 assert!(!record.executed, "{} must not be invoked", tool_id);
17622 assert!(!record.success);
17623 match denial {
17624 MutationDenial::Policy => {
17625 assert_eq!(record.policy.outcome, PermissionOutcome::Deny);
17626 assert!(record.approval.as_ref().is_some_and(|approval| matches!(
17627 &approval.status,
17628 ToolApprovalStatus::NotRequired
17629 )));
17630 }
17631 MutationDenial::Approval => {
17632 assert_eq!(record.policy.outcome, PermissionOutcome::RequiresApproval);
17633 assert!(record.approval.as_ref().is_some_and(|approval| matches!(
17634 &approval.status,
17635 ToolApprovalStatus::Rejected
17636 )));
17637 }
17638 }
17639 assert_eq!(std::fs::read_to_string(&preserved).unwrap(), "preserved");
17640 assert!(!destination.exists());
17641 }
17642
17643 fn recovery_manager_with_fallbacks(
17644 fallbacks: impl IntoIterator<Item = (String, String)>,
17645 ) -> RecoveryManager {
17646 use ai_agents_recovery::{ErrorRecoveryConfig, ToolRecoveryConfig, ToolRetryConfig};
17647
17648 let per_tool = fallbacks
17649 .into_iter()
17650 .map(|(tool, fallback_tool)| {
17651 (
17652 tool,
17653 ToolRetryConfig {
17654 max_retries: 0,
17655 timeout_ms: Some(1_000),
17656 on_failure: ToolFailureAction::Fallback { fallback_tool },
17657 },
17658 )
17659 })
17660 .collect();
17661 RecoveryManager::new(ErrorRecoveryConfig {
17662 tools: ToolRecoveryConfig {
17663 per_tool,
17664 ..Default::default()
17665 },
17666 ..Default::default()
17667 })
17668 }
17669
17670 fn approval_check() -> HITLCheckResult {
17671 HITLCheckResult::required(
17672 ApprovalTrigger::tool("test", serde_json::json!({})),
17673 HashMap::new(),
17674 "Approve?",
17675 None,
17676 )
17677 }
17678
17679 fn agent_with_approval_result(
17680 raw_result: ApprovalResult,
17681 timeout_action: TimeoutAction,
17682 hooks: Arc<ApprovalRecordingHooks>,
17683 ) -> RuntimeAgent {
17684 use ai_agents_hitl::{CallbackHandler, HITLConfig};
17685
17686 let config = HITLConfig {
17687 on_timeout: timeout_action,
17688 ..Default::default()
17689 };
17690 let handler = CallbackHandler::new(move |_| raw_result.clone());
17691 AgentBuilder::new()
17692 .system_prompt("Test HITL hooks.")
17693 .llm(Arc::new(mock_with_response("done")))
17694 .build()
17695 .unwrap()
17696 .with_hooks(hooks)
17697 .with_hitl(HITLEngine::new(config), Arc::new(handler))
17698 }
17699
17700 #[tokio::test]
17701 async fn approval_hooks_expose_direct_effective_decisions_after_raw_results() {
17702 let cases = vec![
17703 (ApprovalResult::Approved, "approved"),
17704 (
17705 ApprovalResult::Rejected {
17706 reason: Some("denied".to_string()),
17707 },
17708 "rejected",
17709 ),
17710 (
17711 ApprovalResult::Modified {
17712 changes: HashMap::from([("value".to_string(), serde_json::json!(2))]),
17713 },
17714 "modified",
17715 ),
17716 ];
17717
17718 for (raw_result, expected) in cases {
17719 let hooks = Arc::new(ApprovalRecordingHooks::new());
17720 let agent =
17721 agent_with_approval_result(raw_result, TimeoutAction::Reject, hooks.clone());
17722
17723 let result = agent.request_hitl_approval(approval_check()).await.unwrap();
17724
17725 assert_eq!(approval_result_name(&result), expected);
17726 assert_correlated_approval_events(&hooks.events(), expected, expected);
17727 }
17728 }
17729
17730 #[tokio::test]
17731 async fn approval_hooks_expose_timeout_policy_decisions() {
17732 for (timeout_action, expected) in [
17733 (TimeoutAction::Approve, "approved"),
17734 (TimeoutAction::Reject, "rejected"),
17735 ] {
17736 let hooks = Arc::new(ApprovalRecordingHooks::new());
17737 let agent =
17738 agent_with_approval_result(ApprovalResult::Timeout, timeout_action, hooks.clone());
17739
17740 let result = agent.request_hitl_approval(approval_check()).await.unwrap();
17741
17742 assert_eq!(approval_result_name(&result), expected);
17743 assert_correlated_approval_events(&hooks.events(), "timeout", expected);
17744 }
17745 }
17746
17747 #[tokio::test]
17748 async fn timeout_error_fires_correlated_resolved_error_before_returning() {
17749 let hooks = Arc::new(ApprovalRecordingHooks::new());
17750 let agent = agent_with_approval_result(
17751 ApprovalResult::Timeout,
17752 TimeoutAction::Error,
17753 hooks.clone(),
17754 );
17755
17756 let error = agent
17757 .request_hitl_approval(approval_check())
17758 .await
17759 .unwrap_err();
17760
17761 assert!(error.to_string().contains("HITL approval timeout"));
17762 assert_correlated_approval_events(&hooks.events(), "timeout", "error");
17763 }
17764
17765 #[tokio::test]
17767 async fn test_integration_yaml_to_chat_basic() {
17768 let mock = mock_with_response("Hello! How can I help you?");
17769 let agent = AgentBuilder::new()
17770 .system_prompt("You are a test assistant.")
17771 .llm(Arc::new(mock))
17772 .build()
17773 .unwrap();
17774
17775 let response = agent.chat("Hi").await.unwrap();
17776 assert!(!response.content.is_empty());
17777 assert_eq!(response.content, "Hello! How can I help you?");
17778 }
17779
17780 #[tokio::test]
17781 async fn stream_events_emit_one_authoritative_final_without_legacy_done() {
17782 let agent = AgentBuilder::new()
17783 .system_prompt("You are a test assistant.")
17784 .llm(Arc::new(mock_with_response(
17785 "Hello from the final response.",
17786 )))
17787 .build()
17788 .unwrap();
17789
17790 let mut stream = agent.chat_stream_events("Hi").await.unwrap();
17791 let mut final_responses = Vec::new();
17792 let mut legacy_done = 0;
17793 while let Some(event) = stream.next().await {
17794 match event {
17795 AgentStreamEvent::Chunk(StreamChunk::Done {}) => legacy_done += 1,
17796 AgentStreamEvent::Chunk(StreamChunk::Error { message }) => {
17797 panic!("unexpected stream error: {message}")
17798 }
17799 AgentStreamEvent::Final(response) => final_responses.push(response),
17800 AgentStreamEvent::Chunk(_) => {}
17801 }
17802 }
17803
17804 assert_eq!(legacy_done, 0);
17805 assert_eq!(final_responses.len(), 1);
17806 let response = final_responses.pop().unwrap();
17807 assert_eq!(response.content, "Hello from the final response.");
17808 assert!(
17809 response
17810 .metadata
17811 .as_ref()
17812 .is_some_and(|metadata| { metadata.contains_key("reasoning") })
17813 );
17814 }
17815
17816 #[tokio::test]
17817 async fn stream_final_content_includes_output_processing_after_provisional_chunks() {
17818 let yaml = r#"
17819name: ProcessedStreamAgent
17820system_prompt: "Answer directly."
17821process:
17822 output:
17823 - type: format
17824 config:
17825 template: "{{ response }} [finalized]"
17826streaming:
17827 enabled: true
17828"#;
17829 let agent = AgentBuilder::from_yaml(yaml)
17830 .unwrap()
17831 .llm(Arc::new(mock_with_response("provisional answer")))
17832 .auto_configure_features()
17833 .unwrap()
17834 .build()
17835 .unwrap();
17836
17837 let mut stream = agent.chat_stream_events("Hi").await.unwrap();
17838 let mut provisional = String::new();
17839 let mut final_content = None;
17840 while let Some(event) = stream.next().await {
17841 match event {
17842 AgentStreamEvent::Chunk(StreamChunk::Content { text }) => {
17843 provisional.push_str(&text)
17844 }
17845 AgentStreamEvent::Chunk(StreamChunk::Error { message }) => {
17846 panic!("unexpected stream error: {message}")
17847 }
17848 AgentStreamEvent::Final(response) => final_content = Some(response.content),
17849 AgentStreamEvent::Chunk(_) => {}
17850 }
17851 }
17852
17853 assert_eq!(provisional, "provisional answer");
17854 assert_eq!(
17855 final_content.as_deref(),
17856 Some("provisional answer [finalized]")
17857 );
17858 }
17859
17860 #[tokio::test]
17861 async fn stream_events_preserve_tool_progress_and_final_tool_calls() {
17862 let agent = AgentBuilder::new()
17863 .system_prompt("Use the echo tool once, then answer.")
17864 .llm(Arc::new(mock_with_responses(vec![
17865 r#"{"tool":"echo","arguments":{"message":"hello"}}"#,
17866 "Echo completed.",
17867 ])))
17868 .tool(Arc::new(ai_agents_tools::EchoTool::new()))
17869 .build()
17870 .unwrap();
17871
17872 let mut stream = agent.chat_stream_events("echo hello").await.unwrap();
17873 let mut starts = 0;
17874 let mut results = 0;
17875 let mut ends = 0;
17876 let mut final_response = None;
17877 while let Some(event) = stream.next().await {
17878 match event {
17879 AgentStreamEvent::Chunk(StreamChunk::ToolCallStart { name, .. }) => {
17880 assert_eq!(name, "echo");
17881 starts += 1;
17882 }
17883 AgentStreamEvent::Chunk(StreamChunk::ToolResult { name, success, .. }) => {
17884 assert_eq!(name, "echo");
17885 assert!(success);
17886 results += 1;
17887 }
17888 AgentStreamEvent::Chunk(StreamChunk::ToolCallEnd { .. }) => ends += 1,
17889 AgentStreamEvent::Chunk(StreamChunk::Error { message }) => {
17890 panic!("unexpected stream error: {message}")
17891 }
17892 AgentStreamEvent::Final(response) => final_response = Some(response),
17893 AgentStreamEvent::Chunk(_) => {}
17894 }
17895 }
17896
17897 assert_eq!((starts, results, ends), (1, 1, 1));
17898 let response = final_response.expect("tool stream must finalize");
17899 assert_eq!(response.content, "Echo completed.");
17900 assert_eq!(
17901 response.tool_calls.as_ref().map(|calls| calls
17902 .iter()
17903 .map(|call| call.name.as_str())
17904 .collect::<Vec<_>>()),
17905 Some(vec!["echo"])
17906 );
17907 }
17908
17909 #[tokio::test]
17910 async fn legacy_stream_still_emits_one_done_chunk() {
17911 let agent = AgentBuilder::new()
17912 .system_prompt("You are a test assistant.")
17913 .llm(Arc::new(mock_with_response(
17914 "Hello from the legacy stream.",
17915 )))
17916 .build()
17917 .unwrap();
17918
17919 let mut stream = agent.chat_stream("Hi").await.unwrap();
17920 let mut done = 0;
17921 while let Some(chunk) = stream.next().await {
17922 match chunk {
17923 StreamChunk::Done {} => done += 1,
17924 StreamChunk::Error { message } => panic!("unexpected stream error: {message}"),
17925 _ => {}
17926 }
17927 }
17928
17929 assert_eq!(done, 1);
17930 }
17931
17932 #[tokio::test]
17934 async fn test_integration_multi_turn_conversation() {
17935 let mock = mock_with_responses(vec![
17936 "Hello! I'm your assistant.",
17937 "The weather is sunny today.",
17938 "Goodbye!",
17939 ]);
17940 let agent = AgentBuilder::new()
17941 .system_prompt("You are helpful.")
17942 .llm(Arc::new(mock))
17943 .build()
17944 .unwrap();
17945
17946 let r1 = agent.chat("Hi").await.unwrap();
17947 assert_eq!(r1.content, "Hello! I'm your assistant.");
17948
17949 let r2 = agent.chat("What's the weather?").await.unwrap();
17950 assert_eq!(r2.content, "The weather is sunny today.");
17951
17952 let r3 = agent.chat("Bye").await.unwrap();
17953 assert_eq!(r3.content, "Goodbye!");
17954
17955 let messages = agent.memory.get_messages(None).await.unwrap();
17957 assert_eq!(messages.len(), 6);
17959 }
17960
17961 #[test]
17962 fn later_approval_preserves_modified_evidence() {
17963 let arguments = serde_json::json!({"dry_run": true});
17964 let mut record = Some(ToolApprovalRecord {
17965 status: ToolApprovalStatus::Modified,
17966 reason: None,
17967 modified_arguments: Some(arguments.clone()),
17968 });
17969
17970 merge_approved_record(&mut record);
17971
17972 let record = record.unwrap();
17973 assert!(matches!(record.status, ToolApprovalStatus::Modified));
17974 assert_eq!(record.modified_arguments, Some(arguments));
17975 }
17976
17977 #[test]
17978 fn approval_binding_rejects_replaced_tool_implementation() {
17979 let reviewed_tool: Arc<dyn ai_agents_core::Tool> = Arc::new(ContextEchoTool);
17980 let same_tool = Arc::clone(&reviewed_tool);
17981 let replacement_tool: Arc<dyn ai_agents_core::Tool> = Arc::new(ContextEchoTool);
17982 let arguments = serde_json::json!({"path": "."});
17983 let versions = ToolDecisionVersions {
17984 policy: 2,
17985 registry: 3,
17986 runtime_control: 4,
17987 state: Some(5),
17988 };
17989 let binding = ToolApprovalBinding {
17990 canonical_id: "context_echo".to_string(),
17991 arguments: arguments.clone(),
17992 confirmation_required: true,
17993 policy_version: versions.policy,
17994 runtime_control_version: versions.runtime_control,
17995 state_generation: versions.state,
17996 reviewed_tool,
17997 };
17998
17999 assert!(!binding.is_stale("context_echo", &arguments, true, versions, &same_tool,));
18000 assert!(binding.is_stale(
18001 "context_echo",
18002 &arguments,
18003 true,
18004 versions,
18005 &replacement_tool,
18006 ));
18007 }
18008
18009 #[tokio::test]
18010 async fn approved_mutation_to_dry_run_remains_executable() {
18011 use ai_agents_hitl::CallbackHandler;
18012
18013 let handler = CallbackHandler::new(|_| ApprovalResult::Modified {
18014 changes: HashMap::from([("dry_run".to_string(), serde_json::json!(true))]),
18015 });
18016 let agent = AgentBuilder::new()
18017 .system_prompt("Test safer approval modifications.")
18018 .llm(Arc::new(mock_with_response("done")))
18019 .tool(Arc::new(ai_agents_tools::FileWriteTool::new()))
18020 .hitl_engine(HITLEngine::new(ai_agents_hitl::HITLConfig::default()))
18021 .approval_handler(Arc::new(handler))
18022 .build()
18023 .unwrap();
18024
18025 let record = agent
18026 .invoke_tool(ToolExecutionRequest::new(
18027 "approved-dry-run",
18028 "file_write",
18029 serde_json::json!({
18030 "path": "./approval-dry-run.txt",
18031 "content": "not written"
18032 }),
18033 ToolCallSource::Manual,
18034 ))
18035 .await
18036 .unwrap();
18037
18038 assert!(record.executed);
18039 assert!(record.success);
18040 assert_eq!(record.executed_arguments["dry_run"], true);
18041 assert!(matches!(
18042 record.approval.as_ref().map(|approval| &approval.status),
18043 Some(ToolApprovalStatus::Modified)
18044 ));
18045 let output: Value = serde_json::from_str(&record.output).unwrap();
18046 assert_eq!(output["mutation_performed"], false);
18047 }
18048
18049 #[tokio::test]
18051 async fn shared_executor_approval_reaches_web_fetch_transport() {
18052 use ai_agents_hitl::{CallbackHandler, HITLConfig};
18053 use ai_agents_tools::{DomainPolicyConfig, ToolPolicyConfig};
18054
18055 let calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
18056 let tool = WebFetchTool::with_transport_and_resolver(
18057 Arc::new(RuntimeWebFetchTransport {
18058 calls: Arc::clone(&calls),
18059 }),
18060 Arc::new(RuntimeWebFetchResolver),
18061 );
18062 let mut security = ToolSecurityConfig {
18063 enabled: true,
18064 fail_closed: true,
18065 ..Default::default()
18066 };
18067 security.tools.insert(
18068 "web_fetch".to_string(),
18069 ToolPolicyConfig {
18070 domains: DomainPolicyConfig {
18071 requires_approval: vec!["approval.test".to_string()],
18072 ..Default::default()
18073 },
18074 allowed_schemes: vec!["https".to_string()],
18075 allowed_ports: vec![443],
18076 ..Default::default()
18077 },
18078 );
18079 let handler = CallbackHandler::new(|_| ApprovalResult::Approved);
18080 let agent = AgentBuilder::new()
18081 .system_prompt("Test approved web fetch execution.")
18082 .llm(Arc::new(mock_with_response("done")))
18083 .tool(Arc::new(tool))
18084 .tool_security(ToolSecurityEngine::new(security))
18085 .build()
18086 .unwrap()
18087 .with_hitl(HITLEngine::new(HITLConfig::default()), Arc::new(handler));
18088
18089 let record = agent
18090 .invoke_tool(ToolExecutionRequest::new(
18091 "approved-web-fetch",
18092 "web_fetch",
18093 serde_json::json!({
18094 "url": "https://approval.test/page",
18095 "cache_ttl_seconds": 0
18096 }),
18097 ToolCallSource::Manual,
18098 ))
18099 .await
18100 .unwrap();
18101
18102 assert!(record.success);
18103 assert!(
18104 record
18105 .approval
18106 .as_ref()
18107 .is_some_and(|approval| matches!(approval.status, ToolApprovalStatus::Approved))
18108 );
18109 assert_eq!(calls.load(Ordering::SeqCst), 1);
18110 }
18111
18112 #[tokio::test]
18113 async fn context_preserves_requested_and_canonical_identity() {
18114 let mock = mock_with_response("hello");
18115 let mut tools = ai_agents_tools::ToolRegistry::new();
18116 tools.register(Arc::new(ContextEchoTool)).unwrap();
18117
18118 let mut security = ToolSecurityConfig {
18119 enabled: true,
18120 fail_closed: true,
18121 ..Default::default()
18122 };
18123 let mut policy = ai_agents_tools::ToolPolicyConfig {
18124 read_paths: vec![".".to_string()],
18125 max_results: Some(7),
18126 ..Default::default()
18127 };
18128 policy
18129 .config
18130 .insert("backend".to_string(), serde_json::json!("memory"));
18131 security.tools.insert("context_echo".to_string(), policy);
18132
18133 let agent = AgentBuilder::new()
18134 .system_prompt("You are helpful.")
18135 .llm(Arc::new(mock))
18136 .tools(tools)
18137 .tool_security(ToolSecurityEngine::new(security))
18138 .build()
18139 .unwrap();
18140
18141 let record = agent
18142 .invoke_tool(ToolExecutionRequest::new(
18143 "ctx-call",
18144 "Context Echo",
18145 serde_json::json!({"path": ".", "max_results": 99}),
18146 ToolCallSource::Manual,
18147 ))
18148 .await
18149 .unwrap();
18150
18151 assert!(record.success);
18152 assert!(matches!(&record.source, ToolCallSource::Manual));
18153 assert_eq!(record.requested_name, "Context Echo");
18154 assert_eq!(record.canonical_id, "context_echo");
18155 assert_eq!(record.policy.outcome, PermissionOutcome::Allow);
18156 assert_eq!(record.executed_arguments["max_results"], 7);
18157 let output: Value = serde_json::from_str(&record.output).unwrap();
18158 assert_eq!(output["requested_name"], "Context Echo");
18159 assert_eq!(output["canonical_id"], "context_echo");
18160 assert_eq!(output["max_results"], 7);
18161 assert_eq!(output["custom_config"]["backend"], "memory");
18162 assert!(record.metadata.contains_key("effective_limits"));
18163 assert!(record.metadata.contains_key("policy_snapshot"));
18164 }
18165
18166 #[tokio::test]
18167 async fn test_runtime_control_cancels_active_tool_call() {
18168 let mock = mock_with_response("hello");
18169 let agent = Arc::new(
18170 AgentBuilder::new()
18171 .system_prompt("You are helpful.")
18172 .llm(Arc::new(mock))
18173 .tool(Arc::new(SlowTool))
18174 .build()
18175 .unwrap(),
18176 );
18177 let control = agent.runtime_control();
18178 let running_agent = Arc::clone(&agent);
18179 let handle = tokio::spawn(async move {
18180 running_agent
18181 .invoke_tool(ToolExecutionRequest::new(
18182 "slow-call",
18183 "slow",
18184 serde_json::json!({}),
18185 ToolCallSource::Manual,
18186 ))
18187 .await
18188 .unwrap()
18189 });
18190
18191 tokio::time::sleep(std::time::Duration::from_millis(100)).await;
18192 control.cancel_all();
18193 let record = handle.await.unwrap();
18194
18195 assert!(record.executed);
18196 assert!(record.cancelled);
18197 assert!(!record.success);
18198 assert!(record.cancellation_reason.is_some());
18199 }
18200
18201 #[tokio::test]
18203 async fn cancelled_tool_does_not_enter_fallback() {
18204 let fallback_calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
18205 let agent = Arc::new(
18206 AgentBuilder::new()
18207 .system_prompt("Test cancellation before fallback.")
18208 .llm(Arc::new(mock_with_response("done")))
18209 .tool(Arc::new(SlowTool))
18210 .tool(Arc::new(RecoveryTestTool {
18211 id: "fallback".to_string(),
18212 succeeds: true,
18213 calls: Arc::clone(&fallback_calls),
18214 max_output_chars: None,
18215 }))
18216 .recovery_manager(recovery_manager_with_fallbacks([(
18217 "slow".to_string(),
18218 "fallback".to_string(),
18219 )]))
18220 .build()
18221 .unwrap(),
18222 );
18223 let control = agent.runtime_control();
18224 let running_agent = Arc::clone(&agent);
18225 let handle = tokio::spawn(async move {
18226 running_agent
18227 .invoke_tool(ToolExecutionRequest::new(
18228 "cancelled-fallback-call",
18229 "slow",
18230 serde_json::json!({}),
18231 ToolCallSource::Manual,
18232 ))
18233 .await
18234 .unwrap()
18235 });
18236
18237 tokio::time::sleep(Duration::from_millis(100)).await;
18238 control.cancel_all();
18239 let record = handle.await.unwrap();
18240
18241 assert!(record.executed);
18242 assert!(record.cancelled);
18243 assert!(!record.success);
18244 assert_eq!(record.canonical_id, "slow");
18245 assert_eq!(fallback_calls.load(Ordering::SeqCst), 0);
18246 assert_eq!(agent.tool_call_history().len(), 1);
18247 }
18248
18249 #[tokio::test]
18250 async fn non_idempotent_tool_calls_are_not_retried() {
18251 use ai_agents_recovery::{ErrorRecoveryConfig, ToolRecoveryConfig, ToolRetryConfig};
18252
18253 let mock = mock_with_response("hello");
18254 let calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
18255 let agent = AgentBuilder::new()
18256 .system_prompt("You are helpful.")
18257 .llm(Arc::new(mock))
18258 .tool(Arc::new(FlakyWriteTool {
18259 calls: Arc::clone(&calls),
18260 }))
18261 .recovery_manager(RecoveryManager::new(ErrorRecoveryConfig {
18262 tools: ToolRecoveryConfig {
18263 default: ToolRetryConfig {
18264 max_retries: 2,
18265 ..Default::default()
18266 },
18267 ..Default::default()
18268 },
18269 ..Default::default()
18270 }))
18271 .build()
18272 .unwrap();
18273
18274 let record = agent
18275 .invoke_tool(ToolExecutionRequest::new(
18276 "flaky-call",
18277 "flaky_write",
18278 serde_json::json!({"path": "./tmp.txt"}),
18279 ToolCallSource::Manual,
18280 ))
18281 .await
18282 .unwrap();
18283
18284 assert!(!record.success);
18285 assert_eq!(calls.load(Ordering::SeqCst), 1);
18286 }
18287
18288 #[tokio::test]
18289 async fn safely_retryable_tool_receives_a_fresh_deadline_per_attempt() {
18290 use ai_agents_recovery::{ErrorRecoveryConfig, ToolRecoveryConfig, ToolRetryConfig};
18291
18292 let calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
18293 let deadlines = Arc::new(parking_lot::Mutex::new(Vec::new()));
18294 let remaining_ms = Arc::new(parking_lot::Mutex::new(Vec::new()));
18295 let agent = AgentBuilder::new()
18296 .system_prompt("Test retry deadlines.")
18297 .llm(Arc::new(mock_with_response("done")))
18298 .tool(Arc::new(RetryDeadlineTool {
18299 calls: Arc::clone(&calls),
18300 deadlines: Arc::clone(&deadlines),
18301 remaining_ms: Arc::clone(&remaining_ms),
18302 }))
18303 .recovery_manager(RecoveryManager::new(ErrorRecoveryConfig {
18304 tools: ToolRecoveryConfig {
18305 per_tool: HashMap::from([(
18306 "retry_deadline".to_string(),
18307 ToolRetryConfig {
18308 max_retries: 1,
18309 ..Default::default()
18310 },
18311 )]),
18312 ..Default::default()
18313 },
18314 ..Default::default()
18315 }))
18316 .build()
18317 .unwrap();
18318
18319 let record = agent
18320 .invoke_tool(ToolExecutionRequest::new(
18321 "retry-deadline-call",
18322 "retry_deadline",
18323 serde_json::json!({}),
18324 ToolCallSource::Manual,
18325 ))
18326 .await
18327 .unwrap();
18328
18329 assert!(record.executed);
18330 assert!(record.success);
18331 assert_eq!(calls.load(Ordering::SeqCst), 2);
18332 let deadlines = deadlines.lock();
18333 assert_eq!(deadlines.len(), 2);
18334 assert!(
18335 deadlines[1] > deadlines[0],
18336 "retry inherited the first invocation deadline"
18337 );
18338 let remaining_ms = remaining_ms.lock();
18339 assert_eq!(remaining_ms.len(), 2);
18340 assert!(
18341 remaining_ms
18342 .iter()
18343 .all(|remaining| (800..=1_000).contains(remaining))
18344 );
18345 }
18346
18347 #[tokio::test]
18349 async fn call_classification_timeout_controls_deadline_and_timer() {
18350 let calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
18351 let remaining_ms = Arc::new(parking_lot::Mutex::new(Vec::new()));
18352 let agent = AgentBuilder::new()
18353 .system_prompt("Test call-level timeout.")
18354 .llm(Arc::new(mock_with_response("done")))
18355 .tool(Arc::new(ClassifiedTimeoutTool {
18356 id: "classified_timeout",
18357 calls: Arc::clone(&calls),
18358 timeout_ms: 100,
18359 sleep_ms: 150,
18360 requires_approval: false,
18361 remaining_ms: Arc::clone(&remaining_ms),
18362 }))
18363 .build()
18364 .unwrap();
18365
18366 let started = Instant::now();
18367 let record = agent
18368 .invoke_tool(ToolExecutionRequest::new(
18369 "classified-timeout-call",
18370 "classified_timeout",
18371 serde_json::json!({}),
18372 ToolCallSource::Manual,
18373 ))
18374 .await
18375 .unwrap();
18376
18377 assert!(record.executed);
18378 assert!(record.timed_out);
18379 assert!(!record.success);
18380 assert_eq!(calls.load(Ordering::SeqCst), 1);
18381 assert!(started.elapsed() < Duration::from_secs(1));
18382 let remaining_ms = remaining_ms.lock();
18383 assert_eq!(remaining_ms.len(), 1);
18384 assert!((1..=100).contains(&remaining_ms[0]));
18385 }
18386
18387 #[tokio::test]
18389 async fn recovery_timeout_only_lowers_call_and_policy_timeouts() {
18390 use ai_agents_recovery::{ErrorRecoveryConfig, ToolRecoveryConfig, ToolRetryConfig};
18391
18392 let calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
18393 let remaining_ms = Arc::new(parking_lot::Mutex::new(Vec::new()));
18394 let agent = AgentBuilder::new()
18395 .system_prompt("Test recovery timeout.")
18396 .llm(Arc::new(mock_with_response("done")))
18397 .tool(Arc::new(ClassifiedTimeoutTool {
18398 id: "recovery_timeout",
18399 calls: Arc::clone(&calls),
18400 timeout_ms: 1_000,
18401 sleep_ms: 150,
18402 requires_approval: false,
18403 remaining_ms: Arc::clone(&remaining_ms),
18404 }))
18405 .recovery_manager(RecoveryManager::new(ErrorRecoveryConfig {
18406 tools: ToolRecoveryConfig {
18407 per_tool: HashMap::from([(
18408 "recovery_timeout".to_string(),
18409 ToolRetryConfig {
18410 timeout_ms: Some(100),
18411 ..Default::default()
18412 },
18413 )]),
18414 ..Default::default()
18415 },
18416 ..Default::default()
18417 }))
18418 .build()
18419 .unwrap();
18420
18421 let started = Instant::now();
18422 let record = agent
18423 .invoke_tool(ToolExecutionRequest::new(
18424 "recovery-timeout-call",
18425 "recovery_timeout",
18426 serde_json::json!({}),
18427 ToolCallSource::Manual,
18428 ))
18429 .await
18430 .unwrap();
18431
18432 assert!(record.executed);
18433 assert!(record.timed_out);
18434 assert!(!record.success);
18435 assert_eq!(calls.load(Ordering::SeqCst), 1);
18436 assert!(started.elapsed() < Duration::from_secs(1));
18437 assert_eq!(record.metadata["effective_limits"]["timeout_ms"], 100);
18438 let remaining_ms = remaining_ms.lock();
18439 assert_eq!(remaining_ms.len(), 1);
18440 assert!((1..=100).contains(&remaining_ms[0]));
18441 }
18442
18443 #[tokio::test]
18445 async fn recovery_default_timeout_controls_deadline_and_timer() {
18446 use ai_agents_recovery::{ErrorRecoveryConfig, ToolRecoveryConfig, ToolRetryConfig};
18447
18448 let calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
18449 let remaining_ms = Arc::new(parking_lot::Mutex::new(Vec::new()));
18450 let agent = AgentBuilder::new()
18451 .system_prompt("Test default recovery timeout.")
18452 .llm(Arc::new(mock_with_response("done")))
18453 .tool(Arc::new(ClassifiedTimeoutTool {
18454 id: "default_recovery_timeout",
18455 calls: Arc::clone(&calls),
18456 timeout_ms: 1_000,
18457 sleep_ms: 150,
18458 requires_approval: false,
18459 remaining_ms: Arc::clone(&remaining_ms),
18460 }))
18461 .recovery_manager(RecoveryManager::new(ErrorRecoveryConfig {
18462 tools: ToolRecoveryConfig {
18463 default: ToolRetryConfig {
18464 timeout_ms: Some(100),
18465 ..Default::default()
18466 },
18467 ..Default::default()
18468 },
18469 ..Default::default()
18470 }))
18471 .build()
18472 .unwrap();
18473
18474 let started = Instant::now();
18475 let record = agent
18476 .invoke_tool(ToolExecutionRequest::new(
18477 "default-recovery-timeout-call",
18478 "default_recovery_timeout",
18479 serde_json::json!({}),
18480 ToolCallSource::Manual,
18481 ))
18482 .await
18483 .unwrap();
18484
18485 assert!(record.executed);
18486 assert!(record.timed_out);
18487 assert!(!record.success);
18488 assert_eq!(calls.load(Ordering::SeqCst), 1);
18489 assert!(started.elapsed() < Duration::from_secs(1));
18490 assert_eq!(record.metadata["effective_limits"]["timeout_ms"], 100);
18491 let remaining_ms = remaining_ms.lock();
18492 assert_eq!(remaining_ms.len(), 1);
18493 assert!((1..=100).contains(&remaining_ms[0]));
18494 }
18495
18496 #[test]
18498 fn recovery_timeout_cannot_widen_security_baseline() {
18499 let security_engine = ToolSecurityEngine::new(ToolSecurityConfig {
18500 default_timeout_ms: 100,
18501 ..Default::default()
18502 });
18503 let safety = ToolSafetyMetadata::compute();
18504 let mut classification = ToolCallClassification::from_metadata(&safety);
18505 classification.timeout_ms = Some(500);
18506
18507 let (limits, timeout) = RuntimeAgent::effective_tool_limits(
18508 &security_engine,
18509 "recovery_cannot_widen",
18510 &safety,
18511 &classification,
18512 Some(1_000),
18513 )
18514 .unwrap();
18515
18516 assert_eq!(limits.timeout_ms, Some(100));
18517 assert_eq!(timeout.timer, Duration::from_millis(100));
18518 }
18519
18520 #[tokio::test]
18522 async fn invalid_call_timeout_stops_before_approval_or_tool_invocation() {
18523 let tool_calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
18524 let approval_calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
18525 let remaining_ms = Arc::new(parking_lot::Mutex::new(Vec::new()));
18526 let mut security = ToolSecurityConfig {
18527 enabled: true,
18528 ..Default::default()
18529 };
18530 security.tools.insert(
18531 "invalid_call_timeout".to_string(),
18532 ai_agents_tools::ToolPolicyConfig {
18533 require_confirmation: true,
18534 ..Default::default()
18535 },
18536 );
18537 let agent = AgentBuilder::new()
18538 .system_prompt("Test invalid call timeout.")
18539 .llm(Arc::new(mock_with_response("done")))
18540 .tool(Arc::new(ClassifiedTimeoutTool {
18541 id: "invalid_call_timeout",
18542 calls: Arc::clone(&tool_calls),
18543 timeout_ms: u64::MAX,
18544 sleep_ms: 0,
18545 requires_approval: false,
18546 remaining_ms,
18547 }))
18548 .tool_security(ToolSecurityEngine::new(security))
18549 .hitl_engine(HITLEngine::new(ai_agents_hitl::HITLConfig::default()))
18550 .approval_handler(Arc::new(CountingApprovalHandler {
18551 calls: Arc::clone(&approval_calls),
18552 }))
18553 .build()
18554 .unwrap();
18555
18556 let error = agent
18557 .invoke_tool(ToolExecutionRequest::new(
18558 "invalid-call-timeout",
18559 "invalid_call_timeout",
18560 serde_json::json!({}),
18561 ToolCallSource::Manual,
18562 ))
18563 .await
18564 .unwrap_err();
18565
18566 assert!(error.to_string().contains(
18567 "effective tool timeout_ms must be no greater than 3153600000000000 milliseconds"
18568 ));
18569 assert_eq!(approval_calls.load(Ordering::SeqCst), 0);
18570 assert_eq!(tool_calls.load(Ordering::SeqCst), 0);
18571 }
18572
18573 #[tokio::test]
18575 async fn invalid_modified_call_timeout_stops_before_lock_or_invocation() {
18576 use ai_agents_hitl::CallbackHandler;
18577
18578 let blocker_gate = PathMutationGate::new();
18579 let tool_calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
18580 let hooks = Arc::new(ToolLifecycleRecordingHooks::new());
18581 let handler = CallbackHandler::new(|_| ApprovalResult::Modified {
18582 changes: HashMap::from([("invalid_timeout".to_string(), Value::Bool(true))]),
18583 });
18584 let agent = Arc::new(
18585 AgentBuilder::new()
18586 .system_prompt("Test final call timeout validation.")
18587 .llm(Arc::new(mock_with_response("done")))
18588 .tool(Arc::new(BlockingPathMutationTool {
18589 id: "timeout_lock_blocker",
18590 path_fields: vec![ai_agents_core::PathPolicyBinding::write("path")],
18591 gate: blocker_gate.clone(),
18592 }))
18593 .tool(Arc::new(ApprovalModifiedTimeoutTool {
18594 calls: Arc::clone(&tool_calls),
18595 }))
18596 .hitl_engine(HITLEngine::new(ai_agents_hitl::HITLConfig::default()))
18597 .approval_handler(Arc::new(handler))
18598 .hooks(hooks.clone())
18599 .build()
18600 .unwrap(),
18601 );
18602 let blocking_agent = Arc::clone(&agent);
18603 let blocker = tokio::spawn(async move {
18604 blocking_agent
18605 .invoke_tool(ToolExecutionRequest::new(
18606 "timeout-lock-blocker",
18607 "timeout_lock_blocker",
18608 serde_json::json!({"path": "./shared-timeout.txt"}),
18609 ToolCallSource::Manual,
18610 ))
18611 .await
18612 .unwrap()
18613 });
18614 blocker_gate.wait_until_entered().await;
18615
18616 let record = tokio::time::timeout(
18617 Duration::from_millis(500),
18618 agent.invoke_tool(ToolExecutionRequest::new(
18619 "invalid-modified-timeout",
18620 "approval_modified_timeout",
18621 serde_json::json!({
18622 "path": "./shared-timeout.txt",
18623 "invalid_timeout": false
18624 }),
18625 ToolCallSource::Manual,
18626 )),
18627 )
18628 .await
18629 .expect("final timeout validation must not wait for the held path lock")
18630 .unwrap();
18631
18632 blocker_gate.release();
18633 assert!(blocker.await.unwrap().success);
18634 assert!(!record.executed);
18635 assert!(!record.success);
18636 assert_eq!(record.policy.outcome, PermissionOutcome::Deny);
18637 assert!(record.output.contains(
18638 "effective tool timeout_ms must be no greater than 3153600000000000 milliseconds"
18639 ));
18640 assert_eq!(tool_calls.load(Ordering::SeqCst), 0);
18641 let invalid_request_events = hooks
18642 .events()
18643 .into_iter()
18644 .filter(|event| event.contains("approval_modified_timeout") || event == "error")
18645 .collect::<Vec<_>>();
18646 assert_eq!(
18647 invalid_request_events,
18648 vec![
18649 "start:approval_modified_timeout",
18650 "complete:approval_modified_timeout:false",
18651 "record:approval_modified_timeout:false",
18652 "error"
18653 ]
18654 );
18655 }
18656
18657 #[tokio::test]
18658 async fn side_effecting_tools_are_serialized_per_resource() {
18659 let mock = mock_with_response("hello");
18660 let active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
18661 let max_active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
18662 let agent = Arc::new(
18663 AgentBuilder::new()
18664 .system_prompt("You are helpful.")
18665 .llm(Arc::new(mock))
18666 .tool(Arc::new(LockedWriteTool {
18667 active: Arc::clone(&active),
18668 max_active: Arc::clone(&max_active),
18669 }))
18670 .build()
18671 .unwrap(),
18672 );
18673
18674 let left = {
18675 let agent = Arc::clone(&agent);
18676 tokio::spawn(async move {
18677 agent
18678 .invoke_tool(ToolExecutionRequest::new(
18679 "lock-1",
18680 "locked_write",
18681 serde_json::json!({"path": "./same.txt"}),
18682 ToolCallSource::Manual,
18683 ))
18684 .await
18685 .unwrap()
18686 })
18687 };
18688 let right = {
18689 let agent = Arc::clone(&agent);
18690 tokio::spawn(async move {
18691 agent
18692 .invoke_tool(ToolExecutionRequest::new(
18693 "lock-2",
18694 "locked_write",
18695 serde_json::json!({"path": "./same.txt"}),
18696 ToolCallSource::Manual,
18697 ))
18698 .await
18699 .unwrap()
18700 })
18701 };
18702
18703 let left = left.await.unwrap();
18704 let right = right.await.unwrap();
18705 assert!(left.success);
18706 assert!(right.success);
18707 assert_eq!(max_active.load(Ordering::SeqCst), 1);
18708 }
18709
18710 #[tokio::test]
18711 async fn path_resources_use_shared_global_lock_and_cleanup() {
18712 let active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
18713 let max_active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
18714 let bindings = ai_agents_core::ToolPolicyBindings {
18715 path_fields: vec![
18716 ai_agents_core::PathPolicyBinding::read_write("source_path"),
18717 ai_agents_core::PathPolicyBinding::write("destination_path"),
18718 ],
18719 ..Default::default()
18720 };
18721 let classification = ai_agents_core::ToolCallClassification::from_metadata(
18722 &MultiResourceWriteTool {
18723 active: Arc::clone(&active),
18724 max_active: Arc::clone(&max_active),
18725 }
18726 .safety_metadata(),
18727 );
18728 let left_args = serde_json::json!({
18729 "source_path": "./a/../first.txt",
18730 "destination_path": "./second.txt"
18731 });
18732 let right_args = serde_json::json!({
18733 "source_path": "./second.txt",
18734 "destination_path": "./first.txt"
18735 });
18736 let left_keys = tool_resource_lock_keys(
18737 "multi_resource_write",
18738 &left_args,
18739 &bindings,
18740 &classification,
18741 );
18742 let right_keys = tool_resource_lock_keys(
18743 "multi_resource_write",
18744 &right_args,
18745 &bindings,
18746 &classification,
18747 );
18748 assert_eq!(left_keys, right_keys);
18749 assert_eq!(left_keys, vec!["path-mutation:global".to_string()]);
18750
18751 let locks = new_tool_resource_locks();
18752 let build_agent = || {
18753 AgentBuilder::new()
18754 .system_prompt("Test shared resource locks.")
18755 .llm(Arc::new(mock_with_response("done")))
18756 .tool(Arc::new(MultiResourceWriteTool {
18757 active: Arc::clone(&active),
18758 max_active: Arc::clone(&max_active),
18759 }))
18760 .build()
18761 .unwrap()
18762 .with_shared_resource_locks(Arc::clone(&locks))
18763 };
18764 let left_agent = Arc::new(build_agent());
18765 let right_agent = Arc::new(build_agent());
18766 let left = tokio::spawn(async move {
18767 left_agent
18768 .invoke_tool(ToolExecutionRequest::new(
18769 "multi-left",
18770 "multi_resource_write",
18771 left_args,
18772 ToolCallSource::Manual,
18773 ))
18774 .await
18775 .unwrap()
18776 });
18777 let right = tokio::spawn(async move {
18778 right_agent
18779 .invoke_tool(ToolExecutionRequest::new(
18780 "multi-right",
18781 "multi_resource_write",
18782 right_args,
18783 ToolCallSource::Manual,
18784 ))
18785 .await
18786 .unwrap()
18787 });
18788 let (left, right) = tokio::time::timeout(std::time::Duration::from_secs(2), async {
18789 tokio::join!(left, right)
18790 })
18791 .await
18792 .expect("reversed resource acquisition must not deadlock");
18793
18794 assert!(left.unwrap().success);
18795 assert!(right.unwrap().success);
18796 assert_eq!(max_active.load(Ordering::SeqCst), 1);
18797 assert!(locks.read().is_empty());
18798 }
18799
18800 #[tokio::test]
18801 async fn global_path_lock_serializes_copy_destination_with_file_write() {
18802 assert_path_mutation_pair_serialized(
18803 "copy_path",
18804 CopyPathTool::new().policy_bindings().path_fields,
18805 serde_json::json!({
18806 "source_path": "./source.txt",
18807 "destination_path": "./shared.txt"
18808 }),
18809 "file_write",
18810 FileWriteTool::new().policy_bindings().path_fields,
18811 serde_json::json!({"path": "./shared.txt"}),
18812 )
18813 .await;
18814 }
18815
18816 #[tokio::test]
18817 async fn parent_and_spawned_runtime_share_global_path_lock() {
18818 let workspace = MutationTestWorkspace::new();
18819 let destination = workspace.root.join("spawned.txt");
18820 let parent_gate = PathMutationGate::new();
18821 let parent = Arc::new(
18822 AgentBuilder::from_yaml(
18823 r#"
18824name: LockParent
18825system_prompt: parent
18826llm:
18827 default: default
18828tools:
18829 - parent_path_write
18830spawner:
18831 shared_llms: true
18832"#,
18833 )
18834 .unwrap()
18835 .llm(Arc::new(mock_with_response("done")))
18836 .auto_configure_spawner()
18837 .await
18838 .unwrap()
18839 .tool(Arc::new(BlockingPathMutationTool {
18840 id: "parent_path_write",
18841 path_fields: vec![ai_agents_core::PathPolicyBinding::write("path")],
18842 gate: parent_gate.clone(),
18843 }))
18844 .build()
18845 .unwrap(),
18846 );
18847
18848 let mut child_spec = crate::spec::AgentSpec {
18849 name: "LockChild".to_string(),
18850 system_prompt: "child".to_string(),
18851 tools: Some(vec![crate::spec::ToolEntry::Simple(
18852 "file_write".to_string(),
18853 )]),
18854 ..Default::default()
18855 };
18856 child_spec.tool_security.enabled = true;
18857 child_spec.tool_security.fail_closed = true;
18858 let file_write_policy = ai_agents_tools::ToolPolicyConfig {
18859 write_paths: vec![workspace.root.to_string_lossy().into_owned()],
18860 allow_without_confirmation: true,
18861 ..Default::default()
18862 };
18863 child_spec
18864 .tool_security
18865 .tools
18866 .insert("file_write".to_string(), file_write_policy);
18867 let spawned = parent
18868 .spawner()
18869 .unwrap()
18870 .spawn_from_spec(child_spec)
18871 .await
18872 .unwrap();
18873 assert!(Arc::ptr_eq(
18874 &parent.resource_locks,
18875 &spawned.agent.resource_locks
18876 ));
18877 assert!(!Arc::ptr_eq(
18878 &parent.runtime_control,
18879 &spawned.agent.runtime_control
18880 ));
18881
18882 let parent_call = {
18883 let parent = Arc::clone(&parent);
18884 let destination = destination.clone();
18885 tokio::spawn(async move {
18886 parent
18887 .invoke_tool(ToolExecutionRequest::new(
18888 "parent-lock-holder",
18889 "parent_path_write",
18890 serde_json::json!({"path": destination}),
18891 ToolCallSource::Manual,
18892 ))
18893 .await
18894 .unwrap()
18895 })
18896 };
18897 parent_gate.wait_until_entered().await;
18898
18899 let child_call = {
18900 let child = Arc::clone(&spawned.agent);
18901 let destination = destination.clone();
18902 tokio::spawn(async move {
18903 child
18904 .invoke_tool(ToolExecutionRequest::new(
18905 "spawned-file-write",
18906 "file_write",
18907 serde_json::json!({
18908 "path": destination,
18909 "content": "spawned",
18910 "dry_run": false
18911 }),
18912 ToolCallSource::Manual,
18913 ))
18914 .await
18915 .unwrap()
18916 })
18917 };
18918 wait_for_resource_lock_strong_count(&parent.resource_locks, 2).await;
18919 assert!(!child_call.is_finished());
18920
18921 parent_gate.release();
18922 let (parent_record, child_record) =
18923 tokio::time::timeout(std::time::Duration::from_secs(2), async {
18924 tokio::join!(parent_call, child_call)
18925 })
18926 .await
18927 .expect("parent and spawned path mutations did not finish");
18928 assert!(parent_record.unwrap().success);
18929 assert!(child_record.unwrap().success);
18930 assert_eq!(std::fs::read_to_string(destination).unwrap(), "spawned");
18931 assert!(parent.resource_locks.read().is_empty());
18932 }
18933
18934 #[tokio::test]
18935 async fn cancelled_global_path_lock_waiter_does_not_retain_weak_entry() {
18936 let locks = new_tool_resource_locks();
18937 let holder_gate = PathMutationGate::new();
18938 let waiter_gate = PathMutationGate::new();
18939 waiter_gate.release();
18940 let holder = Arc::new(
18941 AgentBuilder::new()
18942 .system_prompt("Hold the global path lock.")
18943 .llm(Arc::new(mock_with_response("done")))
18944 .tool(Arc::new(BlockingPathMutationTool {
18945 id: "holder_write",
18946 path_fields: vec![ai_agents_core::PathPolicyBinding::write("path")],
18947 gate: holder_gate.clone(),
18948 }))
18949 .build()
18950 .unwrap()
18951 .with_shared_resource_locks(Arc::clone(&locks)),
18952 );
18953 let waiter = Arc::new(
18954 AgentBuilder::new()
18955 .system_prompt("Wait for the global path lock.")
18956 .llm(Arc::new(mock_with_response("done")))
18957 .tool(Arc::new(BlockingPathMutationTool {
18958 id: "waiter_write",
18959 path_fields: vec![ai_agents_core::PathPolicyBinding::write("path")],
18960 gate: waiter_gate.clone(),
18961 }))
18962 .build()
18963 .unwrap()
18964 .with_shared_resource_locks(Arc::clone(&locks)),
18965 );
18966
18967 let holder_call = {
18968 let holder = Arc::clone(&holder);
18969 tokio::spawn(async move {
18970 holder
18971 .invoke_tool(ToolExecutionRequest::new(
18972 "holder-call",
18973 "holder_write",
18974 serde_json::json!({"path": "./shared.txt"}),
18975 ToolCallSource::Manual,
18976 ))
18977 .await
18978 .unwrap()
18979 })
18980 };
18981 holder_gate.wait_until_entered().await;
18982
18983 let waiter_call = {
18984 let waiter = Arc::clone(&waiter);
18985 tokio::spawn(async move {
18986 waiter
18987 .invoke_tool(ToolExecutionRequest::new(
18988 "waiter-call",
18989 "waiter_write",
18990 serde_json::json!({"path": "./shared.txt"}),
18991 ToolCallSource::Manual,
18992 ))
18993 .await
18994 .unwrap()
18995 })
18996 };
18997 wait_for_resource_lock_strong_count(&locks, 2).await;
18998 waiter.runtime_control().cancel_all();
18999
19000 let waiter_record = tokio::time::timeout(std::time::Duration::from_secs(2), waiter_call)
19001 .await
19002 .expect("cancelled lock waiter did not finish")
19003 .unwrap();
19004 assert!(!waiter_record.success);
19005 assert!(!waiter_record.executed);
19006 assert!(waiter_record.cancelled);
19007 assert_eq!(
19008 waiter_record.cancellation_reason.as_deref(),
19009 Some("runtime control cancellation")
19010 );
19011 assert!(!waiter_gate.entered.load(Ordering::SeqCst));
19012 assert_eq!(
19013 locks
19014 .read()
19015 .get("path-mutation:global")
19016 .map_or(0, |lock| lock.strong_count()),
19017 1
19018 );
19019
19020 holder_gate.release();
19021 let holder_record = tokio::time::timeout(std::time::Duration::from_secs(2), holder_call)
19022 .await
19023 .expect("lock holder did not finish")
19024 .unwrap();
19025 assert!(holder_record.success);
19026 assert!(locks.read().is_empty());
19027 }
19028
19029 #[tokio::test]
19030 async fn path_mutation_policy_and_approval_denials_do_not_invoke_tools() {
19031 for denial in [MutationDenial::Policy, MutationDenial::Approval] {
19032 let tools: [Arc<dyn Tool>; 3] = [
19033 Arc::new(CopyPathTool::new()),
19034 Arc::new(MovePathTool::new()),
19035 Arc::new(DeletePathTool::new()),
19036 ];
19037 for tool in tools {
19038 assert_path_mutation_denied(tool, denial).await;
19039 }
19040 }
19041 }
19042
19043 #[tokio::test]
19044 async fn policy_denial_keeps_executor_hook_lifecycle_and_record_authority() {
19045 let workspace = MutationTestWorkspace::new();
19046 let target = workspace.root.join("denied.txt");
19047 let hooks = Arc::new(ToolLifecycleRecordingHooks::new());
19048 let agent = AgentBuilder::new()
19049 .system_prompt("Test denied tool hooks.")
19050 .llm(Arc::new(mock_with_response("done")))
19051 .tool(Arc::new(FileWriteTool::new()))
19052 .tool_security(ToolSecurityEngine::new(mutation_denial_security_config(
19053 "file_write",
19054 &workspace.root,
19055 MutationDenial::Policy,
19056 )))
19057 .hooks(hooks.clone())
19058 .build()
19059 .unwrap();
19060
19061 let record = agent
19062 .invoke_tool(ToolExecutionRequest::new(
19063 "denied-hook-call",
19064 "file_write",
19065 serde_json::json!({
19066 "path": target.to_string_lossy(),
19067 "content": "blocked"
19068 }),
19069 ToolCallSource::Manual,
19070 ))
19071 .await
19072 .unwrap();
19073
19074 assert!(!record.executed);
19075 assert!(!record.success);
19076 assert_eq!(record.policy.outcome, PermissionOutcome::Deny);
19077 assert_eq!(
19078 hooks.events(),
19079 vec![
19080 "start:file_write",
19081 "complete:file_write:false",
19082 "record:file_write:false",
19083 "error"
19084 ]
19085 );
19086 assert!(!target.exists());
19087 }
19088
19089 #[tokio::test]
19090 async fn approval_argument_changes_are_rechecked_against_final_scope() {
19091 let active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19092 let max_active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19093 let entered = Arc::new(tokio::sync::Barrier::new(2));
19094 let release = Arc::new(tokio::sync::Notify::new());
19095 let handler = Arc::new(BlockingApprovalHandler {
19096 entered: Arc::clone(&entered),
19097 release: Arc::clone(&release),
19098 result: ApprovalResult::Modified {
19099 changes: HashMap::from([(
19100 "path".to_string(),
19101 Value::String("./after-approval.txt".to_string()),
19102 )]),
19103 },
19104 });
19105 let agent = Arc::new(
19106 AgentBuilder::new()
19107 .system_prompt("Test final scope validation.")
19108 .llm(Arc::new(mock_with_response("done")))
19109 .tool(Arc::new(LockedWriteTool {
19110 active: Arc::clone(&active),
19111 max_active: Arc::clone(&max_active),
19112 }))
19113 .tool_security(ToolSecurityEngine::new(approval_security_config(true)))
19114 .hitl_engine(HITLEngine::new(ai_agents_hitl::HITLConfig::default()))
19115 .approval_handler(handler)
19116 .build()
19117 .unwrap(),
19118 );
19119 let control = agent.runtime_control();
19120 let running = Arc::clone(&agent);
19121 let call = tokio::spawn(async move {
19122 running
19123 .invoke_tool(ToolExecutionRequest::new(
19124 "approval-scope",
19125 "locked_write",
19126 serde_json::json!({"path": "./before-approval.txt"}),
19127 ToolCallSource::Manual,
19128 ))
19129 .await
19130 .unwrap()
19131 });
19132 entered.wait().await;
19133 let expected_version = control.set_tool_scope(Vec::new());
19134 release.notify_one();
19135 let record = call.await.unwrap();
19136
19137 assert!(!record.executed);
19138 assert!(!record.success);
19139 assert_eq!(record.runtime_config_version, expected_version);
19140 assert_eq!(record.executed_arguments["path"], "./after-approval.txt");
19141 assert_eq!(max_active.load(Ordering::SeqCst), 0);
19142 assert_eq!(
19143 record.metadata["runtime_scope_snapshot"],
19144 serde_json::json!([])
19145 );
19146 }
19147
19148 #[tokio::test]
19149 async fn approval_is_rechecked_against_final_policy_snapshot() {
19150 let active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19151 let max_active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19152 let entered = Arc::new(tokio::sync::Barrier::new(2));
19153 let release = Arc::new(tokio::sync::Notify::new());
19154 let handler = Arc::new(BlockingApprovalHandler {
19155 entered: Arc::clone(&entered),
19156 release: Arc::clone(&release),
19157 result: ApprovalResult::Approved,
19158 });
19159 let agent = Arc::new(
19160 AgentBuilder::new()
19161 .system_prompt("Test final policy validation.")
19162 .llm(Arc::new(mock_with_response("done")))
19163 .tool(Arc::new(LockedWriteTool {
19164 active: Arc::clone(&active),
19165 max_active: Arc::clone(&max_active),
19166 }))
19167 .tool_security(ToolSecurityEngine::new(approval_security_config(true)))
19168 .hitl_engine(HITLEngine::new(ai_agents_hitl::HITLConfig::default()))
19169 .approval_handler(handler)
19170 .build()
19171 .unwrap(),
19172 );
19173 let control = agent.runtime_control();
19174 let running = Arc::clone(&agent);
19175 let call = tokio::spawn(async move {
19176 running
19177 .invoke_tool(ToolExecutionRequest::new(
19178 "approval-policy",
19179 "locked_write",
19180 serde_json::json!({"path": "./policy.txt"}),
19181 ToolCallSource::Manual,
19182 ))
19183 .await
19184 .unwrap()
19185 });
19186 entered.wait().await;
19187 let expected_version = control.set_tool_security(approval_security_config(false));
19188 release.notify_one();
19189 let record = call.await.unwrap();
19190
19191 assert!(!record.executed);
19192 assert!(!record.success);
19193 assert_eq!(record.runtime_config_version, expected_version);
19194 assert_eq!(record.policy.outcome, PermissionOutcome::Unavailable);
19195 assert_eq!(max_active.load(Ordering::SeqCst), 0);
19196 assert!(record.metadata.contains_key("policy_snapshot"));
19197 }
19198
19199 #[test]
19200 fn invalid_live_policy_does_not_replace_snapshot_or_generation() {
19201 let agent = AgentBuilder::new()
19202 .system_prompt("Test runtime policy validation.")
19203 .llm(Arc::new(mock_with_response("done")))
19204 .build()
19205 .unwrap();
19206 let control = agent.runtime_control();
19207 let mut valid = ToolSecurityConfig::default();
19208 valid.tools.insert(
19209 "web_search".to_string(),
19210 ai_agents_tools::ToolPolicyConfig {
19211 max_results: Some(5),
19212 ..Default::default()
19213 },
19214 );
19215 let generation = control.try_set_tool_security(valid).unwrap();
19216
19217 let mut invalid = ToolSecurityConfig::default();
19218 invalid.tools.insert(
19219 "web_search".to_string(),
19220 ai_agents_tools::ToolPolicyConfig {
19221 max_results: Some(0),
19222 ..Default::default()
19223 },
19224 );
19225 let error = control.try_set_tool_security(invalid).unwrap_err();
19226
19227 assert!(
19228 error
19229 .to_string()
19230 .contains("max_results must be greater than 0")
19231 );
19232 assert_eq!(control.version(), generation);
19233 assert_eq!(
19234 control
19235 .state
19236 .tool_security_override
19237 .read()
19238 .as_ref()
19239 .unwrap()
19240 .config()
19241 .tools["web_search"]
19242 .max_results,
19243 Some(5)
19244 );
19245 }
19246
19247 #[test]
19249 fn invalid_timeout_config_stops_before_approval_or_tool_invocation() {
19250 let tool_calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19251 let approval_calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19252 let spec = crate::spec::AgentSpec {
19253 tool_security: ToolSecurityConfig {
19254 enabled: true,
19255 default_timeout_ms: u64::MAX,
19256 ..Default::default()
19257 },
19258 ..Default::default()
19259 };
19260
19261 let result = AgentBuilder::from_spec(spec)
19262 .llm(Arc::new(mock_with_response("done")))
19263 .tool(Arc::new(FlakyWriteTool {
19264 calls: Arc::clone(&tool_calls),
19265 }))
19266 .hitl_engine(HITLEngine::new(ai_agents_hitl::HITLConfig::default()))
19267 .approval_handler(Arc::new(CountingApprovalHandler {
19268 calls: Arc::clone(&approval_calls),
19269 }))
19270 .build();
19271
19272 assert!(result.is_err());
19273 assert_eq!(approval_calls.load(Ordering::SeqCst), 0);
19274 assert_eq!(tool_calls.load(Ordering::SeqCst), 0);
19275 }
19276
19277 #[test]
19279 fn invalid_recovery_timeout_config_stops_before_approval_or_tool_invocation() {
19280 use ai_agents_recovery::{ErrorRecoveryConfig, ToolRecoveryConfig, ToolRetryConfig};
19281
19282 let tool_calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19283 let approval_calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19284 let spec = crate::spec::AgentSpec {
19285 error_recovery: ErrorRecoveryConfig {
19286 tools: ToolRecoveryConfig {
19287 default: ToolRetryConfig {
19288 timeout_ms: Some(u64::MAX),
19289 ..Default::default()
19290 },
19291 ..Default::default()
19292 },
19293 ..Default::default()
19294 },
19295 ..Default::default()
19296 };
19297
19298 let result = AgentBuilder::from_spec(spec)
19299 .llm(Arc::new(mock_with_response("done")))
19300 .tool(Arc::new(FlakyWriteTool {
19301 calls: Arc::clone(&tool_calls),
19302 }))
19303 .hitl_engine(HITLEngine::new(ai_agents_hitl::HITLConfig::default()))
19304 .approval_handler(Arc::new(CountingApprovalHandler {
19305 calls: Arc::clone(&approval_calls),
19306 }))
19307 .build();
19308
19309 assert!(result.is_err());
19310 assert_eq!(approval_calls.load(Ordering::SeqCst), 0);
19311 assert_eq!(tool_calls.load(Ordering::SeqCst), 0);
19312 }
19313
19314 #[test]
19316 fn invalid_timeout_policy_does_not_replace_snapshot_or_generation() {
19317 let agent = AgentBuilder::new()
19318 .system_prompt("Test runtime timeout policy validation.")
19319 .llm(Arc::new(mock_with_response("done")))
19320 .build()
19321 .unwrap();
19322 let control = agent.runtime_control();
19323 let valid = ToolSecurityConfig {
19324 default_timeout_ms: 5_000,
19325 ..Default::default()
19326 };
19327 let generation = control.try_set_tool_security(valid).unwrap();
19328
19329 let invalid = ToolSecurityConfig {
19330 default_timeout_ms: MAX_TOOL_TIMEOUT_MS + 1,
19331 ..Default::default()
19332 };
19333 let error = control.try_set_tool_security(invalid).unwrap_err();
19334
19335 assert!(error.to_string().contains(&format!(
19336 "tool_security.default_timeout_ms must be no greater than {MAX_TOOL_TIMEOUT_MS} milliseconds"
19337 )));
19338 assert_eq!(control.version(), generation);
19339 assert_eq!(
19340 control
19341 .state
19342 .tool_security_override
19343 .read()
19344 .as_ref()
19345 .unwrap()
19346 .config()
19347 .default_timeout_ms,
19348 5_000
19349 );
19350 }
19351
19352 #[test]
19354 fn runtime_tool_timeout_conversion_enforces_the_stable_boundary() {
19355 let timeout = RuntimeAgent::validated_tool_timeout(MAX_TOOL_TIMEOUT_MS).unwrap();
19356 assert_eq!(timeout.timer, Duration::from_millis(MAX_TOOL_TIMEOUT_MS));
19357 assert_eq!(
19358 timeout.deadline_delta,
19359 chrono::Duration::milliseconds(MAX_TOOL_TIMEOUT_MS as i64)
19360 );
19361
19362 for timeout_ms in [MAX_TOOL_TIMEOUT_MS + 1, u64::MAX] {
19363 let error = RuntimeAgent::validated_tool_timeout(timeout_ms).unwrap_err();
19364 assert!(error.to_string().contains(&format!(
19365 "effective tool timeout_ms must be no greater than {MAX_TOOL_TIMEOUT_MS} milliseconds"
19366 )));
19367 }
19368 }
19369
19370 #[tokio::test]
19371 async fn persistent_override_preserves_rate_history_within_generation() {
19372 let calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19373 let agent = AgentBuilder::new()
19374 .system_prompt("Test persistent policy overrides.")
19375 .llm(Arc::new(mock_with_response("done")))
19376 .tool(Arc::new(RecoveryTestTool {
19377 id: "limited_override".to_string(),
19378 succeeds: true,
19379 calls: Arc::clone(&calls),
19380 max_output_chars: None,
19381 }))
19382 .build()
19383 .unwrap();
19384 let mut security = ToolSecurityConfig {
19385 enabled: true,
19386 fail_closed: true,
19387 ..Default::default()
19388 };
19389 let policy = ai_agents_tools::ToolPolicyConfig {
19390 write_paths: vec![".".to_string()],
19391 rate_limit: Some(1),
19392 ..Default::default()
19393 };
19394 security
19395 .tools
19396 .insert("limited_override".to_string(), policy);
19397 let generation = agent.runtime_control().set_tool_security(security);
19398
19399 let first = agent
19400 .invoke_tool(ToolExecutionRequest::new(
19401 "limited-first",
19402 "limited_override",
19403 serde_json::json!({"path": "./limited.txt"}),
19404 ToolCallSource::Manual,
19405 ))
19406 .await
19407 .unwrap();
19408 let second = agent
19409 .invoke_tool(ToolExecutionRequest::new(
19410 "limited-second",
19411 "limited_override",
19412 serde_json::json!({"path": "./limited.txt"}),
19413 ToolCallSource::Manual,
19414 ))
19415 .await
19416 .unwrap();
19417
19418 assert!(first.success);
19419 assert_eq!(first.policy_version, generation);
19420 assert!(!second.executed);
19421 assert!(second.output.contains("Rate limit exceeded"));
19422 assert_eq!(second.policy_version, generation);
19423 assert_eq!(calls.load(Ordering::SeqCst), 1);
19424 }
19425
19426 #[tokio::test]
19427 async fn concurrent_rate_admission_consumes_capacity_atomically() {
19428 let calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19429 let tool = Arc::new(RecoveryTestTool {
19430 id: "atomic_rate".to_string(),
19431 succeeds: true,
19432 calls: Arc::clone(&calls),
19433 max_output_chars: None,
19434 });
19435 let arguments = serde_json::json!({"path": "./atomic-rate.txt"});
19436 let bindings = tool.policy_bindings();
19437 let classification = tool.classify_call(&arguments);
19438 let resource_keys =
19439 tool_resource_lock_keys(tool.id(), &arguments, &bindings, &classification);
19440 let mut security = ToolSecurityConfig {
19441 enabled: true,
19442 fail_closed: true,
19443 ..Default::default()
19444 };
19445 let policy = ai_agents_tools::ToolPolicyConfig {
19446 write_paths: vec![".".to_string()],
19447 rate_limit: Some(1),
19448 ..Default::default()
19449 };
19450 security.tools.insert(tool.id().to_string(), policy);
19451 let agent = Arc::new(
19452 AgentBuilder::new()
19453 .system_prompt("Test atomic rate admission.")
19454 .llm(Arc::new(mock_with_response("done")))
19455 .tool(tool)
19456 .tool_security(ToolSecurityEngine::new(security))
19457 .build()
19458 .unwrap(),
19459 );
19460 let held = agent
19461 .acquire_tool_resource_locks(&resource_keys)
19462 .await
19463 .unwrap();
19464 let left = {
19465 let agent = Arc::clone(&agent);
19466 let arguments = arguments.clone();
19467 tokio::spawn(async move {
19468 agent
19469 .invoke_tool(ToolExecutionRequest::new(
19470 "atomic-rate-left",
19471 "atomic_rate",
19472 arguments,
19473 ToolCallSource::Manual,
19474 ))
19475 .await
19476 .unwrap()
19477 })
19478 };
19479 let right = {
19480 let agent = Arc::clone(&agent);
19481 tokio::spawn(async move {
19482 agent
19483 .invoke_tool(ToolExecutionRequest::new(
19484 "atomic-rate-right",
19485 "atomic_rate",
19486 arguments,
19487 ToolCallSource::Manual,
19488 ))
19489 .await
19490 .unwrap()
19491 })
19492 };
19493 tokio::time::sleep(std::time::Duration::from_millis(25)).await;
19494 drop(held);
19495 let (left, right) = tokio::join!(left, right);
19496 let records = [left.unwrap(), right.unwrap()];
19497
19498 assert_eq!(records.iter().filter(|record| record.success).count(), 1);
19499 assert_eq!(records.iter().filter(|record| record.executed).count(), 1);
19500 assert!(
19501 records.iter().any(|record| {
19502 !record.executed && record.output.contains("Rate limit exceeded")
19503 })
19504 );
19505 assert_eq!(calls.load(Ordering::SeqCst), 1);
19506 }
19507
19508 #[tokio::test]
19509 async fn changed_policy_generation_invalidates_pending_approval() {
19510 let active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19511 let max_active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19512 let entered = Arc::new(tokio::sync::Barrier::new(2));
19513 let release = Arc::new(tokio::sync::Notify::new());
19514 let handler = Arc::new(BlockingApprovalHandler {
19515 entered: Arc::clone(&entered),
19516 release: Arc::clone(&release),
19517 result: ApprovalResult::Approved,
19518 });
19519 let agent = Arc::new(
19520 AgentBuilder::new()
19521 .system_prompt("Test stale approval denial.")
19522 .llm(Arc::new(mock_with_response("done")))
19523 .tool(Arc::new(LockedWriteTool {
19524 active: Arc::clone(&active),
19525 max_active: Arc::clone(&max_active),
19526 }))
19527 .tool_security(ToolSecurityEngine::new(approval_security_config(true)))
19528 .hitl_engine(HITLEngine::new(ai_agents_hitl::HITLConfig::default()))
19529 .approval_handler(handler)
19530 .build()
19531 .unwrap(),
19532 );
19533 let running = Arc::clone(&agent);
19534 let call = tokio::spawn(async move {
19535 running
19536 .invoke_tool(ToolExecutionRequest::new(
19537 "stale-approval",
19538 "locked_write",
19539 serde_json::json!({"path": "./stale.txt"}),
19540 ToolCallSource::Manual,
19541 ))
19542 .await
19543 .unwrap()
19544 });
19545 entered.wait().await;
19546 let generation = agent
19547 .runtime_control()
19548 .set_tool_security(approval_security_config(true));
19549 release.notify_one();
19550 let record = call.await.unwrap();
19551
19552 assert!(!record.executed);
19553 assert!(record.output.contains("Approval became stale"));
19554 assert_eq!(record.policy_version, generation);
19555 assert_eq!(max_active.load(Ordering::SeqCst), 0);
19556 }
19557
19558 #[tokio::test]
19559 async fn final_policy_reapplies_argument_caps_after_approval_changes() {
19560 use ai_agents_hitl::CallbackHandler;
19561
19562 let mut security = ToolSecurityConfig {
19563 enabled: true,
19564 fail_closed: true,
19565 ..Default::default()
19566 };
19567 let policy = ai_agents_tools::ToolPolicyConfig {
19568 read_paths: vec![".".to_string()],
19569 max_results: Some(5),
19570 require_confirmation: true,
19571 ..Default::default()
19572 };
19573 security.tools.insert("context_echo".to_string(), policy);
19574 let handler = CallbackHandler::new(|_| ApprovalResult::Modified {
19575 changes: HashMap::from([("max_results".to_string(), serde_json::json!(99))]),
19576 });
19577 let agent = AgentBuilder::new()
19578 .system_prompt("Test final argument caps.")
19579 .llm(Arc::new(mock_with_response("done")))
19580 .tool(Arc::new(ContextEchoTool))
19581 .tool_security(ToolSecurityEngine::new(security))
19582 .hitl_engine(HITLEngine::new(ai_agents_hitl::HITLConfig::default()))
19583 .approval_handler(Arc::new(handler))
19584 .build()
19585 .unwrap();
19586
19587 let record = agent
19588 .invoke_tool(ToolExecutionRequest::new(
19589 "final-cap",
19590 "context_echo",
19591 serde_json::json!({"path": ".", "max_results": 1}),
19592 ToolCallSource::Manual,
19593 ))
19594 .await
19595 .unwrap();
19596
19597 assert!(record.success);
19598 assert_eq!(record.executed_arguments["max_results"], 5);
19599 assert_eq!(
19600 record.approval.unwrap().modified_arguments.unwrap()["max_results"],
19601 5
19602 );
19603 }
19604
19605 #[tokio::test]
19606 async fn no_binding_writes_use_canonical_fallback_lock() {
19607 let active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19608 let max_active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19609 let agent = Arc::new(
19610 AgentBuilder::new()
19611 .system_prompt("Test fallback resource locks.")
19612 .llm(Arc::new(mock_with_response("done")))
19613 .tool(Arc::new(NoBindingWriteTool {
19614 active: Arc::clone(&active),
19615 max_active: Arc::clone(&max_active),
19616 }))
19617 .build()
19618 .unwrap(),
19619 );
19620 let left = {
19621 let agent = Arc::clone(&agent);
19622 tokio::spawn(async move {
19623 agent
19624 .invoke_tool(ToolExecutionRequest::new(
19625 "no-binding-left",
19626 "no_binding_write",
19627 serde_json::json!({}),
19628 ToolCallSource::Manual,
19629 ))
19630 .await
19631 .unwrap()
19632 })
19633 };
19634 let right = {
19635 let agent = Arc::clone(&agent);
19636 tokio::spawn(async move {
19637 agent
19638 .invoke_tool(ToolExecutionRequest::new(
19639 "no-binding-right",
19640 "no_binding_write",
19641 serde_json::json!({}),
19642 ToolCallSource::Manual,
19643 ))
19644 .await
19645 .unwrap()
19646 })
19647 };
19648 let (left, right) = tokio::join!(left, right);
19649
19650 assert!(left.unwrap().success);
19651 assert!(right.unwrap().success);
19652 assert_eq!(max_active.load(Ordering::SeqCst), 1);
19653 }
19654
19655 #[tokio::test]
19656 async fn parent_and_child_paths_share_a_resource_lock() {
19657 let active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19658 let max_active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19659 let agent = Arc::new(
19660 AgentBuilder::new()
19661 .system_prompt("Test parent child resource locks.")
19662 .llm(Arc::new(mock_with_response("done")))
19663 .tool(Arc::new(LockedWriteTool {
19664 active: Arc::clone(&active),
19665 max_active: Arc::clone(&max_active),
19666 }))
19667 .build()
19668 .unwrap(),
19669 );
19670 let parent = format!("./lock-parent-{}", uuid::Uuid::new_v4());
19671 let child = format!("{}/child.txt", parent);
19672 let left = {
19673 let agent = Arc::clone(&agent);
19674 tokio::spawn(async move {
19675 agent
19676 .invoke_tool(ToolExecutionRequest::new(
19677 "parent-lock",
19678 "locked_write",
19679 serde_json::json!({"path": parent}),
19680 ToolCallSource::Manual,
19681 ))
19682 .await
19683 .unwrap()
19684 })
19685 };
19686 let right = {
19687 let agent = Arc::clone(&agent);
19688 tokio::spawn(async move {
19689 agent
19690 .invoke_tool(ToolExecutionRequest::new(
19691 "child-lock",
19692 "locked_write",
19693 serde_json::json!({"path": child}),
19694 ToolCallSource::Manual,
19695 ))
19696 .await
19697 .unwrap()
19698 })
19699 };
19700 let (left, right) = tokio::join!(left, right);
19701
19702 assert!(left.unwrap().success);
19703 assert!(right.unwrap().success);
19704 assert_eq!(max_active.load(Ordering::SeqCst), 1);
19705 }
19706
19707 #[tokio::test]
19708 async fn tool_hooks_can_reenter_after_resource_guards_are_dropped() {
19709 let calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19710 let hooks = Arc::new(ReentrantToolHooks {
19711 agent: parking_lot::Mutex::new(None),
19712 invoked: AtomicBool::new(false),
19713 nested_success: AtomicBool::new(false),
19714 });
19715 let agent = Arc::new(
19716 AgentBuilder::new()
19717 .system_prompt("Test hook reentrancy.")
19718 .llm(Arc::new(mock_with_response("done")))
19719 .tool(Arc::new(RecoveryTestTool {
19720 id: "reentrant_write".to_string(),
19721 succeeds: true,
19722 calls: Arc::clone(&calls),
19723 max_output_chars: None,
19724 }))
19725 .hooks(hooks.clone())
19726 .build()
19727 .unwrap(),
19728 );
19729 *hooks.agent.lock() = Some(Arc::downgrade(&agent));
19730 let record = tokio::time::timeout(
19731 std::time::Duration::from_secs(2),
19732 agent.invoke_tool(ToolExecutionRequest::new(
19733 "outer-hook-call",
19734 "reentrant_write",
19735 serde_json::json!({"path": "./hook.txt"}),
19736 ToolCallSource::Manual,
19737 )),
19738 )
19739 .await
19740 .expect("tool completion hook must not retain resource guards")
19741 .unwrap();
19742
19743 assert!(record.success);
19744 assert!(hooks.nested_success.load(Ordering::SeqCst));
19745 assert_eq!(calls.load(Ordering::SeqCst), 2);
19746 }
19747
19748 #[tokio::test]
19750 async fn fallback_finalizes_original_record_before_shared_execution() {
19751 let primary_calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19752 let fallback_calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19753 let hooks = Arc::new(ToolLifecycleRecordingHooks::new());
19754 let agent = AgentBuilder::new()
19755 .system_prompt("Test fallback execution.")
19756 .llm(Arc::new(mock_with_response("done")))
19757 .tool(Arc::new(RecoveryTestTool {
19758 id: "primary".to_string(),
19759 succeeds: false,
19760 calls: Arc::clone(&primary_calls),
19761 max_output_chars: None,
19762 }))
19763 .tool(Arc::new(RecoveryTestTool {
19764 id: "fallback".to_string(),
19765 succeeds: true,
19766 calls: Arc::clone(&fallback_calls),
19767 max_output_chars: None,
19768 }))
19769 .recovery_manager(recovery_manager_with_fallbacks([(
19770 "primary".to_string(),
19771 "fallback".to_string(),
19772 )]))
19773 .hooks(hooks.clone())
19774 .build()
19775 .unwrap();
19776 let record = tokio::time::timeout(
19777 std::time::Duration::from_secs(2),
19778 agent.invoke_tool(ToolExecutionRequest::new(
19779 "fallback-call",
19780 "primary",
19781 serde_json::json!({"path": "./shared.txt"}),
19782 ToolCallSource::Manual,
19783 )),
19784 )
19785 .await
19786 .expect("fallback must not retain the primary resource guard")
19787 .unwrap();
19788
19789 assert_eq!(
19790 hooks.events(),
19791 vec![
19792 "start:primary",
19793 "complete:primary:false",
19794 "record:primary:true",
19795 "error",
19796 "start:fallback",
19797 "complete:fallback:true",
19798 "record:fallback:true",
19799 ]
19800 );
19801 let records = hooks.records();
19802 assert_eq!(records.len(), 2);
19803 let original = &records[0];
19804 assert_eq!(original.canonical_id, "primary");
19805 assert!(matches!(original.source, ToolCallSource::Manual));
19806 assert!(original.executed);
19807 assert!(!original.success);
19808
19809 let fallback = &records[1];
19810 assert_eq!(fallback.canonical_id, "fallback");
19811 assert_eq!(fallback.call_id, "fallback-call");
19812 assert!(matches!(
19813 &fallback.source,
19814 ToolCallSource::Fallback { original_tool } if original_tool == "primary"
19815 ));
19816 assert!(fallback.executed);
19817 assert!(fallback.success);
19818 assert_eq!(record.canonical_id, fallback.canonical_id);
19819 assert_eq!(record.output, fallback.output);
19820
19821 let history = agent.tool_call_history();
19822 assert_eq!(
19823 history
19824 .iter()
19825 .map(|entry| entry.tool_id.as_str())
19826 .collect::<Vec<_>>(),
19827 vec!["primary", "fallback"]
19828 );
19829 assert_eq!(history[0].result.get("success"), Some(&Value::Bool(false)));
19830 assert_eq!(primary_calls.load(Ordering::SeqCst), 1);
19831 assert_eq!(fallback_calls.load(Ordering::SeqCst), 1);
19832 }
19833
19834 #[tokio::test]
19836 async fn self_fallback_cycle_is_denied_before_reinvocation() {
19837 let calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19838 let hooks = Arc::new(ToolLifecycleRecordingHooks::new());
19839 let agent = AgentBuilder::new()
19840 .system_prompt("Test self-fallback cycle admission.")
19841 .llm(Arc::new(mock_with_response("done")))
19842 .tool(Arc::new(RecoveryTestTool {
19843 id: "primary".to_string(),
19844 succeeds: false,
19845 calls: Arc::clone(&calls),
19846 max_output_chars: None,
19847 }))
19848 .recovery_manager(recovery_manager_with_fallbacks([(
19849 "primary".to_string(),
19850 "primary".to_string(),
19851 )]))
19852 .hooks(hooks.clone())
19853 .build()
19854 .unwrap();
19855
19856 let record = tokio::time::timeout(
19857 std::time::Duration::from_secs(2),
19858 agent.invoke_tool(ToolExecutionRequest::new(
19859 "self-fallback-call",
19860 "primary",
19861 serde_json::json!({"path": "./shared.txt"}),
19862 ToolCallSource::Manual,
19863 )),
19864 )
19865 .await
19866 .expect("self fallback must terminate without recursive execution")
19867 .unwrap();
19868
19869 assert_eq!(calls.load(Ordering::SeqCst), 1);
19870 assert_eq!(record.canonical_id, "primary");
19871 assert!(!record.executed);
19872 assert!(!record.success);
19873 assert_eq!(record.policy.outcome, PermissionOutcome::Deny);
19874 assert!(record.output.contains("fallback cycle"));
19875 assert!(matches!(
19876 record.source,
19877 ToolCallSource::Fallback { ref original_tool } if original_tool == "primary"
19878 ));
19879 assert_eq!(
19880 record.metadata.get("fallback_chain"),
19881 Some(&serde_json::json!(["primary"]))
19882 );
19883 assert_eq!(
19884 hooks.events(),
19885 vec![
19886 "start:primary",
19887 "complete:primary:false",
19888 "record:primary:true",
19889 "error",
19890 "complete:primary:false",
19891 "record:primary:false",
19892 "error",
19893 ]
19894 );
19895 assert_eq!(agent.tool_call_history().len(), 2);
19896 }
19897
19898 #[tokio::test]
19900 async fn alias_mediated_fallback_cycle_is_denied_canonically() {
19901 let primary_calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19902 let secondary_calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19903 let hooks = Arc::new(ToolLifecycleRecordingHooks::new());
19904 let agent = AgentBuilder::new()
19905 .system_prompt("Test canonical fallback cycle admission.")
19906 .llm(Arc::new(mock_with_response("done")))
19907 .tool(Arc::new(RecoveryTestTool {
19908 id: "primary".to_string(),
19909 succeeds: false,
19910 calls: Arc::clone(&primary_calls),
19911 max_output_chars: None,
19912 }))
19913 .tool(Arc::new(RecoveryTestTool {
19914 id: "secondary".to_string(),
19915 succeeds: false,
19916 calls: Arc::clone(&secondary_calls),
19917 max_output_chars: None,
19918 }))
19919 .recovery_manager(recovery_manager_with_fallbacks([
19920 ("primary".to_string(), "secondary".to_string()),
19921 ("secondary".to_string(), "primary alias".to_string()),
19922 ]))
19923 .hooks(hooks.clone())
19924 .build()
19925 .unwrap();
19926 agent.tools.set_tool_aliases(
19927 "primary",
19928 ToolAliases::new().with_name("en", "primary alias"),
19929 );
19930
19931 let record = tokio::time::timeout(
19932 std::time::Duration::from_secs(2),
19933 agent.invoke_tool(ToolExecutionRequest::new(
19934 "alias-fallback-call",
19935 "primary",
19936 serde_json::json!({"path": "./shared.txt"}),
19937 ToolCallSource::Manual,
19938 )),
19939 )
19940 .await
19941 .expect("alias-mediated fallback cycle must terminate")
19942 .unwrap();
19943
19944 assert_eq!(primary_calls.load(Ordering::SeqCst), 1);
19945 assert_eq!(secondary_calls.load(Ordering::SeqCst), 1);
19946 assert_eq!(record.requested_name, "primary alias");
19947 assert_eq!(record.canonical_id, "primary");
19948 assert!(!record.executed);
19949 assert!(record.output.contains("fallback cycle"));
19950 assert_eq!(
19951 record.metadata.get("fallback_chain"),
19952 Some(&serde_json::json!(["primary", "secondary"]))
19953 );
19954 assert_eq!(hooks.records().len(), 3);
19955 assert_eq!(agent.tool_call_history().len(), 3);
19956 }
19957
19958 #[tokio::test]
19960 async fn final_canonical_drift_cannot_bypass_fallback_ancestry() {
19961 let primary_calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19962 let secondary_calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
19963 let provider = Arc::new(DriftingFallbackProvider {
19964 refreshed: AtomicBool::new(false),
19965 primary_calls: Arc::clone(&primary_calls),
19966 secondary_calls: Arc::clone(&secondary_calls),
19967 });
19968 let registry = ToolRegistry::new();
19969 registry.register_provider(provider).await.unwrap();
19970 let lifecycle = Arc::new(ToolLifecycleRecordingHooks::new());
19971 let hooks = Arc::new(RefreshFallbackProviderHooks {
19972 agent: parking_lot::Mutex::new(None),
19973 lifecycle: Arc::clone(&lifecycle),
19974 });
19975 let agent = Arc::new(
19976 AgentBuilder::new()
19977 .system_prompt("Test final canonical fallback admission.")
19978 .llm(Arc::new(mock_with_response("done")))
19979 .tools(registry)
19980 .recovery_manager(recovery_manager_with_fallbacks([(
19981 "primary".to_string(),
19982 "fallback alias".to_string(),
19983 )]))
19984 .hooks(hooks.clone())
19985 .build()
19986 .unwrap(),
19987 );
19988 *hooks.agent.lock() = Some(Arc::downgrade(&agent));
19989
19990 let record = agent
19991 .invoke_tool(ToolExecutionRequest::new(
19992 "drifting-fallback-call",
19993 "primary",
19994 serde_json::json!({"path": "./shared.txt"}),
19995 ToolCallSource::Manual,
19996 ))
19997 .await
19998 .unwrap();
19999
20000 assert_eq!(primary_calls.load(Ordering::SeqCst), 1);
20001 assert_eq!(secondary_calls.load(Ordering::SeqCst), 0);
20002 assert_eq!(record.canonical_id, "secondary");
20003 assert!(!record.executed);
20004 assert_eq!(record.policy.outcome, PermissionOutcome::Deny);
20005 assert!(record.output.contains("fallback cycle"));
20006 assert_eq!(
20007 record.metadata.get("fallback_chain"),
20008 Some(&serde_json::json!(["primary", "secondary"]))
20009 );
20010 assert_eq!(
20011 record.metadata.get("final_resolved_canonical_id"),
20012 Some(&serde_json::json!("primary"))
20013 );
20014 assert_eq!(
20015 lifecycle.events(),
20016 vec![
20017 "start:primary",
20018 "complete:primary:false",
20019 "record:primary:true",
20020 "error",
20021 "start:secondary",
20022 "complete:secondary:false",
20023 "record:secondary:false",
20024 "error",
20025 ]
20026 );
20027 let records = lifecycle.records();
20028 assert_eq!(records.len(), 2);
20029 assert_eq!(records[1].canonical_id, "secondary");
20030 assert_eq!(
20031 records[1].metadata.get("final_resolved_canonical_id"),
20032 Some(&serde_json::json!("primary"))
20033 );
20034 let history = agent.tool_call_history();
20035 assert_eq!(
20036 history
20037 .iter()
20038 .map(|entry| entry.tool_id.as_str())
20039 .collect::<Vec<_>>(),
20040 vec!["primary", "secondary"]
20041 );
20042 }
20043
20044 #[tokio::test]
20046 async fn acyclic_fallback_chain_is_denied_after_the_hop_limit() {
20047 let tool_count = MAX_TOOL_FALLBACK_HOPS + 2;
20048 let calls = (0..tool_count)
20049 .map(|_| Arc::new(std::sync::atomic::AtomicUsize::new(0)))
20050 .collect::<Vec<_>>();
20051 let mut builder = AgentBuilder::new()
20052 .system_prompt("Test bounded acyclic fallback admission.")
20053 .llm(Arc::new(mock_with_response("done")));
20054 for (index, counter) in calls.iter().enumerate() {
20055 builder = builder.tool(Arc::new(RecoveryTestTool {
20056 id: format!("fallback_{index}"),
20057 succeeds: false,
20058 calls: Arc::clone(counter),
20059 max_output_chars: None,
20060 }));
20061 }
20062 let fallbacks = (0..tool_count - 1).map(|index| {
20063 (
20064 format!("fallback_{index}"),
20065 format!("fallback_{}", index + 1),
20066 )
20067 });
20068 let agent = builder
20069 .recovery_manager(recovery_manager_with_fallbacks(fallbacks))
20070 .build()
20071 .unwrap();
20072
20073 let record = tokio::time::timeout(
20074 std::time::Duration::from_secs(2),
20075 agent.invoke_tool(ToolExecutionRequest::new(
20076 "bounded-fallback-call",
20077 "fallback_0",
20078 serde_json::json!({"path": "./shared.txt"}),
20079 ToolCallSource::Manual,
20080 )),
20081 )
20082 .await
20083 .expect("bounded fallback chain must terminate")
20084 .unwrap();
20085
20086 for counter in calls.iter().take(MAX_TOOL_FALLBACK_HOPS + 1) {
20087 assert_eq!(counter.load(Ordering::SeqCst), 1);
20088 }
20089 assert_eq!(calls[MAX_TOOL_FALLBACK_HOPS + 1].load(Ordering::SeqCst), 0);
20090 assert_eq!(
20091 record.canonical_id,
20092 format!("fallback_{}", MAX_TOOL_FALLBACK_HOPS + 1)
20093 );
20094 assert!(!record.executed);
20095 assert!(record.output.contains("maximum of 16 hops"));
20096 assert_eq!(agent.tool_call_history().len(), tool_count);
20097 }
20098
20099 #[tokio::test]
20100 async fn diagnostics_without_provider_records_unavailable_without_execution() {
20101 let mock = mock_with_response("hello");
20102 let yaml = r#"
20103name: DiagnosticsNoProviderAgent
20104system_prompt: "Review diagnostics."
20105tools: [diagnostics]
20106"#;
20107 let agent = AgentBuilder::from_yaml(yaml)
20108 .unwrap()
20109 .llm(Arc::new(mock))
20110 .auto_configure_features()
20111 .unwrap()
20112 .build()
20113 .unwrap();
20114
20115 let record = agent
20116 .invoke_tool(ToolExecutionRequest::new(
20117 "diagnostics-call",
20118 "diagnostics",
20119 serde_json::json!({}),
20120 ToolCallSource::Manual,
20121 ))
20122 .await
20123 .unwrap();
20124
20125 assert!(!record.executed);
20126 assert!(!record.success);
20127 assert_eq!(record.policy.outcome, PermissionOutcome::Unavailable);
20128 }
20129
20130 #[tokio::test]
20131 async fn web_search_without_provider_records_unavailable_without_execution() {
20132 let mock = mock_with_response("hello");
20133 let yaml = r#"
20134name: WebSearchNoProviderAgent
20135system_prompt: "You search the web."
20136tools: [web_search]
20137"#;
20138 let agent = AgentBuilder::from_yaml(yaml)
20139 .unwrap()
20140 .llm(Arc::new(mock))
20141 .auto_configure_features()
20142 .unwrap()
20143 .build()
20144 .unwrap();
20145
20146 let record = agent
20147 .invoke_tool(ToolExecutionRequest::new(
20148 "web-search-call",
20149 "web_search",
20150 serde_json::json!({"query": "rust async"}),
20151 ToolCallSource::Manual,
20152 ))
20153 .await
20154 .unwrap();
20155
20156 assert!(!record.executed);
20157 assert!(!record.success);
20158 assert_eq!(record.policy.outcome, PermissionOutcome::Unavailable);
20159 }
20160
20161 #[tokio::test]
20162 async fn unavailable_host_tool_does_not_request_approval() {
20163 let approvals = Arc::new(std::sync::atomic::AtomicUsize::new(0));
20164 let handler = Arc::new(CountingApprovalHandler {
20165 calls: Arc::clone(&approvals),
20166 });
20167 let mut security = ToolSecurityConfig {
20168 enabled: true,
20169 fail_closed: true,
20170 ..Default::default()
20171 };
20172 security.tools.insert(
20173 "web_search".to_string(),
20174 ai_agents_tools::ToolPolicyConfig {
20175 enabled: true,
20176 require_confirmation: true,
20177 ..Default::default()
20178 },
20179 );
20180 let yaml = r#"
20181name: UnavailableApprovalAgent
20182system_prompt: "Search only with approval."
20183tools: [web_search]
20184"#;
20185 let agent = AgentBuilder::from_yaml(yaml)
20186 .unwrap()
20187 .llm(Arc::new(mock_with_response("done")))
20188 .auto_configure_features()
20189 .unwrap()
20190 .tool_security(ToolSecurityEngine::new(security))
20191 .hitl_engine(HITLEngine::new(ai_agents_hitl::HITLConfig::default()))
20192 .approval_handler(handler)
20193 .build()
20194 .unwrap();
20195
20196 let record = agent
20197 .invoke_tool(ToolExecutionRequest::new(
20198 "unavailable-before-approval",
20199 "web_search",
20200 serde_json::json!({"query": "rust async"}),
20201 ToolCallSource::Manual,
20202 ))
20203 .await
20204 .unwrap();
20205
20206 assert_eq!(approvals.load(Ordering::SeqCst), 0);
20207 assert!(!record.executed);
20208 assert!(!record.success);
20209 assert_eq!(record.policy.outcome, PermissionOutcome::Unavailable);
20210 assert!(
20211 record
20212 .approval
20213 .as_ref()
20214 .is_some_and(|approval| matches!(approval.status, ToolApprovalStatus::Unavailable))
20215 );
20216 }
20217
20218 #[tokio::test]
20219 async fn test_spawner_section_does_not_grant_core_tools_when_top_level_tools_omitted() {
20220 let mock = mock_with_response("hello");
20221 let yaml = r#"
20222name: SpawnerNoGrantAgent
20223system_prompt: "You manage agents."
20224spawner:
20225 max_agents: 2
20226"#;
20227 let agent = AgentBuilder::from_yaml(yaml)
20228 .unwrap()
20229 .llm(Arc::new(mock))
20230 .auto_configure_features()
20231 .unwrap()
20232 .auto_configure_spawner()
20233 .await
20234 .unwrap()
20235 .build()
20236 .unwrap();
20237
20238 let available = agent.get_available_tool_ids().await.unwrap();
20239 assert!(available.is_empty());
20240 }
20241
20242 #[tokio::test]
20243 async fn test_spawner_section_does_not_grant_core_tools_when_top_level_tools_empty() {
20244 let mock = mock_with_response("hello");
20245 let yaml = r#"
20246name: EmptySpawnerNoGrantAgent
20247system_prompt: "You manage agents."
20248tools: []
20249spawner:
20250 max_agents: 2
20251"#;
20252 let agent = AgentBuilder::from_yaml(yaml)
20253 .unwrap()
20254 .llm(Arc::new(mock))
20255 .auto_configure_features()
20256 .unwrap()
20257 .auto_configure_spawner()
20258 .await
20259 .unwrap()
20260 .build()
20261 .unwrap();
20262
20263 let available = agent.get_available_tool_ids().await.unwrap();
20264 assert!(available.is_empty());
20265 }
20266
20267 #[tokio::test]
20268 async fn test_management_tools_flag_grants_core_tools_when_top_level_tools_empty() {
20269 let mock = mock_with_response("hello");
20270 let yaml = r#"
20271name: ManagementGrantAgent
20272system_prompt: "You manage agents."
20273tools: []
20274spawner:
20275 management_tools: true
20276"#;
20277 let agent = AgentBuilder::from_yaml(yaml)
20278 .unwrap()
20279 .llm(Arc::new(mock))
20280 .auto_configure_features()
20281 .unwrap()
20282 .auto_configure_spawner()
20283 .await
20284 .unwrap()
20285 .build()
20286 .unwrap();
20287
20288 let available = agent.get_available_tool_ids().await.unwrap();
20289 assert_eq!(available.len(), 4);
20290 assert!(available.contains(&"spawn_agent".to_string()));
20291 assert!(available.contains(&"send_agent_message".to_string()));
20292 assert!(available.contains(&"list_agents".to_string()));
20293 assert!(available.contains(&"remove_agent".to_string()));
20294 }
20295
20296 #[tokio::test]
20297 async fn test_management_tools_flag_grants_core_tools_when_top_level_tools_omitted() {
20298 let mock = mock_with_response("hello");
20299 let yaml = r#"
20300name: ManagementOmittedToolsGrantAgent
20301system_prompt: "You manage agents."
20302spawner:
20303 management_tools: true
20304"#;
20305 let agent = AgentBuilder::from_yaml(yaml)
20306 .unwrap()
20307 .llm(Arc::new(mock))
20308 .auto_configure_features()
20309 .unwrap()
20310 .auto_configure_spawner()
20311 .await
20312 .unwrap()
20313 .build()
20314 .unwrap();
20315
20316 let available = agent.get_available_tool_ids().await.unwrap();
20317 assert_eq!(available.len(), 4);
20318 assert!(available.contains(&"spawn_agent".to_string()));
20319 assert!(available.contains(&"send_agent_message".to_string()));
20320 assert!(available.contains(&"list_agents".to_string()));
20321 assert!(available.contains(&"remove_agent".to_string()));
20322 }
20323
20324 #[tokio::test]
20325 async fn test_management_tools_selected_grants_only_selected_tools() {
20326 let mock = mock_with_response("hello");
20327 let yaml = r#"
20328name: ManagementSelectedGrantAgent
20329system_prompt: "You manage agents."
20330tools: []
20331spawner:
20332 management_tools:
20333 - spawn_agent
20334 - send_agent_message
20335 - list_agents
20336"#;
20337 let agent = AgentBuilder::from_yaml(yaml)
20338 .unwrap()
20339 .llm(Arc::new(mock))
20340 .auto_configure_features()
20341 .unwrap()
20342 .auto_configure_spawner()
20343 .await
20344 .unwrap()
20345 .build()
20346 .unwrap();
20347
20348 let available = agent.get_available_tool_ids().await.unwrap();
20349 assert_eq!(available.len(), 3);
20350 assert!(available.contains(&"spawn_agent".to_string()));
20351 assert!(available.contains(&"send_agent_message".to_string()));
20352 assert!(available.contains(&"list_agents".to_string()));
20353 assert!(!available.contains(&"remove_agent".to_string()));
20354 }
20355
20356 #[tokio::test]
20357 async fn test_orchestration_tools_flag_grants_tools_when_top_level_tools_empty() {
20358 let mock = mock_with_response("hello");
20359 let yaml = r#"
20360name: OrchestrationGrantAgent
20361system_prompt: "You coordinate agents."
20362llms:
20363 default:
20364 provider: openai
20365 model: gpt-4
20366 router:
20367 provider: openai
20368 model: gpt-4
20369llm:
20370 default: default
20371 router: router
20372tools: []
20373spawner:
20374 orchestration_tools: true
20375"#;
20376 let agent = AgentBuilder::from_yaml(yaml)
20377 .unwrap()
20378 .llm(Arc::new(mock))
20379 .auto_configure_features()
20380 .unwrap()
20381 .auto_configure_spawner()
20382 .await
20383 .unwrap()
20384 .build()
20385 .unwrap();
20386
20387 let available = agent.get_available_tool_ids().await.unwrap();
20388 assert_eq!(available.len(), 5);
20389 assert!(available.contains(&"route_to_agent".to_string()));
20390 assert!(available.contains(&"pipeline_process".to_string()));
20391 assert!(available.contains(&"concurrent_ask".to_string()));
20392 assert!(available.contains(&"group_discussion".to_string()));
20393 assert!(available.contains(&"handoff_conversation".to_string()));
20394 }
20395
20396 #[tokio::test]
20397 async fn test_persona_evolve_flag_grants_tool_when_top_level_tools_empty() {
20398 let mock = mock_with_response("hello");
20399 let yaml = r#"
20400name: PersonaGrantAgent
20401system_prompt: "You can evolve persona."
20402llm:
20403 provider: openai
20404 model: gpt-4
20405tools: []
20406persona:
20407 identity:
20408 name: "Guide"
20409 role: "Helper"
20410 evolution:
20411 enabled: true
20412 allow_llm_evolve: true
20413 mutable_fields:
20414 - traits.personality
20415"#;
20416 let agent = AgentBuilder::from_yaml(yaml)
20417 .unwrap()
20418 .llm(Arc::new(mock))
20419 .build()
20420 .unwrap();
20421
20422 let available = agent.get_available_tool_ids().await.unwrap();
20423 assert_eq!(available, vec!["persona_evolve".to_string()]);
20424 }
20425
20426 #[tokio::test]
20427 async fn test_persona_evolve_flag_grants_tool_when_top_level_tools_omitted() {
20428 let mock = mock_with_response("hello");
20429 let yaml = r#"
20430name: PersonaOmittedToolsGrantAgent
20431system_prompt: "You can evolve persona."
20432llm:
20433 provider: openai
20434 model: gpt-4
20435persona:
20436 identity:
20437 name: "Guide"
20438 role: "Helper"
20439 evolution:
20440 enabled: true
20441 allow_llm_evolve: true
20442 mutable_fields:
20443 - traits.personality
20444"#;
20445 let agent = AgentBuilder::from_yaml(yaml)
20446 .unwrap()
20447 .llm(Arc::new(mock))
20448 .build()
20449 .unwrap();
20450
20451 let available = agent.get_available_tool_ids().await.unwrap();
20452 assert_eq!(available, vec!["persona_evolve".to_string()]);
20453 }
20454
20455 #[tokio::test]
20456 async fn test_omitted_yaml_tools_exposes_no_tools() {
20457 let mock = mock_with_response("hello");
20458 let yaml = r#"
20459name: NoToolsAgent
20460system_prompt: "You are helpful."
20461"#;
20462 let agent = AgentBuilder::from_yaml(yaml)
20463 .unwrap()
20464 .llm(Arc::new(mock))
20465 .auto_configure_features()
20466 .unwrap()
20467 .build()
20468 .unwrap();
20469
20470 let available = agent.get_available_tool_ids().await.unwrap();
20471 assert!(available.is_empty());
20472 }
20473
20474 #[tokio::test]
20475 async fn runtime_scope_cannot_widen_omitted_or_empty_yaml_grants() {
20476 for tools in ["", "tools: []"] {
20477 let yaml = format!(
20478 r#"
20479name: RuntimeScopeNoGrantAgent
20480system_prompt: "No ordinary tools are granted."
20481{tools}
20482"#
20483 );
20484 let agent = AgentBuilder::from_yaml(&yaml)
20485 .unwrap()
20486 .llm(Arc::new(mock_with_response("done")))
20487 .auto_configure_features()
20488 .unwrap()
20489 .build()
20490 .unwrap();
20491
20492 agent
20493 .runtime_control()
20494 .set_tool_scope(vec!["calculator".to_string()]);
20495
20496 assert!(agent.get_available_tool_ids().await.unwrap().is_empty());
20497 }
20498 }
20499
20500 #[tokio::test]
20501 async fn runtime_scope_widening_attempt_keeps_only_declared_tools() {
20502 let yaml = r#"
20503name: RuntimeScopeWideningAgent
20504system_prompt: "Runtime scope cannot add authority."
20505tools: [calculator]
20506"#;
20507 let agent = AgentBuilder::from_yaml(yaml)
20508 .unwrap()
20509 .llm(Arc::new(mock_with_response("done")))
20510 .auto_configure_features()
20511 .unwrap()
20512 .build()
20513 .unwrap();
20514
20515 agent
20516 .runtime_control()
20517 .set_tool_scope(vec!["calculator".to_string(), "datetime".to_string()]);
20518
20519 assert_eq!(
20520 agent.get_available_tool_ids().await.unwrap(),
20521 vec!["calculator".to_string()]
20522 );
20523 }
20524
20525 #[tokio::test]
20526 async fn runtime_scope_is_canonical_unique_ordered_and_clear_restores_declared_grant() {
20527 let yaml = r#"
20528name: RuntimeScopeIntersectionAgent
20529system_prompt: "Use only declared tools."
20530tools: [calculator, datetime]
20531"#;
20532 let agent = AgentBuilder::from_yaml(yaml)
20533 .unwrap()
20534 .llm(Arc::new(mock_with_response("done")))
20535 .auto_configure_features()
20536 .unwrap()
20537 .build()
20538 .unwrap();
20539 let mut aliases = ai_agents_tools::ToolAliases::default();
20540 aliases
20541 .names
20542 .insert("en".to_string(), "calculate_alias".to_string());
20543 agent.tools.set_tool_aliases("calculator", aliases);
20544 let control = agent.runtime_control();
20545
20546 control.set_tool_scope(vec![
20547 "datetime".to_string(),
20548 "calculate_alias".to_string(),
20549 "calculator".to_string(),
20550 "unknown".to_string(),
20551 "datetime".to_string(),
20552 ]);
20553 assert_eq!(
20554 agent.get_available_tool_ids().await.unwrap(),
20555 vec!["calculator".to_string(), "datetime".to_string()]
20556 );
20557
20558 control.set_tool_scope(vec!["datetime".to_string()]);
20559 assert_eq!(
20560 agent.get_available_tool_ids().await.unwrap(),
20561 vec!["datetime".to_string()]
20562 );
20563
20564 control.clear_tool_scope_override();
20565 assert_eq!(
20566 agent.get_available_tool_ids().await.unwrap(),
20567 vec!["calculator".to_string(), "datetime".to_string()]
20568 );
20569 }
20570
20571 #[tokio::test]
20572 async fn runtime_scope_preserves_programmatic_registration_as_declared_grant() {
20573 let agent = AgentBuilder::new()
20574 .system_prompt("Use registered tools.")
20575 .llm(Arc::new(mock_with_response("done")))
20576 .tool(Arc::new(ContextEchoTool))
20577 .tool(Arc::new(SlowTool))
20578 .build()
20579 .unwrap();
20580
20581 agent.runtime_control().set_tool_scope(vec![
20582 "Context Echo".to_string(),
20583 "context_echo".to_string(),
20584 "unknown".to_string(),
20585 ]);
20586
20587 assert_eq!(
20588 agent.get_available_tool_ids().await.unwrap(),
20589 vec!["context_echo".to_string()]
20590 );
20591 }
20592
20593 #[tokio::test]
20594 async fn nested_state_scopes_intersect_every_ancestor_with_aliases() {
20595 let yaml = r#"
20596name: NestedStateScopeAgent
20597system_prompt: "Honor every state scope."
20598tools: [calculator, datetime, echo]
20599states:
20600 initial: root
20601 states:
20602 root:
20603 tools: [calculate_alias, datetime]
20604 initial: middle
20605 states:
20606 middle:
20607 initial: leaf
20608 states:
20609 leaf:
20610 tools: [datetime_alias, echo]
20611"#;
20612 let agent = AgentBuilder::from_yaml(yaml)
20613 .unwrap()
20614 .llm(Arc::new(mock_with_response("done")))
20615 .auto_configure_features()
20616 .unwrap()
20617 .build()
20618 .unwrap();
20619 let mut calculator_aliases = ai_agents_tools::ToolAliases::default();
20620 calculator_aliases
20621 .names
20622 .insert("en".to_string(), "calculate_alias".to_string());
20623 agent
20624 .tools
20625 .set_tool_aliases("calculator", calculator_aliases);
20626 let mut datetime_aliases = ai_agents_tools::ToolAliases::default();
20627 datetime_aliases
20628 .names
20629 .insert("en".to_string(), "datetime_alias".to_string());
20630 agent.tools.set_tool_aliases("datetime", datetime_aliases);
20631 agent.runtime_control().set_tool_scope(vec![
20632 "unknown".to_string(),
20633 "datetime_alias".to_string(),
20634 "calculate_alias".to_string(),
20635 "datetime".to_string(),
20636 ]);
20637
20638 assert_eq!(agent.current_state().as_deref(), Some("root.middle.leaf"));
20639 assert_eq!(
20640 agent.get_available_tool_ids().await.unwrap(),
20641 vec!["datetime".to_string()]
20642 );
20643 }
20644
20645 #[tokio::test]
20646 async fn ancestor_empty_state_scope_denies_omitted_descendants() {
20647 let yaml = r#"
20648name: NestedEmptyStateScopeAgent
20649system_prompt: "An empty ancestor scope denies all tools."
20650tools: [calculator]
20651states:
20652 initial: root
20653 states:
20654 root:
20655 tools: []
20656 initial: middle
20657 states:
20658 middle:
20659 initial: leaf
20660 states:
20661 leaf: {}
20662"#;
20663 let agent = AgentBuilder::from_yaml(yaml)
20664 .unwrap()
20665 .llm(Arc::new(mock_with_response("done")))
20666 .auto_configure_features()
20667 .unwrap()
20668 .build()
20669 .unwrap();
20670
20671 assert!(agent.get_available_tool_ids().await.unwrap().is_empty());
20672 }
20673
20674 #[tokio::test]
20675 async fn state_change_during_approval_invalidates_the_reviewed_authority() {
20676 let active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
20677 let max_active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
20678 let entered = Arc::new(tokio::sync::Barrier::new(2));
20679 let release = Arc::new(tokio::sync::Notify::new());
20680 let handler = Arc::new(BlockingApprovalHandler {
20681 entered: Arc::clone(&entered),
20682 release: Arc::clone(&release),
20683 result: ApprovalResult::Approved,
20684 });
20685 let yaml = r#"
20686name: ApprovalStateGenerationAgent
20687system_prompt: "State authority may change during approval."
20688tools: [locked_write]
20689states:
20690 initial: first
20691 states:
20692 first:
20693 tools: [locked_write]
20694 second:
20695 tools: [locked_write]
20696"#;
20697 let agent = Arc::new(
20698 AgentBuilder::from_yaml(yaml)
20699 .unwrap()
20700 .llm(Arc::new(mock_with_response("done")))
20701 .tool(Arc::new(LockedWriteTool {
20702 active: Arc::clone(&active),
20703 max_active: Arc::clone(&max_active),
20704 }))
20705 .tool_security(ToolSecurityEngine::new(approval_security_config(true)))
20706 .hitl_engine(HITLEngine::new(ai_agents_hitl::HITLConfig::default()))
20707 .approval_handler(handler)
20708 .build()
20709 .unwrap(),
20710 );
20711 let running = Arc::clone(&agent);
20712 let call = tokio::spawn(async move {
20713 running
20714 .invoke_tool(ToolExecutionRequest::new(
20715 "approval-state-generation",
20716 "locked_write",
20717 serde_json::json!({"path": "./state-generation.txt"}),
20718 ToolCallSource::Manual,
20719 ))
20720 .await
20721 .unwrap()
20722 });
20723
20724 entered.wait().await;
20725 agent.transition_to("second").await.unwrap();
20726 release.notify_one();
20727 let record = call.await.unwrap();
20728
20729 assert!(!record.executed);
20730 assert!(record.output.contains("Approval became stale"));
20731 assert_eq!(max_active.load(Ordering::SeqCst), 0);
20732 }
20733
20734 #[tokio::test]
20735 async fn state_change_while_waiting_for_resource_lock_fails_final_admission() {
20736 let holder_gate = PathMutationGate::new();
20737 let waiter_gate = PathMutationGate::new();
20738 let yaml = r#"
20739name: LockedStateGenerationAgent
20740system_prompt: "State authority must remain stable through admission."
20741tools: [state_lock_holder, state_lock_waiter]
20742states:
20743 initial: first
20744 states:
20745 first:
20746 tools: [state_lock_holder, state_lock_waiter]
20747 second:
20748 tools: [state_lock_holder, state_lock_waiter]
20749"#;
20750 let agent = Arc::new(
20751 AgentBuilder::from_yaml(yaml)
20752 .unwrap()
20753 .llm(Arc::new(mock_with_response("done")))
20754 .tool(Arc::new(BlockingPathMutationTool {
20755 id: "state_lock_holder",
20756 path_fields: vec![ai_agents_core::PathPolicyBinding::write("path")],
20757 gate: holder_gate.clone(),
20758 }))
20759 .tool(Arc::new(BlockingPathMutationTool {
20760 id: "state_lock_waiter",
20761 path_fields: vec![ai_agents_core::PathPolicyBinding::write("path")],
20762 gate: waiter_gate.clone(),
20763 }))
20764 .build()
20765 .unwrap(),
20766 );
20767 let holder_call = {
20768 let agent = Arc::clone(&agent);
20769 tokio::spawn(async move {
20770 agent
20771 .invoke_tool(ToolExecutionRequest::new(
20772 "state-lock-holder",
20773 "state_lock_holder",
20774 serde_json::json!({"path": "./shared-state-path.txt"}),
20775 ToolCallSource::Manual,
20776 ))
20777 .await
20778 .unwrap()
20779 })
20780 };
20781 holder_gate.wait_until_entered().await;
20782 let waiter_call = {
20783 let agent = Arc::clone(&agent);
20784 tokio::spawn(async move {
20785 agent
20786 .invoke_tool(ToolExecutionRequest::new(
20787 "state-lock-waiter",
20788 "state_lock_waiter",
20789 serde_json::json!({"path": "./shared-state-path.txt"}),
20790 ToolCallSource::Manual,
20791 ))
20792 .await
20793 .unwrap()
20794 })
20795 };
20796
20797 wait_for_resource_lock_strong_count(&agent.resource_locks, 2).await;
20798 agent.transition_to("second").await.unwrap();
20799 holder_gate.release();
20800 let holder_record = holder_call.await.unwrap();
20801 let waiter_record = waiter_call.await.unwrap();
20802
20803 assert!(holder_record.success);
20804 assert!(!waiter_record.executed);
20805 assert!(
20806 waiter_record
20807 .output
20808 .contains("state scope changed before admission")
20809 );
20810 assert!(!waiter_gate.entered.load(Ordering::SeqCst));
20811 }
20812
20813 #[tokio::test]
20814 async fn test_state_tools_cannot_widen_top_level_grant() {
20815 let mock = mock_with_response("hello");
20816 let yaml = r#"
20817name: NarrowToolsAgent
20818system_prompt: "You are helpful."
20819tools:
20820 - calculator
20821states:
20822 initial: current
20823 states:
20824 current:
20825 tools: [datetime]
20826"#;
20827 let agent = AgentBuilder::from_yaml(yaml)
20828 .unwrap()
20829 .llm(Arc::new(mock))
20830 .auto_configure_features()
20831 .unwrap()
20832 .build()
20833 .unwrap();
20834
20835 let available = agent.get_available_tool_ids().await.unwrap();
20836 assert!(available.is_empty());
20837 }
20838
20839 #[tokio::test]
20841 async fn test_integration_tool_execution() {
20842 let mock = mock_with_responses(vec![
20844 r#"I'll calculate that for you.
20846{"tool": "calculator", "arguments": {"expression": "2+2"}}"#,
20847 "The answer is 4.",
20849 ]);
20850 let observed = mock.clone();
20851 let mut tools = ai_agents_tools::ToolRegistry::new();
20852 tools
20853 .register(Arc::new(ai_agents_tools::CalculatorTool))
20854 .unwrap();
20855
20856 let agent = AgentBuilder::new()
20857 .system_prompt("You are a calculator assistant.")
20858 .llm(Arc::new(mock))
20859 .tools(tools)
20860 .build()
20861 .unwrap();
20862
20863 let response = agent.chat("What is 2+2?").await.unwrap();
20864
20865 assert_eq!(response.content, "The answer is 4.");
20866 assert_eq!(response.tool_calls.as_ref().map(Vec::len), Some(1));
20867 assert_eq!(
20868 observed.call_count(),
20869 2,
20870 "tool result must trigger a second LLM call"
20871 );
20872 let history = agent.tool_call_history();
20873 assert_eq!(history.len(), 1);
20874 assert_eq!(history[0].tool_id, "calculator");
20875 assert_eq!(
20876 history[0].result.get("result"),
20877 Some(&serde_json::json!(4.0)),
20878 "{:?}",
20879 history[0].result
20880 );
20881 }
20882
20883 #[test]
20886 fn legacy_tool_call_marker_is_plain_text() {
20887 let agent = AgentBuilder::new()
20888 .system_prompt("x")
20889 .llm(Arc::new(mock_with_response("x")))
20890 .build()
20891 .unwrap();
20892 let parsed = agent
20893 .parse_tool_calls(
20894 r#"[TOOL_CALL: {"name": "calculator", "arguments": {"expression": "2+2"}}]"#,
20895 )
20896 .unwrap();
20897 assert!(parsed.is_none());
20898 }
20899
20900 #[tokio::test]
20901 async fn test_tool_hitl_rejection_finalizes_blocking_turn() {
20902 let responses = Arc::new(std::sync::atomic::AtomicUsize::new(0));
20903 let hooks = Arc::new(ResponseCountingHooks {
20904 responses: Arc::clone(&responses),
20905 });
20906 let mock = mock_with_response(r#"{"tool":"echo","arguments":{"message":"hello"}}"#);
20907 let yaml = r#"
20908name: ToolRejectAgent
20909system_prompt: "You use tools when requested."
20910tools:
20911 - echo
20912hitl:
20913 tools:
20914 echo:
20915 require_approval: true
20916 approval_message: "Approve echo?"
20917"#;
20918 let agent = AgentBuilder::from_yaml(yaml)
20919 .unwrap()
20920 .llm(Arc::new(mock))
20921 .auto_configure_features()
20922 .unwrap()
20923 .hooks(hooks)
20924 .build()
20925 .unwrap();
20926
20927 let response = agent.chat("echo hello").await.unwrap();
20928
20929 assert!(
20930 response.content.contains("Operation cancelled"),
20931 "unexpected response: {}",
20932 response.content
20933 );
20934 assert_eq!(responses.load(Ordering::SeqCst), 1);
20935 let messages = agent.memory.get_messages(None).await.unwrap();
20936 assert_eq!(messages.len(), 3);
20937 assert_eq!(messages[0].content, "echo hello");
20938 assert!(messages[1].content.contains("\"tool\":\"echo\""));
20939 assert!(messages[2].content.contains("rejected by the approver"));
20940 }
20941
20942 #[tokio::test]
20943 async fn test_tool_hitl_rejection_finalizes_streaming_turn() {
20944 use futures::StreamExt;
20945
20946 let responses = Arc::new(std::sync::atomic::AtomicUsize::new(0));
20947 let hooks = Arc::new(ResponseCountingHooks {
20948 responses: Arc::clone(&responses),
20949 });
20950 let mock = mock_with_response(r#"{"tool":"echo","arguments":{"message":"hello"}}"#);
20951 let yaml = r#"
20952name: ToolRejectStreamingAgent
20953system_prompt: "You use tools when requested."
20954tools:
20955 - echo
20956streaming:
20957 enabled: true
20958hitl:
20959 tools:
20960 echo:
20961 require_approval: true
20962 approval_message: "Approve echo?"
20963"#;
20964 let agent = AgentBuilder::from_yaml(yaml)
20965 .unwrap()
20966 .llm(Arc::new(mock))
20967 .auto_configure_features()
20968 .unwrap()
20969 .hooks(hooks)
20970 .build()
20971 .unwrap();
20972
20973 let mut stream = agent.chat_stream("echo hello").await.unwrap();
20974 let mut terminal_error = String::new();
20975 let mut done = false;
20976 while let Some(chunk) = stream.next().await {
20977 match chunk {
20978 StreamChunk::Error { message } => terminal_error = message,
20979 StreamChunk::Done {} => {
20980 done = true;
20981 break;
20982 }
20983 _ => {}
20984 }
20985 }
20986
20987 assert!(done);
20988 assert!(
20989 terminal_error.contains("Operation cancelled"),
20990 "unexpected terminal error: {}",
20991 terminal_error
20992 );
20993 assert_eq!(responses.load(Ordering::SeqCst), 1);
20994 let messages = agent.memory.get_messages(None).await.unwrap();
20995 assert_eq!(messages.len(), 3);
20996 assert_eq!(messages[0].content, "echo hello");
20997 assert!(messages[1].content.contains("\"tool\":\"echo\""));
20998 assert!(messages[2].content.contains("rejected by the approver"));
20999 }
21000
21001 #[tokio::test]
21002 async fn tool_hitl_rejection_preserves_legacy_error_but_finalizes_event_stream() {
21003 let mock = mock_with_response(r#"{"tool":"echo","arguments":{"message":"hello"}}"#);
21004 let yaml = r#"
21005name: ToolRejectEventAgent
21006system_prompt: "You use tools when requested."
21007tools:
21008 - echo
21009streaming:
21010 enabled: true
21011hitl:
21012 tools:
21013 echo:
21014 require_approval: true
21015 approval_message: "Approve echo?"
21016"#;
21017 let agent = AgentBuilder::from_yaml(yaml)
21018 .unwrap()
21019 .llm(Arc::new(mock))
21020 .auto_configure_features()
21021 .unwrap()
21022 .build()
21023 .unwrap();
21024
21025 let mut stream = agent.chat_stream_events("echo hello").await.unwrap();
21026 let mut error_seen = false;
21027 let mut final_response = None;
21028 while let Some(event) = stream.next().await {
21029 match event {
21030 AgentStreamEvent::Chunk(StreamChunk::Error { .. }) => error_seen = true,
21031 AgentStreamEvent::Final(response) => final_response = Some(response),
21032 AgentStreamEvent::Chunk(_) => {}
21033 }
21034 }
21035
21036 assert!(!error_seen);
21037 assert!(
21038 final_response
21039 .is_some_and(|response| { response.content.contains("Operation cancelled") })
21040 );
21041 }
21042
21043 #[tokio::test]
21044 async fn test_pre_response_guard_transition_skips_old_state_llm() {
21045 let mock = mock_with_response("Billing state response");
21046 let call_counter = mock.clone();
21047 let yaml = r#"
21048name: OptimizedStateAgent
21049system_prompt: "You route before answering."
21050runtime:
21051 optimization:
21052 enabled: true
21053 pre_response_deterministic_transitions: true
21054states:
21055 initial: greeting
21056 states:
21057 greeting:
21058 prompt: "Old state prompt that should be skipped."
21059 transitions:
21060 - to: billing
21061 guard:
21062 context:
21063 topic:
21064 eq: billing
21065 timing: pre_response
21066 billing:
21067 prompt: "Answer from the billing state."
21068"#;
21069 let agent = AgentBuilder::from_yaml(yaml)
21070 .unwrap()
21071 .llm(Arc::new(mock))
21072 .build()
21073 .unwrap();
21074 agent
21075 .set_context("topic", serde_json::json!("billing"))
21076 .unwrap();
21077
21078 let response = agent.chat("I need billing help").await.unwrap();
21079
21080 assert_eq!(agent.current_state().as_deref(), Some("billing"));
21081 assert_eq!(response.content, "Billing state response");
21082 assert_eq!(call_counter.call_count(), 1);
21083 assert_eq!(agent.actor_facts().len(), 0);
21084 }
21085
21086 #[tokio::test]
21087 async fn test_set_context_supports_dotted_paths_for_pre_response_guards() {
21088 let mock = mock_with_response("Billing state response");
21089 let call_counter = mock.clone();
21090 let yaml = r#"
21091name: OptimizedStateAgent
21092system_prompt: "You route before answering."
21093runtime:
21094 optimization:
21095 enabled: true
21096 pre_response_deterministic_transitions: true
21097context:
21098 request:
21099 type: runtime
21100 default:
21101 topic: general
21102states:
21103 initial: greeting
21104 states:
21105 greeting:
21106 prompt: "Old state prompt that should be skipped."
21107 transitions:
21108 - to: billing
21109 guard:
21110 context:
21111 request.topic:
21112 eq: billing
21113 timing: pre_response
21114 billing:
21115 prompt: "Answer from the billing state."
21116"#;
21117 let agent = AgentBuilder::from_yaml(yaml)
21118 .unwrap()
21119 .llm(Arc::new(mock))
21120 .build()
21121 .unwrap();
21122 agent
21123 .set_context("request.topic", serde_json::json!("billing"))
21124 .unwrap();
21125
21126 let response = agent.chat("I need billing help").await.unwrap();
21127
21128 assert_eq!(agent.current_state().as_deref(), Some("billing"));
21129 assert_eq!(response.content, "Billing state response");
21130 assert_eq!(call_counter.call_count(), 1);
21131 assert_eq!(
21132 agent.get_context().get("request"),
21133 Some(&serde_json::json!({"topic": "billing"}))
21134 );
21135 }
21136
21137 #[tokio::test]
21138 async fn test_pre_response_rejection_does_not_commit_staged_context_or_user() {
21139 let mock = mock_with_response("billing");
21140 let yaml = r#"
21141name: OptimizedStateAgent
21142system_prompt: "You route before answering."
21143runtime:
21144 optimization:
21145 enabled: true
21146 pre_response_deterministic_transitions: true
21147hitl:
21148 states:
21149 billing:
21150 on_enter: require_approval
21151 approval_message: "Approve billing route?"
21152states:
21153 initial: greeting
21154 states:
21155 greeting:
21156 prompt: "Old state prompt."
21157 extract:
21158 - key: topic
21159 description: "Support topic"
21160 transitions:
21161 - to: billing
21162 guard:
21163 context:
21164 topic:
21165 eq: billing
21166 timing: pre_response
21167 run_extractors: true
21168 billing:
21169 prompt: "Billing state."
21170"#;
21171 let agent = AgentBuilder::from_yaml(yaml)
21172 .unwrap()
21173 .llm(Arc::new(mock))
21174 .build()
21175 .unwrap();
21176
21177 let response = agent
21178 .try_pre_response_transition("billing please")
21179 .await
21180 .unwrap();
21181
21182 assert!(response.is_none());
21183 assert_eq!(agent.current_state().as_deref(), Some("greeting"));
21184 assert!(!agent.get_context().contains_key("topic"));
21185 assert_eq!(agent.memory.get_messages(None).await.unwrap().len(), 0);
21186 }
21187
21188 #[tokio::test]
21189 async fn test_pre_response_extractor_commits_context_on_winning_path() {
21190 let mock = mock_with_responses(vec!["billing", "Billing response"]);
21191 let yaml = r#"
21192name: OptimizedStateAgent
21193system_prompt: "You route before answering."
21194runtime:
21195 optimization:
21196 enabled: true
21197 pre_response_deterministic_transitions: true
21198states:
21199 initial: greeting
21200 states:
21201 greeting:
21202 prompt: "Old state prompt."
21203 extract:
21204 - key: topic
21205 description: "Support topic"
21206 transitions:
21207 - to: billing
21208 guard:
21209 context:
21210 topic:
21211 eq: billing
21212 timing: pre_response
21213 run_extractors: true
21214 billing:
21215 prompt: "Billing state."
21216"#;
21217 let agent = AgentBuilder::from_yaml(yaml)
21218 .unwrap()
21219 .llm(Arc::new(mock))
21220 .build()
21221 .unwrap();
21222
21223 let response = agent.chat("billing please").await.unwrap();
21224
21225 assert_eq!(agent.current_state().as_deref(), Some("billing"));
21226 assert_eq!(response.content, "Billing response");
21227 assert_eq!(
21228 agent.get_context().get("topic"),
21229 Some(&serde_json::json!("billing"))
21230 );
21231 }
21232
21233 #[tokio::test]
21234 async fn test_pre_response_extractor_miss_does_not_mutate_context() {
21235 let mock = mock_with_response("__NONE__");
21236 let yaml = r#"
21237name: OptimizedStateAgent
21238system_prompt: "You route before answering."
21239runtime:
21240 optimization:
21241 enabled: true
21242 pre_response_deterministic_transitions: true
21243states:
21244 initial: greeting
21245 states:
21246 greeting:
21247 prompt: "Old state prompt."
21248 extract:
21249 - key: topic
21250 description: "Support topic"
21251 transitions:
21252 - to: billing
21253 guard:
21254 context:
21255 topic:
21256 eq: billing
21257 timing: pre_response
21258 run_extractors: true
21259 billing:
21260 prompt: "Billing state."
21261"#;
21262 let agent = AgentBuilder::from_yaml(yaml)
21263 .unwrap()
21264 .llm(Arc::new(mock))
21265 .build()
21266 .unwrap();
21267
21268 let response = agent.try_pre_response_transition("hello").await.unwrap();
21269
21270 assert!(response.is_none());
21271 assert_eq!(agent.current_state().as_deref(), Some("greeting"));
21272 assert!(!agent.get_context().contains_key("topic"));
21273 }
21274
21275 #[tokio::test]
21276 async fn test_default_guard_transition_stays_post_response() {
21277 let mock = mock_with_responses(vec!["Greeting response", "Billing response"]);
21278 let call_counter = mock.clone();
21279 let yaml = r#"
21280name: TimingAgent
21281system_prompt: "You route carefully."
21282runtime:
21283 optimization:
21284 enabled: true
21285 pre_response_deterministic_transitions: true
21286states:
21287 initial: greeting
21288 states:
21289 greeting:
21290 prompt: "Old state prompt."
21291 transitions:
21292 - to: billing
21293 guard:
21294 context:
21295 topic:
21296 eq: billing
21297 billing:
21298 prompt: "Billing state."
21299"#;
21300 let agent = AgentBuilder::from_yaml(yaml)
21301 .unwrap()
21302 .llm(Arc::new(mock))
21303 .build()
21304 .unwrap();
21305 agent
21306 .set_context("topic", serde_json::json!("billing"))
21307 .unwrap();
21308
21309 let response = agent.chat("billing please").await.unwrap();
21310
21311 assert_eq!(agent.current_state().as_deref(), Some("billing"));
21312 assert_eq!(response.content, "Billing response");
21313 assert_eq!(call_counter.call_count(), 2);
21314 }
21315
21316 #[tokio::test]
21317 async fn test_explicit_post_response_guard_transition_stays_post_response() {
21318 let mock = mock_with_responses(vec!["Greeting response", "Billing response"]);
21319 let call_counter = mock.clone();
21320 let yaml = r#"
21321name: TimingAgent
21322system_prompt: "You route carefully."
21323runtime:
21324 optimization:
21325 enabled: true
21326 pre_response_deterministic_transitions: true
21327states:
21328 initial: greeting
21329 states:
21330 greeting:
21331 prompt: "Old state prompt."
21332 transitions:
21333 - to: billing
21334 guard:
21335 context:
21336 topic:
21337 eq: billing
21338 timing: post_response
21339 billing:
21340 prompt: "Billing state."
21341"#;
21342 let agent = AgentBuilder::from_yaml(yaml)
21343 .unwrap()
21344 .llm(Arc::new(mock))
21345 .build()
21346 .unwrap();
21347 agent
21348 .set_context("topic", serde_json::json!("billing"))
21349 .unwrap();
21350
21351 let response = agent.chat("billing please").await.unwrap();
21352
21353 assert_eq!(agent.current_state().as_deref(), Some("billing"));
21354 assert_eq!(response.content, "Billing response");
21355 assert_eq!(call_counter.call_count(), 2);
21356 }
21357
21358 #[tokio::test]
21359 async fn test_pre_response_extractors_are_transition_scoped() {
21360 let mock = mock_with_responses(vec!["billing", "Billing response"]);
21361 let yaml = r#"
21362name: ScopedExtractorAgent
21363system_prompt: "You route carefully."
21364runtime:
21365 optimization:
21366 enabled: true
21367 pre_response_deterministic_transitions: true
21368states:
21369 initial: greeting
21370 states:
21371 greeting:
21372 prompt: "Old state prompt."
21373 extract:
21374 - key: topic
21375 description: "Support topic"
21376 transitions:
21377 - to: wrong
21378 guard:
21379 context:
21380 topic:
21381 eq: billing
21382 timing: pre_response
21383 - to: billing
21384 guard:
21385 context:
21386 topic:
21387 eq: billing
21388 timing: pre_response
21389 run_extractors: true
21390 wrong:
21391 prompt: "Wrong state."
21392 billing:
21393 prompt: "Billing state."
21394"#;
21395 let agent = AgentBuilder::from_yaml(yaml)
21396 .unwrap()
21397 .llm(Arc::new(mock))
21398 .build()
21399 .unwrap();
21400
21401 let response = agent.chat("billing please").await.unwrap();
21402
21403 assert_eq!(agent.current_state().as_deref(), Some("billing"));
21404 assert_eq!(response.content, "Billing response");
21405 }
21406
21407 #[tokio::test]
21408 async fn test_pre_response_resolved_intent_routes_early() {
21409 let mock = mock_with_response("Billing response");
21410 let yaml = r#"
21411name: IntentAgent
21412system_prompt: "You route carefully."
21413runtime:
21414 optimization:
21415 enabled: true
21416 pre_response_deterministic_transitions: true
21417states:
21418 initial: greeting
21419 states:
21420 greeting:
21421 prompt: "Old state prompt."
21422 transitions:
21423 - to: billing
21424 intent: billing
21425 timing: pre_response
21426 billing:
21427 prompt: "Billing state."
21428"#;
21429 let agent = AgentBuilder::from_yaml(yaml)
21430 .unwrap()
21431 .llm(Arc::new(mock))
21432 .build()
21433 .unwrap();
21434 agent
21435 .set_context("resolved_intent", serde_json::json!("billing"))
21436 .unwrap();
21437
21438 let response = agent
21439 .try_pre_response_transition("I need billing help")
21440 .await
21441 .unwrap()
21442 .unwrap();
21443
21444 assert_eq!(agent.current_state().as_deref(), Some("billing"));
21445 assert_eq!(response.content, "Billing response");
21446 }
21447
21448 #[tokio::test]
21449 async fn test_background_overflow_error_surfaces() {
21450 let mut config = RuntimeConfig::default();
21451 config.optimization.enabled = true;
21452 config.optimization.post_turn.max_background_tasks = 1;
21453 config.optimization.post_turn.on_background_overflow = BackgroundOverflowPolicy::Error;
21454 let policy = crate::optimization::MaintenanceTaskPolicy {
21455 mode: MaintenanceMode::Background,
21456 await_before_next_turn: AwaitBeforeNextTurn::Always,
21457 };
21458 let agent = AgentBuilder::new()
21459 .system_prompt("You are helpful.")
21460 .llm(Arc::new(mock_with_response("ok")))
21461 .build()
21462 .unwrap()
21463 .with_runtime_config(config);
21464 agent
21465 .background_maintenance
21466 .spawn(None, async { std::future::pending::<Result<()>>().await })
21467 .unwrap();
21468
21469 let result = agent
21470 .spawn_or_handle_background(None, async { Ok(()) }, "facts", &policy)
21471 .await;
21472
21473 assert!(result.is_err());
21474 }
21475
21476 #[tokio::test]
21477 async fn test_speculative_reasoning_low_cap_uses_serial_reasoning() {
21478 let default_mock = mock_with_response("Plain draft response");
21479 let router_mock = mock_with_response("cot");
21480 let router_counter = router_mock.clone();
21481 let yaml = r#"
21482name: ReasoningReservationAgent
21483system_prompt: "You answer plainly unless reasoning wins."
21484llm:
21485 default: default
21486 router: router
21487observability:
21488 enabled: true
21489 export:
21490 write_raw_events: true
21491reasoning:
21492 mode: auto
21493 judge_llm: router
21494runtime:
21495 optimization:
21496 enabled: true
21497 max_speculative_llm_calls_per_turn: 1
21498 speculative_reasoning_auto: true
21499 max_parallel_runtime_tasks: 2
21500"#;
21501 let agent = AgentBuilder::from_yaml(yaml)
21502 .unwrap()
21503 .llm_alias("default", Arc::new(default_mock))
21504 .llm_alias("router", Arc::new(router_mock))
21505 .build()
21506 .unwrap();
21507
21508 let response = agent.chat("hello").await.unwrap();
21509
21510 assert_eq!(response.content, "Plain draft response");
21511 assert_eq!(router_counter.call_count(), 1);
21512 let events = agent.observability().unwrap().raw_events();
21513 assert!(!events.iter().any(|event| {
21514 event.dimensions.get("commit_behavior") == Some(&"reasoning_decision".to_string())
21515 }));
21516 }
21517
21518 #[tokio::test]
21519 async fn test_forced_reasoning_skips_plain_speculative_draft() {
21520 let mock = mock_with_response("Reasoned response");
21521 let yaml = r#"
21522name: ForcedReasoningAgent
21523system_prompt: "You reason before answering."
21524observability:
21525 enabled: true
21526 export:
21527 write_raw_events: true
21528reasoning:
21529 mode: cot
21530runtime:
21531 optimization:
21532 enabled: true
21533 max_speculative_llm_calls_per_turn: 2
21534 speculative_state_transitions: true
21535 max_parallel_runtime_tasks: 2
21536states:
21537 initial: triage
21538 states:
21539 triage:
21540 prompt: "Answer from triage."
21541 transitions:
21542 - to: billing
21543 guard:
21544 context:
21545 route:
21546 eq: billing
21547 timing: parallel
21548 billing:
21549 prompt: "Billing state."
21550"#;
21551 let agent = AgentBuilder::from_yaml(yaml)
21552 .unwrap()
21553 .llm(Arc::new(mock))
21554 .build()
21555 .unwrap();
21556
21557 let response = agent.chat("hello").await.unwrap();
21558
21559 assert_eq!(response.content, "Reasoned response");
21560 let events = agent.observability().unwrap().raw_events();
21561 assert!(
21562 !events
21563 .iter()
21564 .any(|event| event.dimensions.contains_key("branch_status"))
21565 );
21566 }
21567
21568 #[tokio::test]
21569 async fn test_speculative_skill_low_cap_uses_serial_skill_route() {
21570 let default_mock = mock_with_response("Skill committed response");
21571 let router_mock = mock_with_response("helper");
21572 let router_counter = router_mock.clone();
21573 let yaml = r#"
21574name: SkillReservationAgent
21575system_prompt: "Use skills when they match."
21576llm:
21577 default: default
21578 router: router
21579observability:
21580 enabled: true
21581 export:
21582 write_raw_events: true
21583runtime:
21584 optimization:
21585 enabled: true
21586 max_speculative_llm_calls_per_turn: 1
21587 speculative_skill_routing: true
21588 max_parallel_runtime_tasks: 2
21589skills:
21590 - id: helper
21591 description: "Answer helper requests"
21592 trigger: "User asks for helper"
21593 steps:
21594 - prompt: "Answer the helper request: {{ user_input }}"
21595"#;
21596 let agent = AgentBuilder::from_yaml(yaml)
21597 .unwrap()
21598 .llm_alias("default", Arc::new(default_mock))
21599 .llm_alias("router", Arc::new(router_mock))
21600 .build()
21601 .unwrap();
21602
21603 let response = agent.chat("please use helper").await.unwrap();
21604
21605 assert_eq!(response.content, "Skill committed response");
21606 assert_eq!(router_counter.call_count(), 1);
21607 let events = agent.observability().unwrap().raw_events();
21608 assert!(
21609 !events
21610 .iter()
21611 .any(|event| event.dimensions.contains_key("branch_status"))
21612 );
21613 }
21614
21615 #[tokio::test]
21616 async fn test_parallel_transition_low_cap_allows_deterministic_route() {
21617 let mock = mock_with_response("unused");
21618 let call_counter = mock.clone();
21619 let yaml = r#"
21620name: ParallelTransitionLowCapAgent
21621system_prompt: "Route before stale responses when safe."
21622runtime:
21623 optimization:
21624 enabled: true
21625 max_speculative_llm_calls_per_turn: 1
21626 speculative_state_transitions: true
21627 max_parallel_runtime_tasks: 2
21628states:
21629 initial: triage
21630 states:
21631 triage:
21632 prompt: "Triage state."
21633 transitions:
21634 - to: billing
21635 guard:
21636 context:
21637 route:
21638 eq: billing
21639 timing: parallel
21640 billing:
21641 prompt: "Billing state."
21642"#;
21643 let agent = AgentBuilder::from_yaml(yaml)
21644 .unwrap()
21645 .llm(Arc::new(mock))
21646 .build()
21647 .unwrap();
21648 agent
21649 .set_context("route", serde_json::json!("billing"))
21650 .unwrap();
21651 agent.update_active_turn_context("billing help", HashMap::new());
21652 assert!(
21653 agent.reserve_active_speculative_llm_call(
21654 RuntimeOptimizationKind::ParallelStateTransition
21655 )
21656 );
21657
21658 let selection = agent
21659 .select_parallel_transition_candidate("billing help")
21660 .await
21661 .unwrap();
21662 agent.end_root_turn();
21663
21664 match selection {
21665 ParallelTransitionSelection::Candidate(candidate) => {
21666 assert_eq!(candidate.target(), "billing");
21667 }
21668 ParallelTransitionSelection::NoMatch => panic!("deterministic route did not match"),
21669 ParallelTransitionSelection::ReservationExhausted => {
21670 panic!("deterministic route consumed LLM budget")
21671 }
21672 }
21673 assert_eq!(call_counter.call_count(), 0);
21674 }
21675
21676 #[tokio::test]
21677 async fn speculative_transition_drops_loser_before_state_actions() {
21678 let lock = Arc::new(tokio::sync::Mutex::new(()));
21679 let first_started = Arc::new(tokio::sync::Notify::new());
21680 let first_dropped = Arc::new(AtomicBool::new(false));
21681 let committed_after_drop = Arc::new(AtomicBool::new(false));
21682 let default = Arc::new(FirstCallLockingProvider {
21683 lock,
21684 first_started: Arc::clone(&first_started),
21685 first_dropped: Arc::clone(&first_dropped),
21686 committed_after_drop: Arc::clone(&committed_after_drop),
21687 calls: AtomicU64::new(0),
21688 });
21689 let router = Arc::new(RoutingAfterProviderStart {
21690 provider_started: first_started,
21691 });
21692 let yaml = r#"
21693name: SpeculativeCancellationAgent
21694system_prompt: "Route before committed work."
21695llm:
21696 default: default
21697 router: router
21698runtime:
21699 optimization:
21700 enabled: true
21701 max_speculative_llm_calls_per_turn: 2
21702 speculative_state_transitions: true
21703 max_parallel_runtime_tasks: 2
21704states:
21705 initial: triage
21706 states:
21707 triage:
21708 prompt: "Triage state."
21709 transitions:
21710 - to: technical
21711 when: "The request needs technical support"
21712 timing: parallel
21713 technical:
21714 prompt: "Technical state."
21715 on_enter:
21716 - prompt: "Prepare technical context."
21717 llm: default
21718 store_as: preparation
21719"#;
21720 let agent = AgentBuilder::from_yaml(yaml)
21721 .unwrap()
21722 .llm_alias("default", default)
21723 .llm_alias("router", router)
21724 .build()
21725 .unwrap();
21726
21727 let response = tokio::time::timeout(
21728 std::time::Duration::from_secs(2),
21729 agent.chat("I cannot log in because of AUTH-17."),
21730 )
21731 .await
21732 .expect("committed work must not wait on the losing provider future")
21733 .unwrap();
21734
21735 assert_eq!(response.content, "Committed technical response.");
21736 assert_eq!(agent.current_state().as_deref(), Some("technical"));
21737 assert!(first_dropped.load(Ordering::SeqCst));
21738 assert!(committed_after_drop.load(Ordering::SeqCst));
21739 }
21740
21741 #[tokio::test]
21742 async fn buffered_transition_drops_stale_stream_before_redispatch() {
21743 use futures::StreamExt;
21744
21745 let lock = Arc::new(tokio::sync::Mutex::new(()));
21746 let stream_started = Arc::new(tokio::sync::Notify::new());
21747 let stream_dropped = Arc::new(AtomicBool::new(false));
21748 let committed_after_drop = Arc::new(AtomicBool::new(false));
21749 let default = Arc::new(BufferedLockingProvider {
21750 lock,
21751 stream_started: Arc::clone(&stream_started),
21752 stream_dropped: Arc::clone(&stream_dropped),
21753 committed_after_drop: Arc::clone(&committed_after_drop),
21754 });
21755 let router = Arc::new(RoutingAfterProviderStart {
21756 provider_started: stream_started,
21757 });
21758 let yaml = r#"
21759name: BufferedCancellationAgent
21760system_prompt: "Hide stale streamed output."
21761llm:
21762 default: default
21763 router: router
21764streaming:
21765 enabled: true
21766 buffer_size: 8
21767runtime:
21768 optimization:
21769 enabled: true
21770 max_speculative_llm_calls_per_turn: 2
21771 speculative_state_transitions: true
21772 streaming_policy: buffer_until_routing_done
21773 max_parallel_runtime_tasks: 2
21774states:
21775 initial: triage
21776 states:
21777 triage:
21778 prompt: "Triage state."
21779 transitions:
21780 - to: technical
21781 when: "The request needs technical support"
21782 timing: parallel
21783 technical:
21784 prompt: "Technical state."
21785"#;
21786 let agent = AgentBuilder::from_yaml(yaml)
21787 .unwrap()
21788 .llm_alias("default", default)
21789 .llm_alias("router", router)
21790 .build()
21791 .unwrap();
21792
21793 let content = tokio::time::timeout(std::time::Duration::from_secs(2), async {
21794 let mut stream = agent
21795 .chat_stream("AUTH-17 needs technical help.")
21796 .await
21797 .unwrap();
21798 let mut content = String::new();
21799 while let Some(chunk) = stream.next().await {
21800 match chunk {
21801 StreamChunk::Content { text } => content.push_str(&text),
21802 StreamChunk::Done {} => break,
21803 StreamChunk::Error { message } => panic!("unexpected stream error: {message}"),
21804 _ => {}
21805 }
21806 }
21807 content
21808 })
21809 .await
21810 .expect("redispatch must not wait on the stale streaming future");
21811
21812 assert_eq!(content, "Committed technical response.");
21813 assert_eq!(agent.current_state().as_deref(), Some("technical"));
21814 assert!(stream_dropped.load(Ordering::SeqCst));
21815 assert!(committed_after_drop.load(Ordering::SeqCst));
21816 }
21817
21818 #[tokio::test]
21819 async fn buffered_transition_drops_established_stream_before_redispatch() {
21820 use futures::StreamExt;
21821
21822 let stream_started = Arc::new(tokio::sync::Notify::new());
21823 let stream_dropped = Arc::new(AtomicBool::new(false));
21824 let stream_dropped_notify = Arc::new(tokio::sync::Notify::new());
21825 let committed_after_drop = Arc::new(AtomicBool::new(false));
21826 let default = Arc::new(EstablishedStreamProvider {
21827 stream_started: Arc::clone(&stream_started),
21828 stream_dropped: Arc::clone(&stream_dropped),
21829 stream_dropped_notify,
21830 committed_after_drop: Arc::clone(&committed_after_drop),
21831 });
21832 let router = Arc::new(RoutingAfterProviderStart {
21833 provider_started: stream_started,
21834 });
21835 let yaml = r#"
21836name: EstablishedStreamCancellationAgent
21837system_prompt: "Hide stale streamed output."
21838llm:
21839 default: default
21840 router: router
21841streaming:
21842 enabled: true
21843 buffer_size: 8
21844runtime:
21845 optimization:
21846 enabled: true
21847 max_speculative_llm_calls_per_turn: 2
21848 speculative_state_transitions: true
21849 streaming_policy: buffer_until_routing_done
21850 max_parallel_runtime_tasks: 2
21851states:
21852 initial: triage
21853 states:
21854 triage:
21855 prompt: "Triage state."
21856 transitions:
21857 - to: technical
21858 when: "The request needs technical support"
21859 timing: parallel
21860 technical:
21861 prompt: "Technical state."
21862"#;
21863 let agent = AgentBuilder::from_yaml(yaml)
21864 .unwrap()
21865 .llm_alias("default", default)
21866 .llm_alias("router", router)
21867 .build()
21868 .unwrap();
21869
21870 let content = tokio::time::timeout(std::time::Duration::from_secs(2), async {
21871 let mut stream = agent
21872 .chat_stream("AUTH-17 needs technical help.")
21873 .await
21874 .unwrap();
21875 let mut content = String::new();
21876 while let Some(chunk) = stream.next().await {
21877 match chunk {
21878 StreamChunk::Content { text } => content.push_str(&text),
21879 StreamChunk::Done {} => break,
21880 StreamChunk::Error { message } => panic!("unexpected stream error: {message}"),
21881 _ => {}
21882 }
21883 }
21884 content
21885 })
21886 .await
21887 .expect("redispatch must wait for the established stale stream to be dropped");
21888
21889 assert_eq!(content, "Committed technical response.");
21890 assert_eq!(agent.current_state().as_deref(), Some("technical"));
21891 assert!(stream_dropped.load(Ordering::SeqCst));
21892 assert!(committed_after_drop.load(Ordering::SeqCst));
21893 }
21894
21895 #[tokio::test]
21896 async fn test_buffered_streaming_transition_reservation_falls_back() {
21897 use futures::StreamExt;
21898
21899 let mock = mock_with_responses(vec![
21900 "Serial streaming response",
21901 "Serial streaming response",
21902 ]);
21903 let router_mock = mock_with_response("1");
21904 let router_counter = router_mock.clone();
21905 let yaml = r#"
21906name: BufferedReservationFallbackAgent
21907system_prompt: "Stream normally if speculative routing cannot be evaluated."
21908llm:
21909 default: default
21910 router: router
21911observability:
21912 enabled: true
21913 export:
21914 write_raw_events: true
21915streaming:
21916 enabled: true
21917 buffer_size: 8
21918runtime:
21919 optimization:
21920 enabled: true
21921 max_speculative_llm_calls_per_turn: 1
21922 speculative_state_transitions: true
21923 streaming_policy: buffer_until_routing_done
21924 max_parallel_runtime_tasks: 2
21925states:
21926 initial: triage
21927 states:
21928 triage:
21929 prompt: "Triage state."
21930 transitions:
21931 - to: billing
21932 guard:
21933 context:
21934 route:
21935 eq: billing
21936 when: "User asks about billing"
21937 timing: parallel
21938 billing:
21939 prompt: "Billing state."
21940"#;
21941 let agent = AgentBuilder::from_yaml(yaml)
21942 .unwrap()
21943 .llm_alias("default", Arc::new(mock))
21944 .llm_alias("router", Arc::new(router_mock))
21945 .build()
21946 .unwrap();
21947
21948 let mut stream = agent.chat_stream("hello").await.unwrap();
21949 let mut content = String::new();
21950 let mut error = None;
21951 while let Some(chunk) = stream.next().await {
21952 match chunk {
21953 StreamChunk::Content { text } => content.push_str(&text),
21954 StreamChunk::Error { message } => error = Some(message),
21955 StreamChunk::Done {} => break,
21956 _ => {}
21957 }
21958 }
21959
21960 assert_eq!(error, None);
21961 assert_eq!(content, "Serial streaming response");
21962 assert_eq!(router_counter.call_count(), 0);
21963 let events = agent.observability().unwrap().raw_events();
21964 assert!(events.iter().any(|event| {
21965 event.dimensions.get("branch_status") == Some(&"cancelled".to_string())
21966 && event.dimensions.get("commit_behavior")
21967 == Some(&"transition_decision".to_string())
21968 }));
21969 }
21970
21971 #[tokio::test]
21972 async fn test_blocking_error_cleanup_resets_root_turn_for_next_chat() {
21973 let mut mock = mock_with_response("Recovered response");
21974 mock.set_error("boom");
21975 let mut handle = mock.clone();
21976 let agent = AgentBuilder::new()
21977 .system_prompt("You are helpful.")
21978 .llm(Arc::new(mock))
21979 .build()
21980 .unwrap();
21981
21982 assert!(agent.chat("first").await.is_err());
21983 handle.clear_error();
21984 let response = agent.chat("second").await.unwrap();
21985
21986 assert_eq!(response.content, "Recovered response");
21987 let messages = agent.memory.get_messages(None).await.unwrap();
21988 let user_count = messages
21989 .iter()
21990 .filter(|message| message.role == ai_agents_core::Role::User)
21991 .count();
21992 assert_eq!(user_count, 2);
21993 }
21994
21995 #[tokio::test]
21996 async fn test_streaming_error_cleanup_resets_root_turn_for_next_chat() {
21997 use futures::StreamExt;
21998
21999 let mut mock = mock_with_response("Recovered response");
22000 mock.set_error("stream boom");
22001 let mut handle = mock.clone();
22002 let agent = AgentBuilder::new()
22003 .system_prompt("You are helpful.")
22004 .llm(Arc::new(mock))
22005 .build()
22006 .unwrap();
22007
22008 let mut stream = agent.chat_stream("first").await.unwrap();
22009 let mut saw_error = false;
22010 while let Some(chunk) = stream.next().await {
22011 if matches!(chunk, StreamChunk::Error { .. }) {
22012 saw_error = true;
22013 }
22014 }
22015 assert!(saw_error);
22016
22017 handle.clear_error();
22018 let response = agent.chat("second").await.unwrap();
22019
22020 assert_eq!(response.content, "Recovered response");
22021 let messages = agent.memory.get_messages(None).await.unwrap();
22022 let user_count = messages
22023 .iter()
22024 .filter(|message| message.role == ai_agents_core::Role::User)
22025 .count();
22026 assert_eq!(user_count, 2);
22027 }
22028
22029 #[tokio::test]
22030 async fn test_buffered_streaming_route_miss_releases_buffer_limit() {
22031 use futures::StreamExt;
22032
22033 let mut mock = mock_with_response("one two three");
22034 mock.set_latency(10);
22035 let yaml = r#"
22036name: BufferedMissAgent
22037system_prompt: "You stream safely."
22038llm:
22039 default: default
22040streaming:
22041 enabled: true
22042 buffer_size: 1
22043runtime:
22044 optimization:
22045 enabled: true
22046 max_speculative_llm_calls_per_turn: 2
22047 speculative_state_transitions: true
22048 streaming_policy: buffer_until_routing_done
22049 max_parallel_runtime_tasks: 2
22050states:
22051 initial: triage
22052 states:
22053 triage:
22054 prompt: "Answer from triage."
22055 transitions:
22056 - to: billing
22057 guard:
22058 context:
22059 route:
22060 eq: billing
22061 timing: parallel
22062 billing:
22063 prompt: "Billing state."
22064"#;
22065 let agent = AgentBuilder::from_yaml(yaml)
22066 .unwrap()
22067 .llm_alias("default", Arc::new(mock))
22068 .build()
22069 .unwrap();
22070
22071 let mut stream = agent.chat_stream("hello").await.unwrap();
22072 let mut content = String::new();
22073 let mut error = None;
22074 while let Some(chunk) = stream.next().await {
22075 match chunk {
22076 StreamChunk::Content { text } => content.push_str(&text),
22077 StreamChunk::Error { message } => error = Some(message),
22078 StreamChunk::Done {} => break,
22079 _ => {}
22080 }
22081 }
22082
22083 assert_eq!(error, None);
22084 assert_eq!(content, "one two three");
22085 }
22086
22087 #[tokio::test]
22088 async fn test_buffered_streaming_main_failure_finalizes_branch() {
22089 use futures::StreamExt;
22090
22091 let mock = mock_with_response("one two");
22092 let mut router_mock = mock_with_response("0");
22093 router_mock.set_latency(50);
22094 let yaml = r#"
22095name: BufferedFailureAgent
22096system_prompt: "You stream safely."
22097llm:
22098 default: default
22099 router: router
22100observability:
22101 enabled: true
22102 export:
22103 write_raw_events: true
22104streaming:
22105 enabled: true
22106 buffer_size: 1
22107runtime:
22108 optimization:
22109 enabled: true
22110 max_speculative_llm_calls_per_turn: 2
22111 speculative_state_transitions: true
22112 streaming_policy: buffer_until_routing_done
22113 max_parallel_runtime_tasks: 2
22114states:
22115 initial: triage
22116 states:
22117 triage:
22118 prompt: "Ask for the category."
22119 transitions:
22120 - to: billing
22121 when: "User asks about billing"
22122 timing: parallel
22123 billing:
22124 prompt: "Billing state."
22125"#;
22126 let agent = AgentBuilder::from_yaml(yaml)
22127 .unwrap()
22128 .llm_alias("default", Arc::new(mock))
22129 .llm_alias("router", Arc::new(router_mock))
22130 .build()
22131 .unwrap();
22132
22133 let mut stream = agent.chat_stream("hello").await.unwrap();
22134 let mut error = String::new();
22135 while let Some(chunk) = stream.next().await {
22136 if let StreamChunk::Error { message } = chunk {
22137 error = message;
22138 }
22139 }
22140
22141 assert!(
22142 error.contains("stream buffer filled"),
22143 "unexpected stream error: {}",
22144 error
22145 );
22146 let events = agent.observability().unwrap().raw_events();
22147 assert!(events.iter().any(|event| {
22148 event.dimensions.get("branch_status") == Some(&"failed".to_string())
22149 && event.dimensions.get("commit_behavior") == Some(&"final_response".to_string())
22150 && event.dimensions.get("optimization")
22151 == Some(&"buffered_streaming_routing".to_string())
22152 }));
22153 }
22154
22155 #[tokio::test]
22156 async fn test_streaming_preflight_does_not_emit_old_state_content() {
22157 use futures::StreamExt;
22158
22159 let mock = mock_with_response("Billing streamed response");
22160 let yaml = r#"
22161name: StreamingOptimizedAgent
22162system_prompt: "You route before streaming."
22163runtime:
22164 optimization:
22165 enabled: true
22166 pre_response_deterministic_transitions: true
22167streaming:
22168 enabled: true
22169states:
22170 initial: greeting
22171 states:
22172 greeting:
22173 prompt: "OLD_STATE_SENTINEL"
22174 transitions:
22175 - to: billing
22176 guard:
22177 context:
22178 topic:
22179 eq: billing
22180 timing: pre_response
22181 billing:
22182 prompt: "Billing state."
22183"#;
22184 let agent = AgentBuilder::from_yaml(yaml)
22185 .unwrap()
22186 .llm(Arc::new(mock))
22187 .build()
22188 .unwrap();
22189 agent
22190 .set_context("topic", serde_json::json!("billing"))
22191 .unwrap();
22192
22193 let mut stream = agent.chat_stream("billing please").await.unwrap();
22194 let mut content = String::new();
22195 while let Some(chunk) = stream.next().await {
22196 match chunk {
22197 StreamChunk::Content { text } => content.push_str(&text),
22198 StreamChunk::Error { message } => panic!("stream error: {}", message),
22199 StreamChunk::Done {} => break,
22200 _ => {}
22201 }
22202 }
22203
22204 assert_eq!(agent.current_state().as_deref(), Some("billing"));
22205 assert!(content.contains("Billing streamed response"));
22206 assert!(!content.contains("OLD_STATE_SENTINEL"));
22207 }
22208
22209 #[tokio::test]
22211 async fn test_integration_state_machine_basic() {
22212 let yaml = r#"
22213name: StateAgent
22214system_prompt: "You are a support agent."
22215states:
22216 initial: greeting
22217 states:
22218 greeting:
22219 prompt: "Welcome the user warmly."
22220 transitions:
22221 - to: support
22222 when: "User needs help"
22223 auto: true
22224 support:
22225 prompt: "Help solve the user's problem."
22226"#;
22227 let mock = mock_with_responses(vec![
22228 "Welcome! How can I help?", "1", "I'll help you with that.", ]);
22232 let builder = AgentBuilder::from_yaml(yaml).unwrap();
22233 let agent = builder.llm(Arc::new(mock)).build().unwrap();
22234
22235 assert_eq!(agent.current_state(), Some("greeting".to_string()));
22236 let _ = agent.chat("I need help").await.unwrap();
22237 }
22240
22241 #[tokio::test]
22243 async fn test_integration_state_on_enter_set_context() {
22244 let yaml = r#"
22245name: ActionAgent
22246system_prompt: "You are helpful."
22247states:
22248 initial: step1
22249 states:
22250 step1:
22251 prompt: "Step 1"
22252 on_exit:
22253 - set_context:
22254 step1_exited: true
22255 transitions:
22256 - to: step2
22257 when: "always"
22258 auto: true
22259 step2:
22260 prompt: "Step 2"
22261 on_enter:
22262 - set_context:
22263 step2_entered: true
22264"#;
22265 let mock = mock_with_responses(vec![
22267 "Processing step 1.",
22268 "0", ]);
22270 let builder = AgentBuilder::from_yaml(yaml).unwrap();
22271 let agent = builder.llm(Arc::new(mock)).build().unwrap();
22272
22273 assert_eq!(agent.current_state(), Some("step1".to_string()));
22274
22275 agent.transition_to("step2").await.unwrap();
22277
22278 assert_eq!(agent.current_state(), Some("step2".to_string()));
22279
22280 let ctx = agent.get_context();
22282 assert_eq!(ctx.get("step1_exited"), Some(&serde_json::json!(true)));
22283 assert_eq!(ctx.get("step2_entered"), Some(&serde_json::json!(true)));
22284 }
22285
22286 #[tokio::test]
22287 async fn state_action_tool_preserves_source_in_stored_record() {
22288 let yaml = r#"
22289name: StateActionToolAgent
22290system_prompt: "You are helpful."
22291tools:
22292 - context_echo
22293states:
22294 initial: idle
22295 states:
22296 idle:
22297 prompt: "Idle"
22298 active:
22299 prompt: "Active"
22300 on_enter:
22301 - set_context:
22302 action_started: true
22303 - tool: context_echo
22304 args: {}
22305"#;
22306 let agent = AgentBuilder::from_yaml(yaml)
22307 .unwrap()
22308 .llm(Arc::new(mock_with_response("unused")))
22309 .tool(Arc::new(ContextEchoTool))
22310 .build()
22311 .unwrap();
22312
22313 agent.transition_to("active").await.unwrap();
22314
22315 let record: ToolExecutionRecord = serde_json::from_value(
22316 agent
22317 .get_context()
22318 .get("last_tool_record")
22319 .cloned()
22320 .expect("successful state action must store its execution record"),
22321 )
22322 .unwrap();
22323 assert!(record.executed);
22324 assert!(record.success);
22325 assert_eq!(record.canonical_id, "context_echo");
22326 assert!(matches!(
22327 &record.source,
22328 ToolCallSource::StateAction {
22329 state: Some(state),
22330 action_index: 1,
22331 } if state == "active"
22332 ));
22333 }
22334
22335 #[tokio::test]
22336 async fn test_ordinary_transition_uses_on_enter_then_on_reenter() {
22337 let yaml = r#"
22338name: OrdinaryLifecycleAgent
22339system_prompt: "You are helpful."
22340states:
22341 initial: intake
22342 regenerate_on_transition: false
22343 states:
22344 intake:
22345 prompt: "Intake"
22346 transitions:
22347 - to: drafting
22348 guard:
22349 context:
22350 route:
22351 eq: drafting
22352 drafting:
22353 prompt: "Drafting"
22354 on_enter:
22355 - set_context:
22356 draft_version: 1
22357 on_reenter:
22358 - set_context:
22359 draft_version: 2
22360 transitions:
22361 - to: review
22362 guard:
22363 context:
22364 route:
22365 eq: review
22366 review:
22367 prompt: "Review"
22368 on_enter:
22369 - set_context:
22370 review_entry: first
22371 transitions:
22372 - to: drafting
22373 guard:
22374 context:
22375 route:
22376 eq: drafting
22377"#;
22378 let agent = AgentBuilder::from_yaml(yaml)
22379 .unwrap()
22380 .llm(Arc::new(mock_with_responses(vec![
22381 "Intake response",
22382 "Draft response",
22383 "Review response",
22384 ])))
22385 .build()
22386 .unwrap();
22387
22388 agent
22389 .set_context("route", serde_json::json!("drafting"))
22390 .unwrap();
22391 agent.chat("Start a draft").await.unwrap();
22392 assert_eq!(agent.current_state().as_deref(), Some("drafting"));
22393 assert_eq!(
22394 agent.get_context().get("draft_version"),
22395 Some(&serde_json::json!(1))
22396 );
22397
22398 agent
22399 .set_context("route", serde_json::json!("review"))
22400 .unwrap();
22401 agent.chat("Review this").await.unwrap();
22402 assert_eq!(agent.current_state().as_deref(), Some("review"));
22403 assert_eq!(
22404 agent.get_context().get("review_entry"),
22405 Some(&serde_json::json!("first"))
22406 );
22407
22408 agent
22409 .set_context("route", serde_json::json!("drafting"))
22410 .unwrap();
22411 agent.chat("Revise this").await.unwrap();
22412 assert_eq!(agent.current_state().as_deref(), Some("drafting"));
22413 assert_eq!(
22414 agent.get_context().get("draft_version"),
22415 Some(&serde_json::json!(2))
22416 );
22417 }
22418
22419 #[tokio::test]
22420 async fn test_manual_transition_uses_on_enter_then_on_reenter() {
22421 let yaml = r#"
22422name: ManualLifecycleAgent
22423system_prompt: "You are helpful."
22424states:
22425 initial: intake
22426 states:
22427 intake:
22428 prompt: "Intake"
22429 drafting:
22430 prompt: "Drafting"
22431 on_enter:
22432 - set_context:
22433 draft_version: 1
22434 on_reenter:
22435 - set_context:
22436 draft_version: 2
22437 review:
22438 prompt: "Review"
22439"#;
22440 let agent = AgentBuilder::from_yaml(yaml)
22441 .unwrap()
22442 .llm(Arc::new(mock_with_response("unused")))
22443 .build()
22444 .unwrap();
22445
22446 assert!(!agent.get_context().contains_key("draft_version"));
22447 agent.transition_to("drafting").await.unwrap();
22448 assert_eq!(agent.current_state().as_deref(), Some("drafting"));
22449 assert_eq!(
22450 agent.get_context().get("draft_version"),
22451 Some(&serde_json::json!(1))
22452 );
22453
22454 agent.transition_to("review").await.unwrap();
22455 agent.transition_to("drafting").await.unwrap();
22456 assert_eq!(agent.current_state().as_deref(), Some("drafting"));
22457 assert_eq!(
22458 agent.get_context().get("draft_version"),
22459 Some(&serde_json::json!(2))
22460 );
22461 }
22462
22463 #[tokio::test]
22464 async fn test_timeout_transition_uses_on_enter_then_on_reenter() {
22465 let yaml = r#"
22466name: TimeoutLifecycleAgent
22467system_prompt: "You are helpful."
22468states:
22469 initial: intake
22470 regenerate_on_transition: false
22471 states:
22472 intake:
22473 prompt: "Intake"
22474 max_turns: 1
22475 timeout_to: drafting
22476 drafting:
22477 prompt: "Drafting"
22478 max_turns: 1
22479 timeout_to: review
22480 on_enter:
22481 - set_context:
22482 draft_version: 1
22483 on_reenter:
22484 - set_context:
22485 draft_version: 2
22486 review:
22487 prompt: "Review"
22488 max_turns: 1
22489 timeout_to: drafting
22490 on_enter:
22491 - set_context:
22492 review_entry: first
22493"#;
22494 let agent = AgentBuilder::from_yaml(yaml)
22495 .unwrap()
22496 .llm(Arc::new(mock_with_responses(vec![
22497 "Intake",
22498 "First draft",
22499 "Review",
22500 "Revised draft",
22501 ])))
22502 .build()
22503 .unwrap();
22504
22505 agent.chat("First turn").await.unwrap();
22506 assert_eq!(agent.current_state().as_deref(), Some("intake"));
22507 assert!(!agent.get_context().contains_key("draft_version"));
22508
22509 agent.chat("Second turn").await.unwrap();
22510 assert_eq!(agent.current_state().as_deref(), Some("drafting"));
22511 assert_eq!(
22512 agent.get_context().get("draft_version"),
22513 Some(&serde_json::json!(1))
22514 );
22515
22516 agent.chat("Third turn").await.unwrap();
22517 assert_eq!(agent.current_state().as_deref(), Some("review"));
22518 assert_eq!(
22519 agent.get_context().get("review_entry"),
22520 Some(&serde_json::json!("first"))
22521 );
22522
22523 agent.chat("Fourth turn").await.unwrap();
22524 assert_eq!(agent.current_state().as_deref(), Some("drafting"));
22525 assert_eq!(
22526 agent.get_context().get("draft_version"),
22527 Some(&serde_json::json!(2))
22528 );
22529 }
22530
22531 #[tokio::test]
22533 async fn test_integration_process_normalize() {
22534 let yaml = r#"
22535name: ProcessAgent
22536system_prompt: "You are helpful."
22537process:
22538 input:
22539 - type: normalize
22540 config:
22541 trim: true
22542 collapse_whitespace: true
22543"#;
22544 let mock = mock_with_response("Got your message.");
22545 let builder = AgentBuilder::from_yaml(yaml).unwrap();
22546 let agent = builder.llm(Arc::new(mock.clone())).build().unwrap();
22547
22548 let _ = agent.chat(" hello world ").await.unwrap();
22549
22550 let history = mock.call_history();
22552 assert!(!history.is_empty());
22553 let last_call = history.last().unwrap();
22555 let user_msg = last_call
22556 .messages
22557 .iter()
22558 .find(|m| m.role == ai_agents_core::Role::User)
22559 .unwrap();
22560 assert_eq!(user_msg.content, "hello world");
22561 }
22562
22563 #[tokio::test]
22567 async fn test_integration_memory_compression() {
22568 let yaml = r#"
22569name: MemoryAgent
22570system_prompt: "You are helpful."
22571memory:
22572 type: compacting
22573 max_messages: 100
22574 compress_threshold: 5
22575 max_recent_messages: 3
22576 summarize_batch_size: 2
22577"#;
22578 let responses: Vec<&str> = (0..8).map(|_| "Response from assistant.").collect();
22580 let mock = mock_with_responses(responses);
22581 let builder = AgentBuilder::from_yaml(yaml).unwrap();
22582 let agent = builder.llm(Arc::new(mock)).build().unwrap();
22583
22584 for i in 0..6 {
22586 let _ = agent.chat(&format!("Message {}", i)).await.unwrap();
22587 }
22588
22589 let messages = agent.memory.get_messages(None).await.unwrap();
22592 assert!(messages.len() <= 12); }
22596
22597 #[tokio::test]
22599 async fn test_integration_multi_llm_registry() {
22600 let mut mock_default = MockLLMProvider::new("default");
22601 mock_default.set_response("Default LLM response.");
22602 let mut mock_router = MockLLMProvider::new("router");
22603 mock_router.set_response("Router response.");
22604
22605 let agent = AgentBuilder::new()
22606 .system_prompt("You are helpful.")
22607 .llm_alias("default", Arc::new(mock_default))
22608 .llm_alias("router", Arc::new(mock_router))
22609 .build()
22610 .unwrap();
22611
22612 let response = agent.chat("Hello").await.unwrap();
22613 assert_eq!(response.content, "Default LLM response.");
22614 }
22615
22616 #[tokio::test]
22618 async fn test_integration_agent_reset() {
22619 let mock = mock_with_responses(vec!["Hello!", "Hello again!"]);
22620 let agent = AgentBuilder::new()
22621 .system_prompt("You are helpful.")
22622 .llm(Arc::new(mock))
22623 .build()
22624 .unwrap();
22625
22626 let _ = agent.chat("Hi").await.unwrap();
22627 let messages = agent.memory.get_messages(None).await.unwrap();
22628 assert_eq!(messages.len(), 2); agent.reset().await.unwrap();
22631 let messages = agent.memory.get_messages(None).await.unwrap();
22632 assert_eq!(messages.len(), 0);
22633 }
22634
22635 #[tokio::test]
22637 async fn test_integration_process_validate_reject() {
22638 use ai_agents_process::{ProcessConfig, ProcessProcessor};
22639
22640 let validate_config = ai_agents_process::ValidateStage {
22641 id: Some("length_check".to_string()),
22642 condition: None,
22643 config: ai_agents_process::ValidateConfig {
22644 rules: vec![ai_agents_process::ValidationRule::MinLength {
22645 min_length: 10,
22646 on_fail: ai_agents_process::ValidationAction {
22647 action: ai_agents_process::ValidationActionType::Reject,
22648 message: None,
22649 },
22650 }],
22651 ..Default::default()
22652 },
22653 };
22654 let process_config = ProcessConfig {
22655 input: vec![ai_agents_process::ProcessStage::Validate(validate_config)],
22656 ..Default::default()
22657 };
22658 let processor = ProcessProcessor::new(process_config);
22659
22660 let mock = mock_with_response("Should not reach here.");
22661 let agent = AgentBuilder::new()
22662 .system_prompt("You are helpful.")
22663 .llm(Arc::new(mock))
22664 .process_processor(processor)
22665 .build()
22666 .unwrap();
22667
22668 let response = agent.chat("Hi").await.unwrap();
22669 assert!(
22671 response.content.contains("rejected")
22672 || response.content.contains("Input rejected")
22673 || response.content.contains("too short")
22674 || response.content.contains("Too short")
22675 || response.content.len() < 50, "Expected rejection response, got: {}",
22677 response.content
22678 );
22679 }
22680
22681 #[tokio::test]
22683 async fn test_llm_fallback_on_failure() {
22684 use ai_agents_recovery::{ErrorRecoveryConfig, LLMFailureAction, LLMRecoveryConfig};
22685
22686 let mut primary = MockLLMProvider::new("primary");
22687 primary.set_error("Primary LLM is unavailable");
22688
22689 let mut fallback = MockLLMProvider::new("fallback");
22690 fallback.set_response("Fallback response works!");
22691
22692 let agent = AgentBuilder::new()
22693 .system_prompt("You are helpful.")
22694 .llm_alias("default", Arc::new(primary))
22695 .llm_alias("backup", Arc::new(fallback))
22696 .recovery_manager(RecoveryManager::new(ErrorRecoveryConfig {
22697 llm: LLMRecoveryConfig {
22698 on_failure: LLMFailureAction::FallbackLlm {
22699 fallback_llm: "backup".to_string(),
22700 },
22701 ..Default::default()
22702 },
22703 ..Default::default()
22704 }))
22705 .build()
22706 .unwrap();
22707
22708 let response = agent.chat("Hello").await.unwrap();
22709 assert!(
22710 response.content.contains("Fallback response"),
22711 "Expected fallback response, got: {}",
22712 response.content
22713 );
22714 }
22715
22716 #[tokio::test]
22718 async fn test_llm_fallback_response_static_message() {
22719 use ai_agents_recovery::{ErrorRecoveryConfig, LLMFailureAction, LLMRecoveryConfig};
22720
22721 let mut primary = MockLLMProvider::new("primary");
22722 primary.set_error("Primary LLM is unavailable");
22723
22724 let agent = AgentBuilder::new()
22725 .system_prompt("You are helpful.")
22726 .llm(Arc::new(primary))
22727 .recovery_manager(RecoveryManager::new(ErrorRecoveryConfig {
22728 llm: LLMRecoveryConfig {
22729 on_failure: LLMFailureAction::FallbackResponse {
22730 message: "I am temporarily unavailable. Please try again later."
22731 .to_string(),
22732 },
22733 ..Default::default()
22734 },
22735 ..Default::default()
22736 }))
22737 .build()
22738 .unwrap();
22739
22740 let response = agent.chat("Hello").await.unwrap();
22741 assert!(
22742 response.content.contains("temporarily unavailable"),
22743 "Expected static fallback message, got: {}",
22744 response.content
22745 );
22746 }
22747
22748 #[tokio::test]
22751 async fn test_tool_failure_skip() {
22752 use ai_agents_recovery::{
22753 ErrorRecoveryConfig, ToolFailureAction, ToolRecoveryConfig, ToolRetryConfig,
22754 };
22755
22756 let mock = mock_with_responses(vec![
22757 r#"{"tool": "calculator", "arguments": {"expression": "not a number +"}}"#,
22758 "The calculation was skipped, but I can still help you.",
22759 ]);
22760 let observed = mock.clone();
22761 let mut tools = ai_agents_tools::ToolRegistry::new();
22762 tools
22763 .register(Arc::new(ai_agents_tools::CalculatorTool))
22764 .unwrap();
22765
22766 let agent = AgentBuilder::new()
22767 .system_prompt("You are helpful.")
22768 .llm(Arc::new(mock))
22769 .tools(tools)
22770 .recovery_manager(RecoveryManager::new(ErrorRecoveryConfig {
22771 tools: ToolRecoveryConfig {
22772 default: ToolRetryConfig {
22773 max_retries: 0,
22774 timeout_ms: None,
22775 on_failure: ToolFailureAction::Skip,
22776 },
22777 ..Default::default()
22778 },
22779 ..Default::default()
22780 }))
22781 .build()
22782 .unwrap();
22783
22784 let response = agent.chat("Compute this").await.unwrap();
22785
22786 assert_eq!(
22787 response.content,
22788 "The calculation was skipped, but I can still help you."
22789 );
22790 assert_eq!(observed.call_count(), 2);
22791 let history = agent.tool_call_history();
22793 assert_eq!(history.len(), 1);
22794 assert_eq!(history[0].tool_id, "calculator");
22795 assert_eq!(
22796 history[0].result.get("skipped"),
22797 Some(&serde_json::json!(true)),
22798 "{:?}",
22799 history[0].result
22800 );
22801 }
22802
22803 #[tokio::test]
22805 async fn test_unregistered_tool_call_records_unavailable_and_continues() {
22806 let mock = mock_with_responses(vec![
22807 r#"{"tool": "nonexistent_tool", "arguments": {}}"#,
22808 "The tool was unavailable, but I can still help you.",
22809 ]);
22810 let observed = mock.clone();
22811
22812 let agent = AgentBuilder::new()
22813 .system_prompt("You are helpful.")
22814 .llm(Arc::new(mock))
22815 .build()
22816 .unwrap();
22817
22818 let response = agent.chat("Use the nonexistent tool").await.unwrap();
22819
22820 assert_eq!(
22821 response.content,
22822 "The tool was unavailable, but I can still help you."
22823 );
22824 assert_eq!(observed.call_count(), 2);
22825 let history = agent.tool_call_history();
22826 assert_eq!(history.len(), 1);
22827 assert_eq!(history[0].tool_id, "nonexistent_tool");
22828 assert_eq!(
22829 history[0].result.pointer("/error/kind"),
22830 Some(&serde_json::json!("tool_unavailable")),
22831 "{:?}",
22832 history[0].result
22833 );
22834 }
22835
22836 fn fallback_llm_recovery(fallback_llm: &str) -> RecoveryManager {
22841 use ai_agents_recovery::{ErrorRecoveryConfig, LLMFailureAction, LLMRecoveryConfig};
22842 RecoveryManager::new(ErrorRecoveryConfig {
22843 llm: LLMRecoveryConfig {
22844 on_failure: LLMFailureAction::FallbackLlm {
22845 fallback_llm: fallback_llm.to_string(),
22846 },
22847 ..Default::default()
22848 },
22849 ..Default::default()
22850 })
22851 }
22852
22853 #[tokio::test]
22854 async fn test_stream_llm_fallback_on_open_failure() {
22855 let mut primary = MockLLMProvider::new("primary");
22856 primary.set_error("Primary LLM is unavailable");
22857 let mut fallback = MockLLMProvider::new("fallback");
22858 fallback.set_response("Fallback response works!");
22859
22860 let agent = AgentBuilder::new()
22861 .system_prompt("You are helpful.")
22862 .llm_alias("default", Arc::new(primary))
22863 .llm_alias("backup", Arc::new(fallback))
22864 .recovery_manager(fallback_llm_recovery("backup"))
22865 .build()
22866 .unwrap();
22867
22868 let (content, chunks, final_response) = collect_stream_events(&agent, "Hello").await;
22869 assert!(
22870 !chunks.iter().any(StreamChunk::is_error),
22871 "fallback must not surface as a stream error: {chunks:?}"
22872 );
22873 let final_response = final_response.expect("Final must be emitted after fallback");
22874 assert!(content.contains("Fallback response"));
22875 assert!(final_response.content.contains("Fallback response"));
22876 }
22877
22878 #[tokio::test]
22879 async fn test_stream_llm_fallback_response_static_message() {
22880 use ai_agents_recovery::{ErrorRecoveryConfig, LLMFailureAction, LLMRecoveryConfig};
22881
22882 let mut primary = MockLLMProvider::new("primary");
22883 primary.set_error("Primary LLM is unavailable");
22884
22885 let agent = AgentBuilder::new()
22886 .system_prompt("You are helpful.")
22887 .llm(Arc::new(primary))
22888 .recovery_manager(RecoveryManager::new(ErrorRecoveryConfig {
22889 llm: LLMRecoveryConfig {
22890 on_failure: LLMFailureAction::FallbackResponse {
22891 message: "Service is temporarily unavailable.".to_string(),
22892 },
22893 ..Default::default()
22894 },
22895 ..Default::default()
22896 }))
22897 .build()
22898 .unwrap();
22899
22900 let (content, chunks, final_response) = collect_stream_events(&agent, "Hello").await;
22901 assert!(!chunks.iter().any(StreamChunk::is_error));
22902 let content_chunks = chunks.iter().filter(|c| c.is_content()).count();
22903 assert_eq!(content_chunks, 1, "static fallback is one content chunk");
22904 assert_eq!(content, "Service is temporarily unavailable.");
22905 assert_eq!(
22906 final_response.expect("Final").content,
22907 "Service is temporarily unavailable."
22908 );
22909 }
22910
22911 struct FailOnceStreamProvider {
22913 remaining_failures: Arc<std::sync::atomic::AtomicUsize>,
22914 open_attempts: Arc<std::sync::atomic::AtomicUsize>,
22915 }
22916
22917 #[async_trait]
22918 impl LLMProvider for FailOnceStreamProvider {
22919 async fn complete(
22920 &self,
22921 _messages: &[ChatMessage],
22922 _config: Option<&LLMConfig>,
22923 ) -> std::result::Result<LLMResponse, LLMError> {
22924 Ok(LLMResponse::new("blocking path", FinishReason::Stop))
22925 }
22926
22927 async fn complete_stream(
22928 &self,
22929 _messages: &[ChatMessage],
22930 _config: Option<&LLMConfig>,
22931 ) -> std::result::Result<
22932 Box<dyn Stream<Item = std::result::Result<LLMChunk, LLMError>> + Unpin + Send>,
22933 LLMError,
22934 > {
22935 self.open_attempts.fetch_add(1, Ordering::SeqCst);
22936 if self
22937 .remaining_failures
22938 .fetch_update(Ordering::SeqCst, Ordering::SeqCst, |n| n.checked_sub(1))
22939 .is_ok()
22940 {
22941 return Err(LLMError::Network("connection reset".to_string()));
22942 }
22943 Ok(Box::new(futures::stream::iter(vec![Ok(LLMChunk::new(
22944 "Recovered after retry",
22945 true,
22946 ))])))
22947 }
22948
22949 fn provider_name(&self) -> &str {
22950 "fail-once-stream"
22951 }
22952
22953 fn supports(&self, feature: LLMFeature) -> bool {
22954 matches!(feature, LLMFeature::Streaming)
22955 }
22956 }
22957
22958 #[tokio::test]
22959 async fn test_stream_llm_retry_then_success() {
22960 use ai_agents_recovery::{BackoffConfig, ErrorRecoveryConfig, RetryConfig};
22961
22962 let open_attempts = Arc::new(std::sync::atomic::AtomicUsize::new(0));
22963 let provider = FailOnceStreamProvider {
22964 remaining_failures: Arc::new(std::sync::atomic::AtomicUsize::new(1)),
22965 open_attempts: Arc::clone(&open_attempts),
22966 };
22967
22968 let agent = AgentBuilder::new()
22969 .system_prompt("You are helpful.")
22970 .llm(Arc::new(provider))
22971 .recovery_manager(RecoveryManager::new(ErrorRecoveryConfig {
22972 default: RetryConfig {
22973 max_retries: 1,
22974 backoff: BackoffConfig {
22975 initial_ms: 1,
22976 max_ms: 1,
22977 ..Default::default()
22978 },
22979 ..Default::default()
22980 },
22981 ..Default::default()
22982 }))
22983 .build()
22984 .unwrap();
22985
22986 let (content, chunks, final_response) = collect_stream_events(&agent, "Hello").await;
22987 assert!(!chunks.iter().any(StreamChunk::is_error), "{chunks:?}");
22988 assert_eq!(open_attempts.load(Ordering::SeqCst), 2);
22989 assert_eq!(content, "Recovered after retry");
22990 assert_eq!(
22991 final_response.expect("Final").content,
22992 "Recovered after retry"
22993 );
22994 }
22995
22996 #[tokio::test]
22997 async fn test_stream_llm_error_action_error_emits_terminal_error() {
22998 let mut primary = MockLLMProvider::new("primary");
22999 primary.set_error("Primary LLM is unavailable");
23000
23001 let agent = AgentBuilder::new()
23002 .system_prompt("You are helpful.")
23003 .llm(Arc::new(primary))
23004 .build()
23005 .unwrap();
23006
23007 let (_, chunks, final_response) = collect_stream_events(&agent, "Hello").await;
23008 assert!(
23009 final_response.is_none(),
23010 "default Error action must not produce Final"
23011 );
23012 assert!(
23013 chunks.iter().any(StreamChunk::is_error),
23014 "default Error action must surface a stream error"
23015 );
23016 }
23017
23018 struct MidStreamFailureProvider;
23020
23021 #[async_trait]
23022 impl LLMProvider for MidStreamFailureProvider {
23023 async fn complete(
23024 &self,
23025 _messages: &[ChatMessage],
23026 _config: Option<&LLMConfig>,
23027 ) -> std::result::Result<LLMResponse, LLMError> {
23028 Ok(LLMResponse::new("blocking path", FinishReason::Stop))
23029 }
23030
23031 async fn complete_stream(
23032 &self,
23033 _messages: &[ChatMessage],
23034 _config: Option<&LLMConfig>,
23035 ) -> std::result::Result<
23036 Box<dyn Stream<Item = std::result::Result<LLMChunk, LLMError>> + Unpin + Send>,
23037 LLMError,
23038 > {
23039 Ok(Box::new(futures::stream::iter(vec![
23040 Ok(LLMChunk::new("Partial ", false)),
23041 Err(LLMError::Network("connection dropped".to_string())),
23042 ])))
23043 }
23044
23045 fn provider_name(&self) -> &str {
23046 "mid-stream-failure"
23047 }
23048
23049 fn supports(&self, feature: LLMFeature) -> bool {
23050 matches!(feature, LLMFeature::Streaming)
23051 }
23052 }
23053
23054 #[tokio::test]
23055 async fn test_stream_mid_stream_failure_is_terminal() {
23056 let mut fallback = MockLLMProvider::new("fallback");
23057 fallback.set_response("Fallback must not run");
23058 let fallback_calls = fallback.clone();
23059
23060 let agent = AgentBuilder::new()
23061 .system_prompt("You are helpful.")
23062 .llm_alias("default", Arc::new(MidStreamFailureProvider))
23063 .llm_alias("backup", Arc::new(fallback))
23064 .recovery_manager(fallback_llm_recovery("backup"))
23065 .build()
23066 .unwrap();
23067
23068 let (content, chunks, final_response) = collect_stream_events(&agent, "Hello").await;
23069 assert_eq!(content, "Partial ");
23070 assert!(chunks.iter().any(StreamChunk::is_error));
23071 assert!(final_response.is_none());
23072 assert_eq!(
23073 fallback_calls.call_count(),
23074 0,
23075 "fallback must not run after a visible delta"
23076 );
23077 }
23078
23079 #[tokio::test]
23080 async fn test_buffered_streaming_draft_uses_fallback_llm() {
23081 use futures::StreamExt;
23082
23083 let mut primary = MockLLMProvider::new("primary");
23084 primary.set_error("Primary LLM is unavailable");
23085 let fallback = mock_with_response("fallback one two");
23086 let yaml = r#"
23087name: BufferedFallbackAgent
23088system_prompt: "You stream safely."
23089llm:
23090 default: default
23091streaming:
23092 enabled: true
23093 buffer_size: 8
23094runtime:
23095 optimization:
23096 enabled: true
23097 max_speculative_llm_calls_per_turn: 2
23098 speculative_state_transitions: true
23099 streaming_policy: buffer_until_routing_done
23100 max_parallel_runtime_tasks: 2
23101states:
23102 initial: triage
23103 states:
23104 triage:
23105 prompt: "Answer from triage."
23106 transitions:
23107 - to: billing
23108 guard:
23109 context:
23110 route:
23111 eq: billing
23112 timing: parallel
23113 billing:
23114 prompt: "Billing state."
23115"#;
23116 let agent = AgentBuilder::from_yaml(yaml)
23117 .unwrap()
23118 .llm_alias("default", Arc::new(primary))
23119 .llm_alias("backup", Arc::new(fallback))
23120 .recovery_manager(fallback_llm_recovery("backup"))
23121 .build()
23122 .unwrap();
23123
23124 let mut stream = agent.chat_stream("hello").await.unwrap();
23125 let mut content = String::new();
23126 let mut error = None;
23127 while let Some(chunk) = stream.next().await {
23128 match chunk {
23129 StreamChunk::Content { text } => content.push_str(&text),
23130 StreamChunk::Error { message } => error = Some(message),
23131 StreamChunk::Done {} => break,
23132 _ => {}
23133 }
23134 }
23135
23136 assert_eq!(error, None);
23137 assert_eq!(content, "fallback one two");
23138 }
23139
23140 #[tokio::test]
23141 async fn parity_llm_fallback_llm() {
23142 let build = || {
23143 let mut primary = MockLLMProvider::new("primary");
23144 primary.set_error("Primary LLM is unavailable");
23145 let mut fallback = MockLLMProvider::new("fallback");
23146 fallback.set_response("Fallback response works!");
23147 AgentBuilder::new()
23148 .system_prompt("You are helpful.")
23149 .llm_alias("default", Arc::new(primary))
23150 .llm_alias("backup", Arc::new(fallback))
23151 .recovery_manager(fallback_llm_recovery("backup"))
23152 .build()
23153 .unwrap()
23154 };
23155 assert_blocking_streaming_parity(build, "Hello").await;
23156 }
23157
23158 #[tokio::test]
23159 async fn parity_llm_fallback_response() {
23160 use ai_agents_recovery::{ErrorRecoveryConfig, LLMFailureAction, LLMRecoveryConfig};
23161 let build = || {
23162 let mut primary = MockLLMProvider::new("primary");
23163 primary.set_error("Primary LLM is unavailable");
23164 AgentBuilder::new()
23165 .system_prompt("You are helpful.")
23166 .llm(Arc::new(primary))
23167 .recovery_manager(RecoveryManager::new(ErrorRecoveryConfig {
23168 llm: LLMRecoveryConfig {
23169 on_failure: LLMFailureAction::FallbackResponse {
23170 message: "Service is temporarily unavailable.".to_string(),
23171 },
23172 ..Default::default()
23173 },
23174 ..Default::default()
23175 }))
23176 .build()
23177 .unwrap()
23178 };
23179 assert_blocking_streaming_parity(build, "Hello").await;
23180 }
23181
23182 #[tokio::test]
23183 async fn parity_basic_chat() {
23184 let build = || {
23185 AgentBuilder::new()
23186 .system_prompt("You are helpful.")
23187 .llm(Arc::new(mock_with_response("Plain answer")))
23188 .build()
23189 .unwrap()
23190 };
23191 assert_blocking_streaming_parity(build, "Hello").await;
23192 }
23193
23194 fn skills_with_parallel_transition_yaml(extra_optimization: &str, streaming: &str) -> String {
23200 format!(
23201 r#"
23202name: SkillsBesideTransitionAgent
23203system_prompt: "Use skills when they match."
23204llm:
23205 default: default
23206 router: router
23207observability:
23208 enabled: true
23209 export:
23210 write_raw_events: true
23211{streaming}
23212runtime:
23213 optimization:
23214 enabled: true
23215 speculative_state_transitions: true
23216{extra_optimization}
23217states:
23218 initial: triage
23219 states:
23220 triage:
23221 prompt: "Triage state."
23222 transitions:
23223 - to: billing
23224 guard:
23225 context:
23226 route:
23227 eq: billing
23228 timing: parallel
23229 billing:
23230 prompt: "Billing state."
23231skills:
23232 - id: helper
23233 description: "Answer helper requests"
23234 trigger: "User asks for helper"
23235 steps:
23236 - prompt: "Answer the helper request: {{{{ user_input }}}}"
23237 llm: skill
23238"#
23239 )
23240 }
23241
23242 struct RoleMocks {
23246 main: MockLLMProvider,
23247 router: MockLLMProvider,
23248 skill: MockLLMProvider,
23249 }
23250
23251 fn role_mocks(main: MockLLMProvider, router: MockLLMProvider) -> RoleMocks {
23252 RoleMocks {
23253 main,
23254 router,
23255 skill: mock_with_response("Skill step response"),
23256 }
23257 }
23258
23259 fn build_skills_beside_transition_agent(yaml: &str, mocks: RoleMocks) -> RuntimeAgent {
23260 AgentBuilder::from_yaml(yaml)
23261 .unwrap()
23262 .llm_alias("default", Arc::new(mocks.main))
23263 .llm_alias("router", Arc::new(mocks.router))
23264 .llm_alias("skill", Arc::new(mocks.skill))
23265 .build()
23266 .unwrap()
23267 }
23268
23269 fn branch_events_with_commit_behavior(agent: &RuntimeAgent, behavior: &str) -> usize {
23270 agent
23271 .observability()
23272 .unwrap()
23273 .raw_events()
23274 .iter()
23275 .filter(|event| event.dimensions.get("commit_behavior") == Some(&behavior.to_string()))
23276 .count()
23277 }
23278
23279 #[tokio::test]
23280 async fn test_speculative_transition_with_skills_and_no_skill_branch_routes_skill_serially() {
23281 let default_mock = mock_with_response("Draft response");
23282 let router_mock = mock_with_response("helper");
23283 let router_counter = router_mock.clone();
23284 let yaml = skills_with_parallel_transition_yaml(
23285 " max_speculative_llm_calls_per_turn: 2\n max_parallel_runtime_tasks: 2",
23286 "",
23287 );
23288 let agent =
23289 build_skills_beside_transition_agent(&yaml, role_mocks(default_mock, router_mock));
23290
23291 let response = agent.chat("please use helper").await.unwrap();
23292
23293 assert_eq!(
23294 response.metadata.as_ref().and_then(|m| m.get("skill_id")),
23295 Some(&serde_json::json!("helper")),
23296 "skill must route even without a skill branch: {response:?}"
23297 );
23298 assert_eq!(router_counter.call_count(), 1);
23299 assert!(branch_events_with_commit_behavior(&agent, "transition_decision") > 0);
23301 assert_eq!(
23302 branch_events_with_commit_behavior(&agent, "skill_selection"),
23303 0
23304 );
23305 }
23306
23307 #[tokio::test]
23308 async fn test_speculative_transition_with_skills_no_match_commits_draft() {
23309 let default_mock = mock_with_response("Draft response");
23310 let router_mock = mock_with_response("none");
23311 let router_counter = router_mock.clone();
23312 let yaml = skills_with_parallel_transition_yaml(
23313 " max_speculative_llm_calls_per_turn: 2\n max_parallel_runtime_tasks: 2",
23314 "",
23315 );
23316 let agent =
23317 build_skills_beside_transition_agent(&yaml, role_mocks(default_mock, router_mock));
23318
23319 let response = agent.chat("just chat").await.unwrap();
23320
23321 assert_eq!(response.content, "Draft response");
23322 assert!(
23323 response
23324 .metadata
23325 .as_ref()
23326 .is_none_or(|m| !m.contains_key("skill_id"))
23327 );
23328 assert_eq!(router_counter.call_count(), 1);
23329 assert!(branch_events_with_commit_behavior(&agent, "final_response") > 0);
23330 }
23331
23332 #[tokio::test]
23333 async fn test_speculative_transition_win_skips_serial_skill_selection() {
23334 let default_mock = mock_with_response("Billing answer");
23335 let router_mock = mock_with_response("none");
23336 let router_counter = router_mock.clone();
23337 let yaml = skills_with_parallel_transition_yaml(
23338 " max_speculative_llm_calls_per_turn: 2\n max_parallel_runtime_tasks: 2",
23339 "",
23340 );
23341 let agent =
23342 build_skills_beside_transition_agent(&yaml, role_mocks(default_mock, router_mock));
23343 agent
23344 .set_context("route", serde_json::json!("billing"))
23345 .unwrap();
23346
23347 let response = agent.chat("billing please").await.unwrap();
23348
23349 assert_eq!(agent.current_state().as_deref(), Some("billing"));
23350 assert_eq!(response.content, "Billing answer");
23351 assert_eq!(router_counter.call_count(), 1);
23353 }
23354
23355 #[tokio::test]
23356 async fn test_speculative_skill_capacity_exhausted_still_routes_skill_serially() {
23357 let default_mock = mock_with_response("Draft response");
23358 let router_mock = mock_with_response("helper");
23359 let router_counter = router_mock.clone();
23360 let yaml = skills_with_parallel_transition_yaml(
23362 " speculative_skill_routing: true\n max_speculative_llm_calls_per_turn: 2\n max_parallel_runtime_tasks: 2",
23363 "",
23364 );
23365 let agent =
23366 build_skills_beside_transition_agent(&yaml, role_mocks(default_mock, router_mock));
23367
23368 let response = agent.chat("please use helper").await.unwrap();
23369
23370 assert_eq!(
23371 response.metadata.as_ref().and_then(|m| m.get("skill_id")),
23372 Some(&serde_json::json!("helper"))
23373 );
23374 assert_eq!(router_counter.call_count(), 1);
23375 assert_eq!(
23376 branch_events_with_commit_behavior(&agent, "skill_selection"),
23377 0
23378 );
23379 }
23380
23381 #[tokio::test]
23382 async fn test_speculative_transition_and_skill_both_enabled_unchanged() {
23383 let default_mock = mock_with_response("Draft response");
23384 let router_mock = mock_with_response("helper");
23385 let router_counter = router_mock.clone();
23386 let yaml = skills_with_parallel_transition_yaml(
23387 " speculative_skill_routing: true\n max_speculative_llm_calls_per_turn: 3\n max_parallel_runtime_tasks: 3",
23388 "",
23389 );
23390 let agent =
23391 build_skills_beside_transition_agent(&yaml, role_mocks(default_mock, router_mock));
23392
23393 let response = agent.chat("please use helper").await.unwrap();
23394
23395 assert_eq!(
23396 response.metadata.as_ref().and_then(|m| m.get("skill_id")),
23397 Some(&serde_json::json!("helper"))
23398 );
23399 assert_eq!(router_counter.call_count(), 1);
23400 assert!(branch_events_with_commit_behavior(&agent, "skill_selection") > 0);
23402 }
23403
23404 const BUFFERED_STREAMING_YAML_FRAGMENT: &str = "streaming:\n enabled: true\n buffer_size: 16";
23405 const BUFFERED_OPTIMIZATION_FRAGMENT: &str = " max_speculative_llm_calls_per_turn: 2\n streaming_policy: buffer_until_routing_done\n max_parallel_runtime_tasks: 2";
23406
23407 #[tokio::test]
23408 async fn test_buffered_streaming_skill_wins_after_transition_miss() {
23409 let mut default_mock = mock_with_response("draft one two");
23410 default_mock.set_latency(10);
23411 let router_mock = mock_with_response("helper");
23412 let router_counter = router_mock.clone();
23413 let yaml = skills_with_parallel_transition_yaml(
23414 BUFFERED_OPTIMIZATION_FRAGMENT,
23415 BUFFERED_STREAMING_YAML_FRAGMENT,
23416 );
23417 let agent =
23418 build_skills_beside_transition_agent(&yaml, role_mocks(default_mock, router_mock));
23419
23420 let (content, chunks, final_response) =
23421 collect_stream_events(&agent, "please use helper").await;
23422
23423 assert!(!chunks.iter().any(StreamChunk::is_error), "{chunks:?}");
23424 assert!(
23425 !content.contains("draft"),
23426 "buffered draft must be discarded when a skill wins: {content:?}"
23427 );
23428 let final_response = final_response.expect("Final");
23429 assert_eq!(
23430 final_response
23431 .metadata
23432 .as_ref()
23433 .and_then(|m| m.get("skill_id")),
23434 Some(&serde_json::json!("helper"))
23435 );
23436 assert_eq!(content, final_response.content);
23437 assert_eq!(router_counter.call_count(), 1);
23438 }
23439
23440 #[tokio::test]
23441 async fn test_buffered_streaming_skill_miss_releases_buffer_and_commits_draft() {
23442 let mut default_mock = mock_with_response("draft one two");
23443 default_mock.set_latency(10);
23444 let router_mock = mock_with_response("none");
23445 let router_counter = router_mock.clone();
23446 let yaml = skills_with_parallel_transition_yaml(
23447 BUFFERED_OPTIMIZATION_FRAGMENT,
23448 BUFFERED_STREAMING_YAML_FRAGMENT,
23449 );
23450 let agent =
23451 build_skills_beside_transition_agent(&yaml, role_mocks(default_mock, router_mock));
23452
23453 let (content, chunks, final_response) = collect_stream_events(&agent, "just chat").await;
23454
23455 assert!(!chunks.iter().any(StreamChunk::is_error), "{chunks:?}");
23456 assert_eq!(content, "draft one two");
23457 assert_eq!(final_response.expect("Final").content, "draft one two");
23458 assert_eq!(router_counter.call_count(), 1);
23459 }
23460
23461 #[tokio::test]
23462 async fn test_buffered_streaming_transition_win_skips_skill_selection() {
23463 let default_mock = mock_with_response("Billing answer");
23464 let router_mock = mock_with_response("none");
23465 let router_counter = router_mock.clone();
23466 let yaml = skills_with_parallel_transition_yaml(
23467 BUFFERED_OPTIMIZATION_FRAGMENT,
23468 BUFFERED_STREAMING_YAML_FRAGMENT,
23469 );
23470 let agent =
23471 build_skills_beside_transition_agent(&yaml, role_mocks(default_mock, router_mock));
23472 agent
23473 .set_context("route", serde_json::json!("billing"))
23474 .unwrap();
23475
23476 let (content, chunks, final_response) =
23477 collect_stream_events(&agent, "billing please").await;
23478
23479 assert!(!chunks.iter().any(StreamChunk::is_error), "{chunks:?}");
23480 assert_eq!(agent.current_state().as_deref(), Some("billing"));
23481 assert_eq!(content, "Billing answer");
23482 assert_eq!(final_response.expect("Final").content, "Billing answer");
23483 assert_eq!(router_counter.call_count(), 1);
23485 }
23486
23487 #[tokio::test]
23488 async fn parity_buffered_policy_with_skills() {
23489 let yaml = skills_with_parallel_transition_yaml(
23490 BUFFERED_OPTIMIZATION_FRAGMENT,
23491 BUFFERED_STREAMING_YAML_FRAGMENT,
23492 );
23493 let build = || {
23494 build_skills_beside_transition_agent(
23495 &yaml,
23496 role_mocks(
23497 mock_with_response("draft one two"),
23498 mock_with_response("helper"),
23499 ),
23500 )
23501 };
23502 let (blocking, _, _) = assert_blocking_streaming_parity(build, "please use helper").await;
23503 assert_eq!(
23504 blocking.metadata.as_ref().and_then(|m| m.get("skill_id")),
23505 Some(&serde_json::json!("helper"))
23506 );
23507 }
23508
23509 #[tokio::test]
23510 async fn parity_buffered_policy_with_cot() {
23511 let yaml = format!(
23512 r#"
23513name: BufferedCotAgent
23514system_prompt: "Think first."
23515llm:
23516 default: default
23517streaming:
23518 enabled: true
23519 buffer_size: 16
23520reasoning:
23521 mode: cot
23522runtime:
23523 optimization:
23524 enabled: true
23525 speculative_state_transitions: true
23526{BUFFERED_OPTIMIZATION_FRAGMENT}
23527states:
23528 initial: triage
23529 states:
23530 triage:
23531 prompt: "Triage state."
23532 transitions:
23533 - to: billing
23534 guard:
23535 context:
23536 route:
23537 eq: billing
23538 timing: parallel
23539 billing:
23540 prompt: "Billing state."
23541"#
23542 );
23543 let build = || {
23544 AgentBuilder::from_yaml(&yaml)
23545 .unwrap()
23546 .llm_alias(
23547 "default",
23548 Arc::new(mock_with_response(
23549 "<thinking>step by step</thinking>Reasoned answer",
23550 )),
23551 )
23552 .build()
23553 .unwrap()
23554 };
23555 let (blocking, streamed, _) = assert_blocking_streaming_parity(build, "hello").await;
23556 assert_eq!(blocking.content, "Reasoned answer");
23557 let mode = streamed
23558 .metadata
23559 .as_ref()
23560 .and_then(|m| m.get("reasoning"))
23561 .and_then(|r| r.get("mode_used"))
23562 .cloned();
23563 assert_eq!(
23565 mode,
23566 Some(serde_json::to_value(ReasoningMode::CoT).unwrap())
23567 );
23568 }
23569
23570 fn post_response_transition_yaml(states_extra: &str, billing_extra: &str) -> String {
23577 format!(
23578 r#"
23579name: PostResponseTransitionAgent
23580system_prompt: "You are helpful."
23581streaming:
23582 enabled: true
23583states:
23584 initial: intake
23585{states_extra}
23586 states:
23587 intake:
23588 prompt: "Intake"
23589 transitions:
23590 - to: billing
23591 guard:
23592 context:
23593 route:
23594 eq: billing
23595 billing:
23596 prompt: "Billing"
23597{billing_extra}
23598"#
23599 )
23600 }
23601
23602 fn build_post_response_transition_agent(yaml: &str, mock: MockLLMProvider) -> RuntimeAgent {
23603 let agent = AgentBuilder::from_yaml(yaml)
23604 .unwrap()
23605 .llm(Arc::new(mock))
23606 .build()
23607 .unwrap();
23608 agent
23609 .set_context("route", serde_json::json!("billing"))
23610 .unwrap();
23611 agent
23612 }
23613
23614 fn count_occurrences(haystack: &str, needle: &str) -> usize {
23615 haystack.matches(needle).count()
23616 }
23617
23618 #[tokio::test]
23619 async fn test_stream_transition_without_regeneration_emits_content_once() {
23620 let yaml = post_response_transition_yaml(" regenerate_on_transition: false", "");
23621 let agent =
23622 build_post_response_transition_agent(&yaml, mock_with_response("Intake answer"));
23623
23624 let (content, chunks, final_response) = collect_stream_events(&agent, "hello").await;
23625
23626 assert!(!chunks.iter().any(StreamChunk::is_error), "{chunks:?}");
23627 assert_eq!(agent.current_state().as_deref(), Some("billing"));
23628 assert_eq!(
23629 count_occurrences(&content, "Intake answer"),
23630 1,
23631 "committed content must not be emitted twice: {content:?}"
23632 );
23633 assert_eq!(final_response.expect("Final").content, content);
23634 assert!(
23635 chunks
23636 .iter()
23637 .any(|c| matches!(c, StreamChunk::StateTransition { .. }))
23638 );
23639 }
23640
23641 #[tokio::test]
23642 async fn test_stream_transition_without_regeneration_buffered_emits_content_once() {
23643 let yaml = post_response_transition_yaml(" regenerate_on_transition: false", "");
23644 let mut mock = mock_with_response("Intake answer");
23645 mock.set_tool_choice(Some(ToolChoice::Auto));
23647 let agent = build_post_response_transition_agent(&yaml, mock);
23648
23649 let (content, chunks, final_response) = collect_stream_events(&agent, "hello").await;
23650
23651 assert!(!chunks.iter().any(StreamChunk::is_error), "{chunks:?}");
23652 assert_eq!(agent.current_state().as_deref(), Some("billing"));
23653 assert_eq!(
23654 count_occurrences(&content, "Intake answer"),
23655 1,
23656 "{content:?}"
23657 );
23658 assert_eq!(final_response.expect("Final").content, content);
23659 }
23660
23661 #[tokio::test]
23662 async fn test_stream_state_regenerate_on_enter_false_emits_content_once() {
23663 let yaml = post_response_transition_yaml("", " regenerate_on_enter: false");
23664 let agent =
23665 build_post_response_transition_agent(&yaml, mock_with_response("Intake answer"));
23666
23667 let (content, chunks, final_response) = collect_stream_events(&agent, "hello").await;
23668
23669 assert!(!chunks.iter().any(StreamChunk::is_error), "{chunks:?}");
23670 assert_eq!(agent.current_state().as_deref(), Some("billing"));
23671 assert_eq!(
23672 count_occurrences(&content, "Intake answer"),
23673 1,
23674 "{content:?}"
23675 );
23676 assert_eq!(final_response.expect("Final").content, content);
23677 }
23678
23679 #[tokio::test]
23680 async fn test_stream_transition_with_regeneration_emits_replacement() {
23681 let yaml = post_response_transition_yaml("", "");
23682 let agent = build_post_response_transition_agent(
23683 &yaml,
23684 mock_with_responses(vec!["Intake answer", "Billing answer"]),
23685 );
23686
23687 let (content, chunks, final_response) = collect_stream_events(&agent, "hello").await;
23688
23689 assert!(!chunks.iter().any(StreamChunk::is_error), "{chunks:?}");
23690 assert_eq!(agent.current_state().as_deref(), Some("billing"));
23691 assert_eq!(count_occurrences(&content, "Intake answer"), 1);
23693 assert_eq!(count_occurrences(&content, "Billing answer"), 1);
23694 assert_eq!(final_response.expect("Final").content, "Billing answer");
23695 }
23696
23697 #[tokio::test]
23698 async fn test_blocking_transition_without_regeneration_unchanged() {
23699 let yaml = post_response_transition_yaml(" regenerate_on_transition: false", "");
23700 let agent =
23701 build_post_response_transition_agent(&yaml, mock_with_response("Intake answer"));
23702
23703 let response = agent.chat("hello").await.unwrap();
23704
23705 assert_eq!(response.content, "Intake answer");
23706 assert_eq!(agent.current_state().as_deref(), Some("billing"));
23707 }
23708
23709 #[tokio::test]
23710 async fn parity_transition_regenerate_off() {
23711 let yaml = post_response_transition_yaml(" regenerate_on_transition: false", "");
23712 let build =
23713 || build_post_response_transition_agent(&yaml, mock_with_response("Intake answer"));
23714 assert_blocking_streaming_parity(build, "hello").await;
23715 }
23716
23717 #[tokio::test]
23718 async fn parity_transition_regenerate_on() {
23719 let yaml = post_response_transition_yaml("", "");
23720 let build = || {
23721 build_post_response_transition_agent(
23722 &yaml,
23723 mock_with_responses(vec!["Intake answer", "Billing answer"]),
23724 )
23725 };
23726 let (blocking, _, _) = assert_blocking_streaming_parity(build, "hello").await;
23727 assert_eq!(blocking.content, "Billing answer");
23728 }
23729
23730 fn rejecting_process_processor() -> ProcessProcessor {
23735 use ai_agents_process::ProcessConfig;
23736 let validate_config = ai_agents_process::ValidateStage {
23737 id: Some("length_check".to_string()),
23738 condition: None,
23739 config: ai_agents_process::ValidateConfig {
23740 rules: vec![ai_agents_process::ValidationRule::MinLength {
23741 min_length: 10,
23742 on_fail: ai_agents_process::ValidationAction {
23743 action: ai_agents_process::ValidationActionType::Reject,
23744 message: None,
23745 },
23746 }],
23747 ..Default::default()
23748 },
23749 };
23750 ProcessProcessor::new(ProcessConfig {
23751 input: vec![ai_agents_process::ProcessStage::Validate(validate_config)],
23752 ..Default::default()
23753 })
23754 }
23755
23756 fn looks_like_rejection(content: &str) -> bool {
23758 content.contains("rejected")
23759 || content.contains("Input rejected")
23760 || content.contains("too short")
23761 || content.contains("Too short")
23762 || content.len() < 50
23763 }
23764
23765 #[tokio::test]
23766 async fn test_stream_input_rejection_is_final_response() {
23767 let responses = Arc::new(std::sync::atomic::AtomicUsize::new(0));
23768 let hooks = Arc::new(ResponseCountingHooks {
23769 responses: Arc::clone(&responses),
23770 });
23771 let mock = mock_with_response("Should not reach here.");
23772 let llm_calls = mock.clone();
23773 let agent = AgentBuilder::new()
23774 .system_prompt("You are helpful.")
23775 .llm(Arc::new(mock))
23776 .process_processor(rejecting_process_processor())
23777 .hooks(hooks.clone())
23778 .build()
23779 .unwrap();
23780
23781 let (content, chunks, final_response) = collect_stream_events(&agent, "Hi").await;
23782
23783 assert!(
23784 !chunks.iter().any(StreamChunk::is_error),
23785 "rejection is a response, not a stream error: {chunks:?}"
23786 );
23787 let final_response = final_response.expect("rejection must finalize as Final");
23788 assert!(
23789 looks_like_rejection(&final_response.content),
23790 "Expected rejection response, got: {}",
23791 final_response.content
23792 );
23793 assert_eq!(content, final_response.content);
23794 assert_eq!(
23795 llm_calls.call_count(),
23796 0,
23797 "rejected input must not reach the LLM"
23798 );
23799 assert_eq!(responses.load(Ordering::SeqCst), 1, "on_response must fire");
23800 }
23801
23802 #[tokio::test]
23803 async fn parity_input_rejection() {
23804 let build = || {
23805 AgentBuilder::new()
23806 .system_prompt("You are helpful.")
23807 .llm(Arc::new(mock_with_response("Should not reach here.")))
23808 .process_processor(rejecting_process_processor())
23809 .build()
23810 .unwrap()
23811 };
23812 let (blocking, _, _) = assert_blocking_streaming_parity(build, "Hi").await;
23813 assert!(
23814 looks_like_rejection(&blocking.content),
23815 "{}",
23816 blocking.content
23817 );
23818 }
23819
23820 fn pre_response_transition_yaml(streaming_policy: &str) -> String {
23821 format!(
23822 r#"
23823name: StreamingPreflightAgent
23824system_prompt: "You route before streaming."
23825runtime:
23826 optimization:
23827 enabled: true
23828 pre_response_deterministic_transitions: true
23829 streaming_policy: {streaming_policy}
23830streaming:
23831 enabled: true
23832 buffer_size: 16
23833states:
23834 initial: greeting
23835 states:
23836 greeting:
23837 prompt: "OLD_STATE_SENTINEL"
23838 transitions:
23839 - to: billing
23840 guard:
23841 context:
23842 topic:
23843 eq: billing
23844 timing: pre_response
23845 billing:
23846 prompt: "Billing state."
23847"#
23848 )
23849 }
23850
23851 fn build_pre_response_transition_agent(yaml: &str) -> RuntimeAgent {
23852 let agent = AgentBuilder::from_yaml(yaml)
23853 .unwrap()
23854 .llm(Arc::new(mock_with_response("Billing streamed response")))
23855 .build()
23856 .unwrap();
23857 agent
23858 .set_context("topic", serde_json::json!("billing"))
23859 .unwrap();
23860 agent
23861 }
23862
23863 #[tokio::test]
23864 async fn test_stream_buffered_policy_runs_pre_response_deterministic_transition() {
23865 let yaml = pre_response_transition_yaml("buffer_until_routing_done");
23866 let agent = build_pre_response_transition_agent(&yaml);
23867
23868 let (content, chunks, final_response) =
23869 collect_stream_events(&agent, "billing please").await;
23870
23871 assert!(!chunks.iter().any(StreamChunk::is_error), "{chunks:?}");
23872 assert_eq!(agent.current_state().as_deref(), Some("billing"));
23873 assert!(content.contains("Billing streamed response"));
23874 assert!(!content.contains("OLD_STATE_SENTINEL"));
23875 assert_eq!(final_response.expect("Final").content, content);
23876 }
23877
23878 #[tokio::test]
23879 async fn test_stream_disabled_policy_skips_preflight() {
23880 let yaml = pre_response_transition_yaml("disabled");
23881 let agent = build_pre_response_transition_agent(&yaml);
23882
23883 let (_, chunks, final_response) = collect_stream_events(&agent, "billing please").await;
23884
23885 assert!(!chunks.iter().any(StreamChunk::is_error), "{chunks:?}");
23886 assert!(final_response.is_some());
23887 assert_eq!(agent.current_state().as_deref(), Some("greeting"));
23891 }
23892
23893 #[tokio::test]
23894 async fn parity_pre_response_transition_buffered_policy() {
23895 let yaml = pre_response_transition_yaml("buffer_until_routing_done");
23896 let build = || build_pre_response_transition_agent(&yaml);
23897 assert_blocking_streaming_parity(build, "billing please").await;
23898 }
23899
23900 fn calculator_agent_with(mock: MockLLMProvider) -> RuntimeAgent {
23905 let mut tools = ai_agents_tools::ToolRegistry::new();
23906 tools
23907 .register(Arc::new(ai_agents_tools::CalculatorTool))
23908 .unwrap();
23909 AgentBuilder::new()
23910 .system_prompt("You are a calculator assistant.")
23911 .llm(Arc::new(mock))
23912 .tools(tools)
23913 .build()
23914 .unwrap()
23915 }
23916
23917 #[tokio::test]
23918 async fn test_stream_tool_start_events_precede_results_for_batch() {
23919 let mock = mock_with_responses(vec![
23920 r#"[{"tool": "calculator", "arguments": {"expression": "1+1"}}, {"tool": "calculator", "arguments": {"expression": "2+2"}}]"#,
23921 "Both answers are ready.",
23922 ]);
23923 let agent = calculator_agent_with(mock);
23924
23925 let (_, chunks, final_response) = collect_stream_events(&agent, "compute both").await;
23926
23927 assert!(!chunks.iter().any(StreamChunk::is_error), "{chunks:?}");
23928 let final_response = final_response.expect("Final");
23929 assert_eq!(final_response.tool_calls.as_ref().map(Vec::len), Some(2));
23930
23931 let tool_events: Vec<&StreamChunk> = chunks
23932 .iter()
23933 .filter(|c| {
23934 matches!(
23935 c,
23936 StreamChunk::ToolCallStart { .. }
23937 | StreamChunk::ToolResult { .. }
23938 | StreamChunk::ToolCallEnd { .. }
23939 )
23940 })
23941 .collect();
23942 assert_eq!(tool_events.len(), 6, "{tool_events:?}");
23943 assert!(matches!(tool_events[0], StreamChunk::ToolCallStart { .. }));
23945 assert!(matches!(tool_events[1], StreamChunk::ToolCallStart { .. }));
23946 assert!(matches!(
23947 tool_events[2],
23948 StreamChunk::ToolResult { success: true, .. }
23949 ));
23950 assert!(matches!(tool_events[3], StreamChunk::ToolCallEnd { .. }));
23951 assert!(matches!(
23952 tool_events[4],
23953 StreamChunk::ToolResult { success: true, .. }
23954 ));
23955 assert!(matches!(tool_events[5], StreamChunk::ToolCallEnd { .. }));
23956 }
23957
23958 #[tokio::test]
23959 async fn test_stream_clarification_final_carries_options_and_detection() {
23960 let responses = || {
23961 vec![
23962 r#"{"is_ambiguous":true,"confidence":0.2,"ambiguity_type":"missing_target","reasoning":"target missing","what_is_unclear":["target"],"detected_language":"en"}"#,
23963 r#"{"question":"What should I send?","options":["report","invoice"]}"#,
23964 ]
23965 };
23966 let (blocking_agent, _) = state_disambiguation_agent(responses(), true, None, true);
23967 let (streaming_agent, _) = state_disambiguation_agent(responses(), true, None, true);
23968
23969 let blocking = blocking_agent.chat("Send it").await.unwrap();
23970 let (_, chunks, streamed) = collect_stream_events(&streaming_agent, "Send it").await;
23971 assert!(!chunks.iter().any(StreamChunk::is_error), "{chunks:?}");
23972 let streamed = streamed.expect("clarification must finalize as Final");
23973
23974 assert_eq!(streamed.content, "What should I send?");
23975 let streamed_meta = streamed
23976 .metadata
23977 .as_ref()
23978 .and_then(|m| m.get("disambiguation"))
23979 .cloned()
23980 .expect("disambiguation metadata");
23981 for key in ["status", "options", "clarifying", "detection"] {
23982 assert!(
23983 streamed_meta.get(key).is_some(),
23984 "missing {key}: {streamed_meta}"
23985 );
23986 }
23987 assert_eq!(
23988 streamed_meta.get("detection").and_then(|d| d.get("type")),
23989 Some(&serde_json::json!("missing_target"))
23990 );
23991 assert_eq!(
23992 blocking
23993 .metadata
23994 .as_ref()
23995 .and_then(|m| m.get("disambiguation")),
23996 Some(&streamed_meta),
23997 "blocking and streaming clarification metadata must be identical"
23998 );
23999 }
24000
24001 struct FailingMemory {
24003 messages: parking_lot::RwLock<Vec<ChatMessage>>,
24004 fail_on_add: usize,
24005 adds: std::sync::atomic::AtomicUsize,
24006 }
24007
24008 #[async_trait]
24009 impl ai_agents_core::Memory for FailingMemory {
24010 async fn add_message(&self, message: ChatMessage) -> Result<()> {
24011 let n = self.adds.fetch_add(1, Ordering::SeqCst) + 1;
24012 if n == self.fail_on_add {
24013 return Err(AgentError::Other(format!(
24014 "simulated memory failure on add #{n}"
24015 )));
24016 }
24017 self.messages.write().push(message);
24018 Ok(())
24019 }
24020
24021 async fn get_messages(&self, limit: Option<usize>) -> Result<Vec<ChatMessage>> {
24022 let messages = self.messages.read();
24023 Ok(match limit {
24024 Some(n) if n < messages.len() => messages[messages.len() - n..].to_vec(),
24025 _ => messages.clone(),
24026 })
24027 }
24028
24029 async fn clear(&self) -> Result<()> {
24030 self.messages.write().clear();
24031 Ok(())
24032 }
24033
24034 fn len(&self) -> usize {
24035 self.messages.read().len()
24036 }
24037
24038 async fn restore(&self, snapshot: ai_agents_core::MemorySnapshot) -> Result<()> {
24039 *self.messages.write() = snapshot.messages;
24040 Ok(())
24041 }
24042 }
24043
24044 impl ai_agents_memory::Memory for FailingMemory {}
24045
24046 struct ReadFailingMemory {
24048 messages: parking_lot::RwLock<Vec<ChatMessage>>,
24049 fail_on_read: Option<usize>,
24050 reads: std::sync::atomic::AtomicUsize,
24051 }
24052
24053 impl ReadFailingMemory {
24054 fn new(fail_reads: bool) -> Self {
24056 Self::fail_on_read(fail_reads.then_some(1))
24057 }
24058
24059 fn fail_on_read(fail_on_read: Option<usize>) -> Self {
24061 Self {
24062 messages: parking_lot::RwLock::new(Vec::new()),
24063 fail_on_read,
24064 reads: std::sync::atomic::AtomicUsize::new(0),
24065 }
24066 }
24067
24068 fn read_count(&self) -> usize {
24070 self.reads.load(Ordering::SeqCst)
24071 }
24072 }
24073
24074 #[async_trait]
24075 impl ai_agents_core::Memory for ReadFailingMemory {
24076 async fn add_message(&self, message: ChatMessage) -> Result<()> {
24078 self.messages.write().push(message);
24079 Ok(())
24080 }
24081
24082 async fn get_messages(&self, limit: Option<usize>) -> Result<Vec<ChatMessage>> {
24084 let read = self.reads.fetch_add(1, Ordering::SeqCst) + 1;
24085 if self.fail_on_read == Some(read) {
24086 return Err(AgentError::Other(
24087 "simulated scope memory failure".to_string(),
24088 ));
24089 }
24090 let messages = self.messages.read();
24091 Ok(match limit {
24092 Some(n) if n < messages.len() => messages[messages.len() - n..].to_vec(),
24093 _ => messages.clone(),
24094 })
24095 }
24096
24097 async fn clear(&self) -> Result<()> {
24099 self.messages.write().clear();
24100 Ok(())
24101 }
24102
24103 fn len(&self) -> usize {
24105 self.messages.read().len()
24106 }
24107
24108 async fn restore(&self, snapshot: ai_agents_core::MemorySnapshot) -> Result<()> {
24110 *self.messages.write() = snapshot.messages;
24111 Ok(())
24112 }
24113 }
24114
24115 impl ai_agents_memory::Memory for ReadFailingMemory {}
24116
24117 fn scoped_planning_agent(
24119 memory: Arc<ReadFailingMemory>,
24120 planner: MockLLMProvider,
24121 ) -> RuntimeAgent {
24122 let yaml = r#"
24123name: ScopedPlanningAgent
24124system_prompt: "Plan safely."
24125reasoning:
24126 mode: plan_and_execute
24127tools: [calculator]
24128states:
24129 initial: current
24130 states:
24131 current:
24132 tools: [calculator]
24133"#;
24134 AgentBuilder::from_yaml(yaml)
24135 .unwrap()
24136 .llm(Arc::new(planner))
24137 .tool(Arc::new(CalculatorTool::new()))
24138 .tool(Arc::new(FileWriteTool::new()))
24139 .memory(memory)
24140 .build()
24141 .unwrap()
24142 }
24143
24144 fn scoped_disambiguation_agent(
24146 memory: Arc<ReadFailingMemory>,
24147 router: MockLLMProvider,
24148 include_available_tools: bool,
24149 ) -> RuntimeAgent {
24150 let yaml = format!(
24151 r#"
24152name: ScopedDisambiguationAgent
24153system_prompt: "Clarify safely."
24154llm:
24155 default: default
24156 router: router
24157disambiguation:
24158 enabled: true
24159 context:
24160 recent_messages: 0
24161 include_available_tools: {include_available_tools}
24162tools: [calculator]
24163states:
24164 initial: current
24165 states:
24166 current:
24167 tools: [calculator]
24168"#
24169 );
24170 AgentBuilder::from_yaml(&yaml)
24171 .unwrap()
24172 .llm_alias("default", Arc::new(mock_with_response("unused")))
24173 .llm_alias("router", Arc::new(router))
24174 .tool(Arc::new(CalculatorTool::new()))
24175 .tool(Arc::new(FileWriteTool::new()))
24176 .memory(memory)
24177 .build()
24178 .unwrap()
24179 }
24180
24181 #[tokio::test]
24182 async fn planning_scope_failure_stops_before_the_planner_in_blocking_and_streaming_turns() {
24183 use futures::StreamExt;
24184
24185 let direct_memory = Arc::new(ReadFailingMemory::new(true));
24186 let direct_planner = mock_with_response(r#"{"steps":[]}"#);
24187 let direct_calls = direct_planner.clone();
24188 let direct = scoped_planning_agent(direct_memory.clone(), direct_planner);
24189 let error = direct.generate_plan("plan this").await.unwrap_err();
24190 assert!(error.to_string().contains("simulated scope memory failure"));
24191 assert_eq!(direct_memory.read_count(), 1);
24192 assert_eq!(direct_calls.call_count(), 0);
24193
24194 let blocking_memory = Arc::new(ReadFailingMemory::new(true));
24195 let blocking_planner = mock_with_response(r#"{"steps":[]}"#);
24196 let blocking_calls = blocking_planner.clone();
24197 let blocking = scoped_planning_agent(blocking_memory.clone(), blocking_planner);
24198 let error = blocking.chat("plan this").await.unwrap_err();
24199 assert!(error.to_string().contains("simulated scope memory failure"));
24200 assert_eq!(blocking_memory.read_count(), 1);
24201 assert_eq!(blocking_calls.call_count(), 0);
24202 assert!(blocking.tool_call_history.read().is_empty());
24203
24204 let streaming_memory = Arc::new(ReadFailingMemory::new(true));
24205 let streaming_planner = mock_with_response(r#"{"steps":[]}"#);
24206 let streaming_calls = streaming_planner.clone();
24207 let streaming = scoped_planning_agent(streaming_memory.clone(), streaming_planner);
24208 let (_, chunks, final_response) = collect_stream_events(&streaming, "plan this").await;
24209 assert!(final_response.is_none());
24210 assert!(chunks.iter().any(|chunk| matches!(
24211 chunk,
24212 StreamChunk::Error { message } if message.contains("simulated scope memory failure")
24213 )));
24214 assert!(!chunks.iter().any(StreamChunk::is_done));
24215 assert_eq!(streaming_memory.read_count(), 1);
24216 assert_eq!(streaming_calls.call_count(), 0);
24217 assert!(streaming.tool_call_history.read().is_empty());
24218
24219 let legacy_memory = Arc::new(ReadFailingMemory::new(true));
24220 let legacy_planner = mock_with_response(r#"{"steps":[]}"#);
24221 let legacy_calls = legacy_planner.clone();
24222 let legacy = scoped_planning_agent(legacy_memory.clone(), legacy_planner);
24223 let mut stream = legacy.chat_stream("plan this").await.unwrap();
24224 let mut chunks = Vec::new();
24225 while let Some(chunk) = stream.next().await {
24226 chunks.push(chunk);
24227 }
24228 assert!(chunks.iter().any(|chunk| matches!(
24229 chunk,
24230 StreamChunk::Error { message } if message.contains("simulated scope memory failure")
24231 )));
24232 assert!(!chunks.iter().any(StreamChunk::is_done));
24233 assert_eq!(legacy_memory.read_count(), 1);
24234 assert_eq!(legacy_calls.call_count(), 0);
24235 }
24236
24237 #[tokio::test]
24238 async fn replanning_scope_failure_does_not_issue_a_second_planner_request() {
24239 let memory = Arc::new(ReadFailingMemory::fail_on_read(Some(4)));
24240 let planner = mock_with_response(
24241 r#"{"steps":[{"id":"step1","description":"invalid calculation","action_type":"tool","action_target":"calculator","args":{"expression":"not valid"},"dependencies":[]}]}"#,
24242 );
24243 let planner_calls = planner.clone();
24244 let agent = scoped_planning_agent(memory.clone(), planner);
24245
24246 let result = agent.chat("plan this").await;
24247 assert!(
24248 result.is_err(),
24249 "expected replan scope failure, got {result:?}; reads={}, planner_calls={}",
24250 memory.read_count(),
24251 planner_calls.call_count()
24252 );
24253 let error = result.unwrap_err();
24254 assert!(error.to_string().contains("simulated scope memory failure"));
24255 assert_eq!(memory.read_count(), 4);
24256 assert_eq!(planner_calls.call_count(), 1);
24257 let records = agent.tool_call_history.read();
24258 assert_eq!(records.len(), 1);
24259 assert_eq!(records[0].tool_id, "calculator");
24260 }
24261
24262 #[tokio::test]
24263 async fn planning_prompt_uses_only_the_effective_tool_scope() {
24264 let memory = Arc::new(ReadFailingMemory::new(false));
24265 let planner = mock_with_response(r#"{"steps":[]}"#);
24266 let planner_calls = planner.clone();
24267 let agent = scoped_planning_agent(memory.clone(), planner);
24268
24269 let plan = agent.generate_plan("plan this").await.unwrap();
24270 assert!(!plan.steps.is_empty());
24271 assert_eq!(memory.read_count(), 1);
24272 assert_eq!(planner_calls.call_count(), 1);
24273 let call = planner_calls.last_call().unwrap();
24274 let prompt = &call.messages[0].content;
24275 assert!(prompt.contains("- calculator ("), "{prompt}");
24276 assert!(!prompt.contains("file_write"), "{prompt}");
24277 }
24278
24279 #[tokio::test]
24280 async fn planning_with_no_granted_tools_keeps_a_normal_empty_scope() {
24281 let yaml = r#"
24282name: EmptyPlanningAgent
24283system_prompt: "Plan safely."
24284reasoning:
24285 mode: plan_and_execute
24286tools: []
24287"#;
24288 let planner = mock_with_response(r#"{"steps":[]}"#);
24289 let planner_calls = planner.clone();
24290 let agent = AgentBuilder::from_yaml(yaml)
24291 .unwrap()
24292 .llm(Arc::new(planner))
24293 .tool(Arc::new(CalculatorTool::new()))
24294 .build()
24295 .unwrap();
24296
24297 agent.generate_plan("plan this").await.unwrap();
24298 let call = planner_calls.last_call().unwrap();
24299 let prompt = &call.messages[0].content;
24300 assert!(prompt.contains("Available tools: none"), "{prompt}");
24301 assert!(!prompt.contains("- calculator ("), "{prompt}");
24302 }
24303
24304 #[tokio::test]
24305 async fn planning_filter_can_narrow_the_effective_scope_to_empty() {
24306 let yaml = r#"
24307name: FilteredPlanningAgent
24308system_prompt: "Plan safely."
24309reasoning:
24310 mode: plan_and_execute
24311 planning:
24312 available:
24313 tools: []
24314tools: [calculator]
24315states:
24316 initial: current
24317 states:
24318 current:
24319 tools: [calculator]
24320"#;
24321 let memory = Arc::new(ReadFailingMemory::new(false));
24322 let planner = mock_with_response(r#"{"steps":[]}"#);
24323 let planner_calls = planner.clone();
24324 let agent = AgentBuilder::from_yaml(yaml)
24325 .unwrap()
24326 .llm(Arc::new(planner))
24327 .tool(Arc::new(CalculatorTool::new()))
24328 .memory(memory.clone())
24329 .build()
24330 .unwrap();
24331
24332 agent.generate_plan("plan this").await.unwrap();
24333 assert_eq!(memory.read_count(), 1);
24334 let call = planner_calls.last_call().unwrap();
24335 let prompt = &call.messages[0].content;
24336 assert!(prompt.contains("Available tools: none"), "{prompt}");
24337 assert!(!prompt.contains("- calculator ("), "{prompt}");
24338 }
24339
24340 #[tokio::test]
24341 async fn disambiguation_scope_failure_never_reaches_a_model_or_final_response() {
24342 let direct_memory = Arc::new(ReadFailingMemory::new(true));
24343 let direct_router = mock_with_response("unused");
24344 let direct_calls = direct_router.clone();
24345 let direct = scoped_disambiguation_agent(direct_memory.clone(), direct_router, true);
24346 let error = direct.build_disambiguation_context().await.unwrap_err();
24347 assert!(error.to_string().contains("simulated scope memory failure"));
24348 assert_eq!(direct_memory.read_count(), 1);
24349 assert_eq!(direct_calls.call_count(), 0);
24350
24351 let blocking_memory = Arc::new(ReadFailingMemory::new(true));
24352 let blocking_router = mock_with_response("unused");
24353 let blocking_calls = blocking_router.clone();
24354 let blocking = scoped_disambiguation_agent(blocking_memory.clone(), blocking_router, true);
24355 let error = blocking.chat("send it").await.unwrap_err();
24356 assert!(error.to_string().contains("simulated scope memory failure"));
24357 assert_eq!(blocking_memory.read_count(), 1);
24358 assert_eq!(blocking_calls.call_count(), 0);
24359
24360 let streaming_memory = Arc::new(ReadFailingMemory::new(true));
24361 let streaming_router = mock_with_response("unused");
24362 let streaming_calls = streaming_router.clone();
24363 let streaming =
24364 scoped_disambiguation_agent(streaming_memory.clone(), streaming_router, true);
24365 let (_, chunks, final_response) = collect_stream_events(&streaming, "send it").await;
24366 assert!(final_response.is_none());
24367 assert!(chunks.iter().any(|chunk| matches!(
24368 chunk,
24369 StreamChunk::Error { message } if message.contains("simulated scope memory failure")
24370 )));
24371 assert!(!chunks.iter().any(StreamChunk::is_done));
24372 assert_eq!(streaming_memory.read_count(), 1);
24373 assert_eq!(streaming_calls.call_count(), 0);
24374
24375 let skipped_memory = Arc::new(ReadFailingMemory::new(true));
24376 let skipped_router = mock_with_response("unused");
24377 let skipped = scoped_disambiguation_agent(skipped_memory.clone(), skipped_router, false);
24378 let context = skipped.build_disambiguation_context().await.unwrap();
24379 assert!(context.available_tools.is_empty());
24380 assert_eq!(skipped_memory.read_count(), 0);
24381 }
24382
24383 #[tokio::test]
24384 async fn test_stream_memory_write_failure_surfaces_as_error() {
24385 let yaml = r#"
24388name: TransitionOnToolCallAgent
24389system_prompt: "You are helpful."
24390streaming:
24391 enabled: true
24392states:
24393 initial: intake
24394 states:
24395 intake:
24396 prompt: "Intake"
24397 transitions:
24398 - to: billing
24399 guard:
24400 context:
24401 route:
24402 eq: billing
24403 billing:
24404 prompt: "Billing"
24405"#;
24406 let build = |fail_on_add: usize| {
24407 let mut tools = ai_agents_tools::ToolRegistry::new();
24408 tools
24409 .register(Arc::new(ai_agents_tools::CalculatorTool))
24410 .unwrap();
24411 let agent = AgentBuilder::from_yaml(yaml)
24412 .unwrap()
24413 .llm(Arc::new(mock_with_responses(vec![
24414 r#"{"tool": "calculator", "arguments": {"expression": "1+1"}}"#,
24415 "Billing answer",
24416 ])))
24417 .tools(tools)
24418 .memory(Arc::new(FailingMemory {
24419 messages: parking_lot::RwLock::new(Vec::new()),
24420 fail_on_add,
24421 adds: std::sync::atomic::AtomicUsize::new(0),
24422 }))
24423 .build()
24424 .unwrap();
24425 agent
24426 .set_context("route", serde_json::json!("billing"))
24427 .unwrap();
24428 agent
24429 };
24430
24431 let blocking = build(2).chat("compute").await;
24432 assert!(
24433 blocking.is_err(),
24434 "blocking must surface the memory failure"
24435 );
24436
24437 let (_, chunks, final_response) = collect_stream_events(&build(2), "compute").await;
24438 assert!(
24439 final_response.is_none(),
24440 "streaming must not finalize after a memory failure"
24441 );
24442 assert!(
24443 chunks.iter().any(|c| matches!(c, StreamChunk::Error { message } if message.contains("simulated memory failure"))),
24444 "streaming must surface the memory failure: {chunks:?}"
24445 );
24446
24447 assert!(build(usize::MAX).chat("compute").await.is_ok());
24449 }
24450
24451 #[tokio::test]
24452 async fn parity_tool_execution() {
24453 let build = || {
24454 calculator_agent_with(mock_with_responses(vec![
24455 r#"{"tool": "calculator", "arguments": {"expression": "2+2"}}"#,
24456 "The answer is 4.",
24457 ]))
24458 };
24459 let (blocking, _, chunks) = assert_blocking_streaming_parity(build, "What is 2+2?").await;
24460 assert_eq!(blocking.content, "The answer is 4.");
24461 assert!(
24462 chunks
24463 .iter()
24464 .any(|c| matches!(c, StreamChunk::ToolResult { .. }))
24465 );
24466 }
24467
24468 #[tokio::test]
24469 async fn parity_disambiguation_clarification() {
24470 let build = || {
24471 state_disambiguation_agent(
24472 vec![
24473 r#"{"is_ambiguous":true,"confidence":0.2,"ambiguity_type":"missing_target","reasoning":"target missing","what_is_unclear":["target"],"detected_language":"en"}"#,
24474 r#"{"question":"What should I send?","options":null}"#,
24475 ],
24476 true,
24477 None,
24478 true,
24479 )
24480 .0
24481 };
24482 let (blocking, _, _) = assert_blocking_streaming_parity(build, "Send it").await;
24483 assert_eq!(blocking.content, "What should I send?");
24484 }
24485
24486 #[tokio::test]
24487 async fn runtime_disambiguation_uses_configured_recent_history_projection() {
24488 let yaml = r#"
24489name: DisambiguationContextAgent
24490system_prompt: "Help."
24491llm:
24492 default: default
24493 router: router
24494disambiguation:
24495 enabled: true
24496 detection:
24497 llm: router
24498 context:
24499 recent_messages: 1
24500 include_state: false
24501 include_available_tools: false
24502states:
24503 initial: private_state
24504 states:
24505 private_state:
24506 prompt: "PRIVATE_STATE_PROMPT"
24507"#;
24508 let main = mock_with_responses(vec!["FIRST_MAIN_MARKER", "SECOND_MAIN_MARKER"]);
24509 let router = mock_with_responses(vec![
24510 r#"{"is_ambiguous":false,"confidence":0.9,"ambiguity_type":null,"reasoning":"clear","what_is_unclear":[],"detected_language":"en"}"#,
24511 r#"{"is_ambiguous":false,"confidence":0.9,"ambiguity_type":null,"reasoning":"clear","what_is_unclear":[],"detected_language":"en"}"#,
24512 ]);
24513 let router_calls = router.clone();
24514 let agent = AgentBuilder::from_yaml(yaml)
24515 .unwrap()
24516 .llm_alias("default", Arc::new(main))
24517 .llm_alias("router", Arc::new(router))
24518 .build()
24519 .unwrap();
24520
24521 agent.chat("FIRST_USER_MARKER").await.unwrap();
24522 agent.chat("SECOND_USER_MARKER").await.unwrap();
24523
24524 let calls = router_calls.call_history();
24525 assert_eq!(calls.len(), 2);
24526 let second_prompt = &calls[1].messages.last().unwrap().content;
24527 assert!(
24528 second_prompt.contains("FIRST_MAIN_MARKER"),
24529 "{second_prompt}"
24530 );
24531 assert!(
24532 !second_prompt.contains("FIRST_USER_MARKER"),
24533 "{second_prompt}"
24534 );
24535 assert!(
24536 !second_prompt.contains("PRIVATE_STATE_PROMPT"),
24537 "{second_prompt}"
24538 );
24539 }
24540
24541 #[tokio::test]
24542 async fn runtime_zero_history_preserves_pending_clarification_across_turns() {
24543 let yaml = r#"
24544name: ZeroHistoryPendingAgent
24545system_prompt: "Help."
24546llm:
24547 default: default
24548 router: router
24549disambiguation:
24550 enabled: true
24551 context:
24552 recent_messages: 0
24553 include_available_tools: false
24554"#;
24555 let main = mock_with_response("Final answer");
24556 let main_calls = main.clone();
24557 let router = mock_with_responses(vec![
24558 r#"{"is_ambiguous":true,"confidence":0.1,"ambiguity_type":"missing_target","reasoning":"target missing","what_is_unclear":["target"],"detected_language":"en"}"#,
24559 r#"{"question":"Who should receive it?","options":null}"#,
24560 r#"{"status":"answered","selected_option":null,"enriched_input":"Send it to Ada","resolved":{"recipient":"Ada"}}"#,
24561 ]);
24562 let router_calls = router.clone();
24563 let agent = AgentBuilder::from_yaml(yaml)
24564 .unwrap()
24565 .llm_alias("default", Arc::new(main))
24566 .llm_alias("router", Arc::new(router))
24567 .build()
24568 .unwrap();
24569
24570 let clarification = agent.chat("Send it").await.unwrap();
24571 assert_eq!(clarification.content, "Who should receive it?");
24572 let response = agent.chat("Ada").await.unwrap();
24573 assert_eq!(response.content, "Final answer");
24574 assert_eq!(router_calls.call_count(), 3);
24575 assert_eq!(main_calls.call_count(), 1);
24576 let calls = router_calls.call_history();
24577 let parse_prompt = &calls[2].messages.last().unwrap().content;
24578 assert!(
24579 parse_prompt.contains("Who should receive it?"),
24580 "{parse_prompt}"
24581 );
24582 }
24583
24584 #[tokio::test]
24585 async fn parity_reflection_enabled() {
24586 let yaml = r#"
24587name: ReflectionAgent
24588system_prompt: "You are careful."
24589reflection:
24590 enabled: true
24591 criteria:
24592 - "Is the answer helpful?"
24593"#;
24594 let build = || {
24595 AgentBuilder::from_yaml(yaml)
24596 .unwrap()
24597 .llm(Arc::new(mock_with_responses(vec![
24598 "Main answer",
24599 "OVERALL: PASS\nCONFIDENCE: 0.9",
24600 ])))
24601 .build()
24602 .unwrap()
24603 };
24604 let (blocking, streamed, _) = assert_blocking_streaming_parity(build, "hello").await;
24605 assert_eq!(blocking.content, "Main answer");
24606 assert!(metadata_keys(&streamed).contains("reflection"));
24607 }
24608
24609 fn state_reflection_agent(
24610 global_retries: u32,
24611 state_retries: u32,
24612 main: MockLLMProvider,
24613 evaluator: MockLLMProvider,
24614 ) -> RuntimeAgent {
24615 let yaml = format!(
24616 r#"
24617name: StateReflectionAgent
24618system_prompt: "You are careful."
24619llm:
24620 default: default
24621 router: evaluator
24622reflection:
24623 enabled: true
24624 evaluator_llm: evaluator
24625 max_retries: {global_retries}
24626 criteria:
24627 - "Global criterion"
24628states:
24629 initial: active
24630 states:
24631 active:
24632 prompt: "Handle the active state."
24633 reflection:
24634 enabled: true
24635 evaluator_llm: evaluator
24636 max_retries: {state_retries}
24637 criteria:
24638 - "State criterion"
24639"#
24640 );
24641 AgentBuilder::from_yaml(&yaml)
24642 .unwrap()
24643 .llm_alias("default", Arc::new(main))
24644 .llm_alias("evaluator", Arc::new(evaluator))
24645 .build()
24646 .unwrap()
24647 }
24648
24649 #[test]
24650 fn routing_log_reports_the_effective_state_reflection_mode() {
24651 let yaml = r#"
24652name: ReflectionLogAgent
24653system_prompt: "Be concise."
24654reflection:
24655 enabled: true
24656states:
24657 initial: active
24658 states:
24659 active:
24660 prompt: "Answer directly."
24661 reflection:
24662 enabled: false
24663"#;
24664 let agent = AgentBuilder::from_yaml(yaml)
24665 .unwrap()
24666 .llm(Arc::new(mock_with_response("answer")))
24667 .build()
24668 .unwrap();
24669 assert!(matches!(
24670 agent.routing_reflection_mode(),
24671 ReflectionMode::Disabled
24672 ));
24673 }
24674
24675 #[tokio::test]
24676 async fn state_reflection_zero_retries_overrides_global_limit() {
24677 let main = mock_with_response("First answer");
24678 let main_calls = main.clone();
24679 let evaluator = mock_with_response("OVERALL: FAIL\nCONFIDENCE: 0.1");
24680 let evaluator_calls = evaluator.clone();
24681 let agent = state_reflection_agent(2, 0, main, evaluator);
24682
24683 let response = agent.chat("hello").await.unwrap();
24684 let reflection = response
24685 .metadata
24686 .as_ref()
24687 .and_then(|metadata| metadata.get("reflection"))
24688 .expect("reflection metadata");
24689
24690 assert_eq!(response.content, "First answer");
24691 assert_eq!(reflection["attempts"], 1);
24692 assert_eq!(main_calls.call_count(), 1);
24693 assert_eq!(evaluator_calls.call_count(), 1);
24694 }
24695
24696 #[tokio::test]
24697 async fn state_reflection_retry_limit_overrides_zero_global_limit() {
24698 let main = mock_with_responses(vec!["First answer", "Second answer", "Third answer"]);
24699 let main_calls = main.clone();
24700 let evaluator = mock_with_responses(vec![
24701 "OVERALL: FAIL\nCONFIDENCE: 0.1",
24702 "OVERALL: FAIL\nCONFIDENCE: 0.1",
24703 "OVERALL: FAIL\nCONFIDENCE: 0.1",
24704 ]);
24705 let evaluator_calls = evaluator.clone();
24706 let agent = state_reflection_agent(0, 2, main, evaluator);
24707
24708 let response = agent.chat("hello").await.unwrap();
24709 let reflection = response
24710 .metadata
24711 .as_ref()
24712 .and_then(|metadata| metadata.get("reflection"))
24713 .expect("reflection metadata");
24714
24715 assert_eq!(response.content, "Third answer");
24716 assert_eq!(reflection["attempts"], 3);
24717 assert_eq!(main_calls.call_count(), 3);
24718 assert_eq!(evaluator_calls.call_count(), 3);
24719 }
24720
24721 #[tokio::test]
24722 async fn state_reflection_override_is_preserved_in_event_stream_metadata() {
24723 let main = mock_with_responses(vec!["First answer", "Second answer", "Third answer"]);
24724 let main_calls = main.clone();
24725 let evaluator = mock_with_responses(vec![
24726 "OVERALL: FAIL\nCONFIDENCE: 0.1",
24727 "OVERALL: FAIL\nCONFIDENCE: 0.1",
24728 "OVERALL: FAIL\nCONFIDENCE: 0.1",
24729 ]);
24730 let evaluator_calls = evaluator.clone();
24731 let agent = state_reflection_agent(0, 2, main, evaluator);
24732
24733 let (content, chunks, final_response) = collect_stream_events(&agent, "hello").await;
24734 let final_response = final_response.expect("successful stream Final");
24735 let reflection = final_response
24736 .metadata
24737 .as_ref()
24738 .and_then(|metadata| metadata.get("reflection"))
24739 .expect("reflection metadata");
24740
24741 assert_eq!(content, "Third answer");
24742 assert!(!chunks.iter().any(StreamChunk::is_error));
24743 assert_eq!(reflection["attempts"], 3);
24744 assert_eq!(main_calls.call_count(), 3);
24745 assert_eq!(evaluator_calls.call_count(), 3);
24746 }
24747
24748 #[tokio::test]
24749 async fn state_reflection_override_preserves_legacy_stream_completion() {
24750 use futures::StreamExt;
24751
24752 let main = mock_with_response("First answer");
24753 let main_calls = main.clone();
24754 let evaluator = mock_with_response("OVERALL: FAIL\nCONFIDENCE: 0.1");
24755 let evaluator_calls = evaluator.clone();
24756 let agent = state_reflection_agent(2, 0, main, evaluator);
24757 let mut stream = agent.chat_stream("hello").await.unwrap();
24758 let mut content = String::new();
24759 let mut done = 0;
24760 while let Some(chunk) = stream.next().await {
24761 match chunk {
24762 StreamChunk::Content { text } => content.push_str(&text),
24763 StreamChunk::Done {} => done += 1,
24764 StreamChunk::Error { message } => panic!("unexpected error: {message}"),
24765 _ => {}
24766 }
24767 }
24768
24769 assert_eq!(content, "First answer");
24770 assert_eq!(done, 1);
24771 assert_eq!(main_calls.call_count(), 1);
24772 assert_eq!(evaluator_calls.call_count(), 1);
24773 }
24774
24775 #[tokio::test]
24776 async fn reflection_evaluator_error_emits_no_event_stream_final() {
24777 let main = mock_with_response("First answer");
24778 let mut evaluator = MockLLMProvider::new("evaluator");
24779 evaluator.set_error("judge failed");
24780 let agent = state_reflection_agent(0, 2, main, evaluator);
24781
24782 let (_, chunks, final_response) = collect_stream_events(&agent, "hello").await;
24783
24784 assert!(final_response.is_none());
24785 assert!(chunks.iter().any(StreamChunk::is_error));
24786 assert!(!chunks.iter().any(StreamChunk::is_done));
24787 }
24788
24789 #[tokio::test]
24790 async fn parity_cot_hidden_thinking() {
24791 let yaml = r#"
24792name: CotHiddenAgent
24793system_prompt: "Think first."
24794reasoning:
24795 mode: cot
24796 output: hidden
24797"#;
24798 let build = || {
24799 AgentBuilder::from_yaml(yaml)
24800 .unwrap()
24801 .llm(Arc::new(mock_with_response(
24802 "<thinking>step by step</thinking>Visible answer",
24803 )))
24804 .build()
24805 .unwrap()
24806 };
24807 let (blocking, streamed, chunks) = assert_blocking_streaming_parity(build, "hello").await;
24808 assert_eq!(blocking.content, "Visible answer");
24809 assert_eq!(content_chunks(&chunks).concat(), streamed.content);
24811 }
24812
24813 fn content_chunks(chunks: &[StreamChunk]) -> Vec<String> {
24818 chunks
24819 .iter()
24820 .filter_map(|c| match c {
24821 StreamChunk::Content { text } => Some(text.clone()),
24822 _ => None,
24823 })
24824 .collect()
24825 }
24826
24827 fn reflection_auto_agent(main: MockLLMProvider, judge: MockLLMProvider) -> RuntimeAgent {
24828 let yaml = r#"
24829name: ReflectionAutoAgent
24830system_prompt: "You are careful."
24831llm:
24832 default: default
24833 router: router
24834reflection:
24835 enabled: auto
24836 evaluator_llm: router
24837 criteria:
24838 - "Is the answer helpful?"
24839"#;
24840 AgentBuilder::from_yaml(yaml)
24841 .unwrap()
24842 .llm_alias("default", Arc::new(main))
24843 .llm_alias("router", Arc::new(judge))
24844 .build()
24845 .unwrap()
24846 }
24847
24848 #[tokio::test]
24849 async fn test_stream_reflection_auto_buffers_and_calls_judge_once_per_iteration() {
24850 let judge = mock_with_responses(vec!["YES", "OVERALL: PASS\nCONFIDENCE: 0.9"]);
24851 let judge_calls = judge.clone();
24852 let agent = reflection_auto_agent(mock_with_response("Main answer one two"), judge);
24853
24854 let (content, chunks, final_response) = collect_stream_events(&agent, "hello").await;
24855
24856 assert!(!chunks.iter().any(StreamChunk::is_error), "{chunks:?}");
24857 assert_eq!(
24858 content_chunks(&chunks).len(),
24859 1,
24860 "auto reflection must buffer the main response: {chunks:?}"
24861 );
24862 assert_eq!(content, "Main answer one two");
24863 assert_eq!(final_response.expect("Final").content, content);
24864 assert_eq!(judge_calls.call_count(), 2);
24866 }
24867
24868 #[tokio::test]
24869 async fn test_stream_reflection_auto_rewrite_is_streamed() {
24870 let judge = mock_with_responses(vec![
24871 "YES",
24872 "OVERALL: FAIL\nCONFIDENCE: 0.1",
24873 "OVERALL: PASS\nCONFIDENCE: 0.9",
24874 ]);
24875 let agent = reflection_auto_agent(
24876 mock_with_responses(vec!["First attempt", "Improved answer"]),
24877 judge,
24878 );
24879
24880 let (content, chunks, final_response) = collect_stream_events(&agent, "hello").await;
24881
24882 assert!(!chunks.iter().any(StreamChunk::is_error), "{chunks:?}");
24883 assert_eq!(content_chunks(&chunks).len(), 1, "{chunks:?}");
24884 assert_eq!(
24885 content, "Improved answer",
24886 "the rewritten answer is what streams"
24887 );
24888 assert_eq!(final_response.expect("Final").content, "Improved answer");
24889 }
24890
24891 fn reasoning_agent(mode: &str, output: &str) -> RuntimeAgent {
24892 let yaml = format!(
24893 r#"
24894name: ReasoningStreamAgent
24895system_prompt: "Think first."
24896reasoning:
24897 mode: {mode}
24898 output: {output}
24899"#
24900 );
24901 AgentBuilder::from_yaml(&yaml)
24902 .unwrap()
24903 .llm(Arc::new(mock_with_response(
24904 "<thinking>step by step</thinking>Visible answer",
24905 )))
24906 .build()
24907 .unwrap()
24908 }
24909
24910 #[tokio::test]
24911 async fn test_stream_cot_hidden_emits_no_thinking_tags() {
24912 let agent = reasoning_agent("cot", "hidden");
24913 let (content, chunks, final_response) = collect_stream_events(&agent, "hello").await;
24914 assert!(!chunks.iter().any(StreamChunk::is_error), "{chunks:?}");
24915 assert_eq!(content_chunks(&chunks).len(), 1, "{chunks:?}");
24916 assert!(!content.contains("<thinking>"), "{content:?}");
24917 assert_eq!(content, "Visible answer");
24918 assert_eq!(final_response.expect("Final").content, content);
24919 }
24920
24921 #[tokio::test]
24922 async fn test_stream_cot_visible_matches_final_format() {
24923 let agent = reasoning_agent("cot", "visible");
24924 let (content, chunks, final_response) = collect_stream_events(&agent, "hello").await;
24925 assert!(!chunks.iter().any(StreamChunk::is_error), "{chunks:?}");
24926 assert!(content.starts_with("Thinking:"), "{content:?}");
24927 assert!(content.contains("Answer:\nVisible answer"), "{content:?}");
24928 assert_eq!(final_response.expect("Final").content, content);
24929 }
24930
24931 #[tokio::test]
24932 async fn test_stream_react_buffers() {
24933 let agent = reasoning_agent("react", "hidden");
24934 let (content, chunks, final_response) = collect_stream_events(&agent, "hello").await;
24935 assert!(!chunks.iter().any(StreamChunk::is_error), "{chunks:?}");
24936 assert_eq!(content_chunks(&chunks).len(), 1, "{chunks:?}");
24937 assert_eq!(content, "Visible answer");
24938 assert_eq!(final_response.expect("Final").content, content);
24939 }
24940
24941 #[tokio::test]
24942 async fn test_stream_plain_mode_still_streams_deltas() {
24943 let agent = AgentBuilder::new()
24944 .system_prompt("You are helpful.")
24945 .llm(Arc::new(mock_with_response("one two three")))
24946 .build()
24947 .unwrap();
24948 let (content, chunks, final_response) = collect_stream_events(&agent, "hello").await;
24949 assert!(
24950 content_chunks(&chunks).len() >= 2,
24951 "plain turns must keep token-level streaming: {chunks:?}"
24952 );
24953 assert_eq!(content, "one two three");
24954 assert_eq!(final_response.expect("Final").content, content);
24955 }
24956
24957 struct ActorProbeHooks {
24959 seen: parking_lot::Mutex<Option<crate::TurnActorContext>>,
24960 }
24961
24962 #[async_trait]
24963 impl AgentHooks for ActorProbeHooks {
24964 async fn on_message_received(&self, _input: &str) {
24965 *self.seen.lock() = current_turn_actor_context();
24966 }
24967 }
24968
24969 fn actor_probe_agent(hooks: Arc<ActorProbeHooks>) -> RuntimeAgent {
24970 let yaml = r#"
24971name: ActorStreamAgent
24972system_prompt: "You are helpful."
24973observability:
24974 enabled: true
24975 export:
24976 write_raw_events: true
24977"#;
24978 AgentBuilder::from_yaml(yaml)
24979 .unwrap()
24980 .llm(Arc::new(mock_with_response("Hello actor")))
24981 .hooks(hooks)
24982 .build()
24983 .unwrap()
24984 }
24985
24986 async fn collect_actor_stream_final(
24987 agent: &RuntimeAgent,
24988 input: &str,
24989 actor_context: crate::TurnActorContext,
24990 ) -> AgentResponse {
24991 use futures::StreamExt;
24992 let mut events = agent
24993 .chat_stream_events_with_actor_context(input, actor_context)
24994 .await
24995 .expect("stream opens");
24996 let mut final_response = None;
24997 while let Some(event) = events.next().await {
24998 match event {
24999 AgentStreamEvent::Final(response) => final_response = Some(response),
25000 AgentStreamEvent::Chunk(StreamChunk::Error { message }) => {
25001 panic!("unexpected stream error: {message}")
25002 }
25003 AgentStreamEvent::Chunk(_) => {}
25004 }
25005 }
25006 final_response.expect("Final")
25007 }
25008
25009 #[tokio::test]
25010 async fn test_stream_events_with_actor_context_scopes_actor_for_turn() {
25011 let hooks = Arc::new(ActorProbeHooks {
25012 seen: parking_lot::Mutex::new(None),
25013 });
25014 let agent = actor_probe_agent(Arc::clone(&hooks));
25015 let actor_context = crate::TurnActorContext::new().with_origin_actor("customer_42");
25016
25017 let final_response = collect_actor_stream_final(&agent, "hi", actor_context).await;
25018
25019 assert_eq!(final_response.content, "Hello actor");
25020 assert_eq!(
25021 hooks
25022 .seen
25023 .lock()
25024 .as_ref()
25025 .and_then(|context| context.effective_actor_id().map(str::to_string)),
25026 Some("customer_42".to_string()),
25027 "the actor context must be visible inside the streaming turn"
25028 );
25029 assert!(
25030 agent.actor_id().is_none(),
25031 "a turn-scoped actor must not mutate the global actor ID"
25032 );
25033 let events = agent.observability().unwrap().raw_events();
25034 assert!(
25035 events
25036 .iter()
25037 .any(|event| event.dimensions.get("actor") == Some(&"customer_42".to_string())),
25038 "observation events must carry the actor dimension"
25039 );
25040 }
25041
25042 #[tokio::test]
25043 async fn test_stream_events_with_actor_context_matches_blocking_actor_context() {
25044 let actor_context = crate::TurnActorContext::new()
25045 .with_origin_actor("customer_42")
25046 .with_sender_agent("coordinator");
25047
25048 let blocking_hooks = Arc::new(ActorProbeHooks {
25049 seen: parking_lot::Mutex::new(None),
25050 });
25051 let blocking_agent = actor_probe_agent(Arc::clone(&blocking_hooks));
25052 let blocking = blocking_agent
25053 .chat_with_actor_context("hi", actor_context.clone())
25054 .await
25055 .unwrap();
25056
25057 let streaming_hooks = Arc::new(ActorProbeHooks {
25058 seen: parking_lot::Mutex::new(None),
25059 });
25060 let streaming_agent = actor_probe_agent(Arc::clone(&streaming_hooks));
25061 let streamed =
25062 collect_actor_stream_final(&streaming_agent, "hi", actor_context.clone()).await;
25063
25064 assert_eq!(blocking.content, streamed.content);
25065 assert_eq!(metadata_keys(&blocking), metadata_keys(&streamed));
25066 assert_eq!(
25067 *blocking_hooks.seen.lock(),
25068 *streaming_hooks.seen.lock(),
25069 "both entry points must expose the same turn actor context"
25070 );
25071 assert_eq!(*streaming_hooks.seen.lock(), Some(actor_context));
25072 }
25073
25074 #[tokio::test]
25075 async fn test_stream_events_with_actor_context_releases_root_turn_on_drop() {
25076 use futures::StreamExt;
25077 let agent = AgentBuilder::new()
25078 .system_prompt("You are helpful.")
25079 .llm(Arc::new(mock_with_response("one two three")))
25080 .build()
25081 .unwrap();
25082 {
25083 let mut events = agent
25084 .chat_stream_events_with_actor_context(
25085 "hi",
25086 crate::TurnActorContext::new().with_origin_actor("customer_42"),
25087 )
25088 .await
25089 .unwrap();
25090 let _first = events.next().await;
25092 }
25093 let next = tokio::time::timeout(Duration::from_secs(5), agent.chat("next")).await;
25094 assert!(
25095 matches!(next, Ok(Ok(_))),
25096 "the root turn must be released when the actor stream is dropped: {next:?}"
25097 );
25098 }
25099
25100 fn skill_scope_agent_with_router(states: &str) -> (RuntimeAgent, MockLLMProvider) {
25101 let yaml = format!(
25102 r#"
25103name: SkillScopeAgent
25104system_prompt: "Route skills."
25105skills:
25106 - id: alpha
25107 description: "Alpha"
25108 trigger: "alpha"
25109 steps:
25110 - prompt: "alpha {{{{ user_input }}}}"
25111 - id: beta
25112 description: "Beta"
25113 trigger: "beta"
25114 steps:
25115 - prompt: "beta {{{{ user_input }}}}"
25116{states}
25117"#
25118 );
25119 let router = mock_with_response("none");
25120 let calls = router.clone();
25121 let agent = AgentBuilder::from_yaml(&yaml)
25122 .unwrap()
25123 .llm(Arc::new(router))
25124 .build()
25125 .unwrap();
25126 (agent, calls)
25127 }
25128
25129 fn skill_scope_agent(states: &str) -> RuntimeAgent {
25130 skill_scope_agent_with_router(states).0
25131 }
25132
25133 fn available_skill_ids(agent: &RuntimeAgent) -> Vec<String> {
25134 let mut ids: Vec<String> = agent
25135 .get_available_skills()
25136 .into_iter()
25137 .map(|skill| skill.id.clone())
25138 .collect();
25139 ids.sort();
25140 ids
25141 }
25142
25143 #[test]
25144 fn state_skill_scope_characterizes_empty_inheritance_and_unknown_ids() {
25145 assert_eq!(
25146 available_skill_ids(&skill_scope_agent("")),
25147 vec!["alpha", "beta"]
25148 );
25149 assert_eq!(
25150 available_skill_ids(&skill_scope_agent(
25151 "states:\n initial: current\n states:\n current:\n skills: []\n"
25152 )),
25153 vec!["alpha", "beta"]
25154 );
25155 assert_eq!(
25156 available_skill_ids(&skill_scope_agent(
25157 "states:\n initial: current\n states:\n current:\n prompt: current\n"
25158 )),
25159 vec!["alpha", "beta"]
25160 );
25161 assert!(
25162 available_skill_ids(&skill_scope_agent(
25163 "states:\n initial: current\n states:\n current:\n skills: [unknown]\n"
25164 ))
25165 .is_empty()
25166 );
25167 assert_eq!(
25168 available_skill_ids(&skill_scope_agent(
25169 "states:\n initial: parent\n states:\n parent:\n skills: [alpha]\n initial: child\n states:\n child:\n skills: []\n"
25170 )),
25171 vec!["alpha"]
25172 );
25173 assert_eq!(
25174 available_skill_ids(&skill_scope_agent(
25175 "states:\n initial: parent\n states:\n parent:\n skills: [alpha]\n initial: child\n states:\n child:\n skills: [beta]\n"
25176 )),
25177 vec!["alpha", "beta"]
25178 );
25179 assert_eq!(
25180 available_skill_ids(&skill_scope_agent(
25181 "states:\n initial: parent\n states:\n parent:\n skills: [alpha]\n initial: child\n states:\n child:\n inherit_parent: false\n skills: []\n"
25182 )),
25183 vec!["alpha", "beta"]
25184 );
25185 }
25186
25187 #[tokio::test]
25188 async fn state_skill_scope_reaches_the_router_candidate_prompt() {
25189 let (inherited, inherited_calls) = skill_scope_agent_with_router(
25190 "states:\n initial: parent\n states:\n parent:\n skills: [alpha]\n initial: child\n states:\n child:\n skills: []\n",
25191 );
25192 assert!(
25193 inherited
25194 .select_skill_candidate("route")
25195 .await
25196 .unwrap()
25197 .is_none()
25198 );
25199 let inherited_call = inherited_calls.last_call().unwrap();
25200 let inherited_prompt = &inherited_call.messages[0].content;
25201 assert!(inherited_prompt.contains("- alpha:"));
25202 assert!(!inherited_prompt.contains("- beta:"));
25203
25204 let (fallback_all, fallback_calls) = skill_scope_agent_with_router(
25205 "states:\n initial: current\n states:\n current:\n skills: []\n",
25206 );
25207 assert!(
25208 fallback_all
25209 .select_skill_candidate("route")
25210 .await
25211 .unwrap()
25212 .is_none()
25213 );
25214 let fallback_call = fallback_calls.last_call().unwrap();
25215 let fallback_prompt = &fallback_call.messages[0].content;
25216 assert!(fallback_prompt.contains("- alpha:"));
25217 assert!(fallback_prompt.contains("- beta:"));
25218
25219 let (omitted, omitted_calls) = skill_scope_agent_with_router(
25220 "states:\n initial: current\n states:\n current:\n prompt: current\n",
25221 );
25222 assert!(
25223 omitted
25224 .select_skill_candidate("route")
25225 .await
25226 .unwrap()
25227 .is_none()
25228 );
25229 let omitted_call = omitted_calls.last_call().unwrap();
25230 let omitted_prompt = &omitted_call.messages[0].content;
25231 assert!(omitted_prompt.contains("- alpha:"));
25232 assert!(omitted_prompt.contains("- beta:"));
25233
25234 let (unknown, unknown_calls) = skill_scope_agent_with_router(
25235 "states:\n initial: current\n states:\n current:\n skills: [unknown]\n",
25236 );
25237 assert!(
25238 unknown
25239 .select_skill_candidate("route")
25240 .await
25241 .unwrap()
25242 .is_none()
25243 );
25244 assert_eq!(unknown_calls.call_count(), 0);
25245 }
25246
25247 #[tokio::test]
25248 async fn parity_skill_route() {
25249 let yaml = skills_with_parallel_transition_yaml(
25250 " max_speculative_llm_calls_per_turn: 2\n max_parallel_runtime_tasks: 2",
25251 "",
25252 );
25253 let build = || {
25254 build_skills_beside_transition_agent(
25255 &yaml,
25256 role_mocks(
25257 mock_with_response("Draft response"),
25258 mock_with_response("helper"),
25259 ),
25260 )
25261 };
25262 assert_blocking_streaming_parity(build, "please use helper").await;
25263 }
25264
25265 fn required_context_agent(mock: MockLLMProvider, default: bool) -> RuntimeAgent {
25267 let default_yaml = if default {
25268 " default:\n brief: fallback\n"
25269 } else {
25270 ""
25271 };
25272 let yaml = format!(
25273 "name: RequiredContextAgent\nsystem_prompt: 'Voice: {{{{ context.voice.brief }}}}'\ncontext:\n voice:\n type: runtime\n required: true\n{default_yaml}"
25274 );
25275 AgentBuilder::from_yaml(&yaml)
25276 .unwrap()
25277 .llm(Arc::new(mock))
25278 .build()
25279 .unwrap()
25280 }
25281
25282 struct CountingContextProvider {
25283 marker: &'static str,
25284 calls: std::sync::atomic::AtomicUsize,
25285 }
25286
25287 #[async_trait]
25288 impl ContextProvider for CountingContextProvider {
25289 async fn get(&self, _key: &str, _current_context: &Value) -> Result<Value> {
25290 let call = self.calls.fetch_add(1, Ordering::SeqCst) + 1;
25291 Ok(serde_json::json!({"call": call, "marker": self.marker}))
25292 }
25293 }
25294
25295 struct FailOnceContextProvider {
25296 attempts: std::sync::atomic::AtomicUsize,
25297 }
25298
25299 #[async_trait]
25300 impl ContextProvider for FailOnceContextProvider {
25301 async fn get(&self, _key: &str, _current_context: &Value) -> Result<Value> {
25302 if self.attempts.fetch_add(1, Ordering::SeqCst) == 0 {
25303 return Err(AgentError::Other("context initialization failed".into()));
25304 }
25305 Ok(serde_json::json!({"brief": "ready"}))
25306 }
25307 }
25308
25309 fn session_context_agent(
25310 provider: Arc<CountingContextProvider>,
25311 refresh: &str,
25312 ) -> RuntimeAgent {
25313 let yaml = format!(
25314 "name: SessionContextAgent\nsystem_prompt: 'Call: {{{{ context.session_data.call }}}}'\ncontext:\n session_data:\n type: callback\n name: counter\n refresh: {refresh}\n"
25315 );
25316 let agent = AgentBuilder::from_yaml(&yaml)
25317 .unwrap()
25318 .llm(Arc::new(mock_with_response("ok")))
25319 .build()
25320 .unwrap();
25321 agent.register_context_provider("counter", provider);
25322 agent
25323 }
25324
25325 #[tokio::test]
25326 async fn session_context_characterizes_reset_and_restore_lifecycle() {
25327 let original_provider = Arc::new(CountingContextProvider {
25328 marker: "original",
25329 calls: std::sync::atomic::AtomicUsize::new(0),
25330 });
25331 let original = session_context_agent(original_provider.clone(), "per_session");
25332 original.chat("first").await.unwrap();
25333 original.chat("second").await.unwrap();
25334 assert_eq!(original_provider.calls.load(Ordering::SeqCst), 1);
25335 assert_eq!(original.get_context()["session_data"]["call"], 1);
25336 assert_eq!(original.get_context()["session_data"]["marker"], "original");
25337
25338 original.reset().await.unwrap();
25339 original.chat("after reset").await.unwrap();
25340 assert_eq!(original_provider.calls.load(Ordering::SeqCst), 1);
25341 assert_eq!(original.get_context()["session_data"]["call"], 1);
25342 assert_eq!(original.get_context()["session_data"]["marker"], "original");
25343 let snapshot = original.save_state().await.unwrap();
25344
25345 let fresh_provider = Arc::new(CountingContextProvider {
25346 marker: "fresh",
25347 calls: std::sync::atomic::AtomicUsize::new(0),
25348 });
25349 let fresh = session_context_agent(fresh_provider.clone(), "per_session");
25350 fresh.restore_state(snapshot.clone()).await.unwrap();
25351 fresh.chat("fresh restore").await.unwrap();
25352 assert_eq!(fresh_provider.calls.load(Ordering::SeqCst), 1);
25353 assert_eq!(fresh.get_context()["session_data"]["call"], 1);
25354 assert_eq!(fresh.get_context()["session_data"]["marker"], "fresh");
25355
25356 let warm_provider = Arc::new(CountingContextProvider {
25357 marker: "warm",
25358 calls: std::sync::atomic::AtomicUsize::new(0),
25359 });
25360 let warm = session_context_agent(warm_provider.clone(), "per_session");
25361 warm.chat("warmup").await.unwrap();
25362 assert_eq!(warm_provider.calls.load(Ordering::SeqCst), 1);
25363 warm.restore_state(snapshot).await.unwrap();
25364 warm.chat("warm restore").await.unwrap();
25365 assert_eq!(warm_provider.calls.load(Ordering::SeqCst), 1);
25366 assert_eq!(warm.get_context()["session_data"]["call"], 1);
25367 assert_eq!(warm.get_context()["session_data"]["marker"], "original");
25368 }
25369
25370 #[tokio::test]
25371 async fn once_context_is_not_refreshed_by_later_turns_or_reset() {
25372 let provider = Arc::new(CountingContextProvider {
25373 marker: "once",
25374 calls: std::sync::atomic::AtomicUsize::new(0),
25375 });
25376 let agent = session_context_agent(provider.clone(), "once");
25377
25378 agent.chat("first").await.unwrap();
25379 agent.chat("second").await.unwrap();
25380 agent.reset().await.unwrap();
25381 agent.chat("after reset").await.unwrap();
25382
25383 assert_eq!(provider.calls.load(Ordering::SeqCst), 1);
25384 assert_eq!(agent.get_context()["session_data"]["marker"], "once");
25385 }
25386
25387 #[tokio::test]
25388 async fn per_turn_context_refreshes_after_initialization_and_restore() {
25389 let original_provider = Arc::new(CountingContextProvider {
25390 marker: "original-per-turn",
25391 calls: std::sync::atomic::AtomicUsize::new(0),
25392 });
25393 let original = session_context_agent(original_provider.clone(), "per_turn");
25394 original.chat("first").await.unwrap();
25395 assert_eq!(original_provider.calls.load(Ordering::SeqCst), 2);
25396 assert_eq!(original.get_context()["session_data"]["call"], 2);
25397 original.chat("second").await.unwrap();
25398 assert_eq!(original_provider.calls.load(Ordering::SeqCst), 3);
25399 original.reset().await.unwrap();
25400 original.chat("after reset").await.unwrap();
25401 assert_eq!(original_provider.calls.load(Ordering::SeqCst), 4);
25402 let snapshot = original.save_state().await.unwrap();
25403
25404 let fresh_provider = Arc::new(CountingContextProvider {
25405 marker: "fresh-per-turn",
25406 calls: std::sync::atomic::AtomicUsize::new(0),
25407 });
25408 let fresh = session_context_agent(fresh_provider.clone(), "per_turn");
25409 fresh.restore_state(snapshot.clone()).await.unwrap();
25410 fresh.chat("fresh restore").await.unwrap();
25411 assert_eq!(fresh_provider.calls.load(Ordering::SeqCst), 2);
25412 assert_eq!(
25413 fresh.get_context()["session_data"]["marker"],
25414 "fresh-per-turn"
25415 );
25416 assert_eq!(fresh.get_context()["session_data"]["call"], 2);
25417
25418 let warm_provider = Arc::new(CountingContextProvider {
25419 marker: "warm-per-turn",
25420 calls: std::sync::atomic::AtomicUsize::new(0),
25421 });
25422 let warm = session_context_agent(warm_provider.clone(), "per_turn");
25423 warm.chat("warmup").await.unwrap();
25424 assert_eq!(warm_provider.calls.load(Ordering::SeqCst), 2);
25425 warm.restore_state(snapshot).await.unwrap();
25426 warm.chat("warm restore").await.unwrap();
25427 assert_eq!(warm_provider.calls.load(Ordering::SeqCst), 3);
25428 assert_eq!(
25429 warm.get_context()["session_data"]["marker"],
25430 "warm-per-turn"
25431 );
25432 assert_eq!(warm.get_context()["session_data"]["call"], 3);
25433 }
25434
25435 #[tokio::test]
25436 async fn test_context_initialization_retries_after_failure() {
25437 let mock = mock_with_response("Voice response");
25438 let calls = mock.clone();
25439 let yaml = "name: CallbackAgent\nsystem_prompt: 'Voice: {{ context.voice.brief }}'\ncontext:\n voice:\n type: callback\n name: flaky\n";
25440 let agent = AgentBuilder::from_yaml(yaml)
25441 .unwrap()
25442 .llm(Arc::new(mock))
25443 .build()
25444 .unwrap();
25445 let provider = Arc::new(FailOnceContextProvider {
25446 attempts: std::sync::atomic::AtomicUsize::new(0),
25447 });
25448 agent.register_context_provider("flaky", provider.clone());
25449
25450 assert!(agent.chat("first").await.is_err());
25451 assert_eq!(calls.call_count(), 0);
25452 assert_eq!(
25453 agent.chat("second").await.unwrap().content,
25454 "Voice response"
25455 );
25456 assert_eq!(provider.attempts.load(Ordering::SeqCst), 2);
25457 assert_eq!(calls.call_count(), 1);
25458 }
25459
25460 #[tokio::test]
25461 async fn test_required_context_blocks_chat_until_supplied_and_after_removal() {
25462 let mock = mock_with_response("Voice response");
25463 let calls = mock.clone();
25464 let agent = required_context_agent(mock, false);
25465
25466 let error = agent.chat("first").await.unwrap_err();
25467 assert!(
25468 error
25469 .to_string()
25470 .contains("Required context 'voice' not provided")
25471 );
25472 assert_eq!(calls.call_count(), 0);
25473 assert!(agent.memory.get_messages(None).await.unwrap().is_empty());
25474
25475 agent
25476 .set_context("voice.brief", serde_json::json!("ready"))
25477 .unwrap();
25478 assert_eq!(
25479 agent.chat("second").await.unwrap().content,
25480 "Voice response"
25481 );
25482 assert_eq!(calls.call_count(), 1);
25483
25484 agent.remove_context("voice");
25485 let error = agent.chat("third").await.unwrap_err();
25486 assert!(
25487 error
25488 .to_string()
25489 .contains("Required context 'voice' not provided")
25490 );
25491 assert_eq!(calls.call_count(), 1);
25492 }
25493
25494 #[tokio::test]
25495 async fn test_required_context_default_satisfies_presence_check() {
25496 let mock = mock_with_response("Fallback response");
25497 let calls = mock.clone();
25498 let agent = required_context_agent(mock, true);
25499
25500 assert_eq!(
25501 agent.chat("hello").await.unwrap().content,
25502 "Fallback response"
25503 );
25504 assert_eq!(
25505 agent.context_manager().get_path("voice.brief"),
25506 Some(serde_json::json!("fallback"))
25507 );
25508 assert_eq!(calls.call_count(), 1);
25509 }
25510
25511 #[tokio::test]
25512 async fn test_required_context_blocks_legacy_stream_before_model_call() {
25513 use futures::StreamExt;
25514
25515 let mock = mock_with_response("Voice response");
25516 let calls = mock.clone();
25517 let agent = required_context_agent(mock, false);
25518 let mut stream = agent.chat_stream("first").await.unwrap();
25519 assert!(
25520 matches!(stream.next().await, Some(StreamChunk::Error { message }) if message.contains("Required context 'voice' not provided"))
25521 );
25522 assert!(stream.next().await.is_none());
25523 drop(stream);
25524 assert_eq!(calls.call_count(), 0);
25525 assert!(agent.memory.get_messages(None).await.unwrap().is_empty());
25526
25527 agent
25528 .set_context("voice.brief", serde_json::json!("ready"))
25529 .unwrap();
25530 assert_eq!(
25531 agent.chat("second").await.unwrap().content,
25532 "Voice response"
25533 );
25534 }
25535
25536 #[tokio::test]
25537 async fn test_required_context_blocks_event_streams_without_final() {
25538 use futures::StreamExt;
25539
25540 for actor_scoped in [false, true] {
25541 let mock = mock_with_response("Voice response");
25542 let calls = mock.clone();
25543 let agent = required_context_agent(mock, false);
25544 let mut events = if actor_scoped {
25545 agent
25546 .chat_stream_events_with_actor_context(
25547 "first",
25548 crate::TurnActorContext::new().with_origin_actor("caller"),
25549 )
25550 .await
25551 .unwrap()
25552 } else {
25553 agent.chat_stream_events("first").await.unwrap()
25554 };
25555 assert!(
25556 matches!(events.next().await, Some(AgentStreamEvent::Chunk(StreamChunk::Error { message })) if message.contains("Required context 'voice' not provided"))
25557 );
25558 assert!(events.next().await.is_none());
25559 drop(events);
25560 assert_eq!(calls.call_count(), 0);
25561 assert!(agent.memory.get_messages(None).await.unwrap().is_empty());
25562
25563 agent
25564 .set_context("voice.brief", serde_json::json!("ready"))
25565 .unwrap();
25566 assert_eq!(
25567 agent.chat("second").await.unwrap().content,
25568 "Voice response"
25569 );
25570 }
25571 }
25572}