Skip to main content

behest_runtime/
agent.rs

1//! Agent runtime — streaming-first execution kernel.
2//!
3//! [`AgentRuntime`] orchestrates the full agent loop: context building,
4//! model invocation (streaming with non-streaming fallback), tool execution,
5//! session persistence, and event emission.
6
7use std::sync::Arc;
8
9#[cfg(feature = "queue")]
10use crate::event_publisher::EventPublisher;
11
12use chrono::Utc;
13use tokio::sync::broadcast;
14use tracing::{debug, error, warn};
15use uuid::Uuid;
16
17use behest_provider::{FinishReason, Message, TokenUsage};
18
19use super::compaction::{CompactionCircuitBreaker, CompactionService};
20use super::context::ContextPipeline;
21use super::doom_loop::DoomLoopDetector;
22use super::error::{RuntimeError, RuntimeResult};
23use super::event::{AgentEvent, RunStarted};
24use super::extensions::Extensions;
25use super::input::{InputAdmission, InputRecord};
26use super::policy::RuntimePolicy;
27use super::run::{RunId, RunRecord, RunRequest, RunStatus};
28use super::session_gate::SessionGate;
29use super::snapshot::{Snapshot, SnapshotStore};
30use super::store::{RunStore, RuntimeStore};
31use super::tool_runtime::ToolRuntime;
32use super::tool_scope::ScopeGuard;
33use super::turn::{TurnState, TurnTransition};
34use behest_tool::ToolRegistry;
35
36/// Streaming-first agent runtime kernel.
37///
38/// Ties together provider registry, context pipeline, tool runtime,
39/// compaction service, persistent stores, snapshot recovery, session
40/// gating, and input admission into a complete agent execution loop.
41pub struct AgentRuntime {
42    providers: behest_provider::ProviderRegistry,
43    pub(super) context: ContextPipeline,
44    pub(super) tools: ToolRuntime,
45    pub(super) store: Arc<RuntimeStore>,
46    pub(super) policy: RuntimePolicy,
47    pub(super) compaction: CompactionService,
48    session_gate: SessionGate,
49    input_admission: InputAdmission,
50    pub(super) event_tx: broadcast::Sender<AgentEvent>,
51    #[cfg(feature = "queue")]
52    pub(super) event_publisher: Option<Arc<dyn EventPublisher>>,
53    snapshot_store: Option<Arc<dyn SnapshotStore>>,
54    /// Composable, hot-pluggable facade over every pluggable runtime
55    /// element. Constructed empty by [`AgentRuntime::new`] and
56    /// populated as the operator configures the runtime via
57    /// `with_*` setters.
58    pub(super) extensions: Arc<Extensions>,
59}
60
61impl AgentRuntime {
62    /// Creates a new agent runtime from an [`Extensions`] facade and a
63    /// [`RuntimePolicy`].
64    ///
65    /// Internally constructs a [`ProviderRegistry`](behest_provider::ProviderRegistry),
66    /// [`RuntimeStore`], [`ContextPipeline`], and [`ToolRuntime`] from the
67    /// extensions or their defaults. Use
68    /// [`with_tool_registry`](Self::with_tool_registry)
69    /// to populate the tool runtime, and the `with_*` setters for
70    /// optional components (event publisher, snapshot store).
71    #[must_use]
72    pub fn new(extensions: Arc<Extensions>, policy: RuntimePolicy) -> Self {
73        let mut providers = behest_provider::ProviderRegistry::new();
74        for (name, provider) in extensions.chat_providers.snapshot() {
75            let _ = name;
76            providers.register_chat_arc(provider);
77        }
78        for (name, provider) in extensions.embedding_providers.snapshot() {
79            let _ = name;
80            providers.register_embedding_arc(provider);
81        }
82        let store = Arc::new(RuntimeStore::from_extensions(&extensions));
83        let context = ContextPipeline::new();
84        let tools = ToolRuntime::new(ToolRegistry::new(), policy.clone());
85        let (event_tx, _) = broadcast::channel(256);
86        let compaction = CompactionService::new(providers.clone(), policy.compaction.clone());
87        let input_admission = InputAdmission::new(policy.input_admission.clone());
88        Self {
89            providers,
90            context,
91            tools,
92            store,
93            policy,
94            compaction,
95            session_gate: SessionGate::new(),
96            input_admission,
97            event_tx,
98            #[cfg(feature = "queue")]
99            event_publisher: None,
100            snapshot_store: None,
101            extensions,
102        }
103    }
104
105    /// Replaces the tool runtime with one backed by the given registry.
106    #[must_use]
107    pub fn with_tool_registry(mut self, registry: ToolRegistry) -> Self {
108        self.tools = ToolRuntime::new(registry, self.policy.clone());
109        self
110    }
111
112    /// Sets an external event publisher for the agent runtime.
113    ///
114    /// When set, every [`AgentEvent`] emitted during a run will also be
115    /// published to the configured [`EventPublisher`] via fire-and-forget.
116    /// Requires the `queue` feature.
117    #[cfg(feature = "queue")]
118    #[must_use]
119    pub fn with_event_publisher(mut self, publisher: Arc<dyn EventPublisher>) -> Self {
120        // Mirror into the Extensions facade for unified access.
121        let _ = self
122            .extensions
123            .event_publishers
124            .register_or_replace("default", Arc::clone(&publisher));
125        self.event_publisher = Some(publisher);
126        self
127    }
128
129    /// Sets an optional snapshot store for FSM run recovery.
130    ///
131    /// When configured, intermediate run state is persisted after each
132    /// turn, allowing crashed runs to be resumed via [`resume`](Self::resume).
133    #[must_use]
134    pub fn with_snapshot_store(mut self, snapshot_store: Arc<dyn SnapshotStore>) -> Self {
135        let _ = self
136            .extensions
137            .snapshot_stores
138            .register_or_replace("default", Arc::clone(&snapshot_store));
139        self.snapshot_store = Some(snapshot_store);
140        self
141    }
142
143    /// Returns the composable [`Extensions`] facade.
144    ///
145    /// Use this to register, replace, or look up providers, tools,
146    /// context adapters, stores, and other pluggable elements by name.
147    /// The facade is shared with the runtime's internal references, so
148    /// registering here is equivalent to using the `with_*` setters.
149    #[must_use]
150    pub fn extensions(&self) -> &Arc<Extensions> {
151        &self.extensions
152    }
153
154    /// Returns the session gate used to serialize concurrent runs per session.
155    #[must_use]
156    pub fn session_gate(&self) -> &SessionGate {
157        &self.session_gate
158    }
159
160    /// Subscribes to runtime events via a tokio broadcast receiver.
161    ///
162    /// The receiver starts with the oldest event still in the buffer
163    /// (capacity 256). Slow consumers that fall behind will miss events.
164    #[must_use]
165    pub fn subscribe(&self) -> broadcast::Receiver<AgentEvent> {
166        self.event_tx.subscribe()
167    }
168
169    /// Returns the runtime policy.
170    #[must_use]
171    pub fn policy(&self) -> &RuntimePolicy {
172        &self.policy
173    }
174
175    /// Returns the tool runtime.
176    #[must_use]
177    pub fn tools(&self) -> &ToolRuntime {
178        &self.tools
179    }
180
181    /// Returns the provider registry.
182    #[must_use]
183    pub fn providers(&self) -> &behest_provider::ProviderRegistry {
184        &self.providers
185    }
186
187    /// Returns the context pipeline.
188    #[must_use]
189    pub fn context(&self) -> &ContextPipeline {
190        &self.context
191    }
192
193    /// Returns the runtime store.
194    #[must_use]
195    pub fn store(&self) -> &Arc<RuntimeStore> {
196        &self.store
197    }
198
199    /// Returns the compaction service.
200    #[must_use]
201    pub fn compaction(&self) -> &CompactionService {
202        &self.compaction
203    }
204
205    /// Returns the snapshot store, if configured.
206    #[must_use]
207    pub fn snapshot_store(&self) -> Option<&Arc<dyn SnapshotStore>> {
208        self.snapshot_store.as_ref()
209    }
210
211    /// Returns the session store.
212    #[must_use]
213    pub fn sessions(&self) -> &dyn behest_store::SessionStore {
214        self.store.sessions()
215    }
216
217    /// Returns the execution store.
218    #[must_use]
219    pub fn executions(&self) -> &dyn behest_store::ExecutionStore {
220        self.store.executions()
221    }
222
223    /// Returns the run store.
224    #[must_use]
225    pub fn runs(&self) -> &dyn RunStore {
226        self.store.runs()
227    }
228
229    /// Returns the embedding store, if configured.
230    #[must_use]
231    pub fn embeddings(&self) -> Option<&dyn behest_store::EmbeddingStore> {
232        self.store.embeddings()
233    }
234
235    /// Returns the artifact store, if configured.
236    #[must_use]
237    pub fn artifacts(&self) -> Option<&dyn behest_store::ArtifactStore> {
238        self.store.artifacts()
239    }
240
241    /// Executes an agent run to completion.
242    ///
243    /// The run loop:
244    /// 1. Creates or loads a session
245    /// 2. Persists the user message
246    /// 3. Iterates: build context → call model → persist response → execute tools → repeat
247    /// 4. Returns the final run ID and finish reason
248    ///
249    /// # Errors
250    ///
251    /// Returns `RuntimeError` on provider, store, context, or policy violations.
252    #[allow(clippy::too_many_lines)]
253    pub async fn run(&self, request: RunRequest) -> RuntimeResult<RunOutput> {
254        let run_id = request.run_id.unwrap_or_default();
255        let session_id = self.store.ensure_session(request.session_id).await?;
256
257        // Acquire per-session lock — prevents concurrent runs from
258        // interleaving writes to the same session.
259        let _session_guard = self
260            .session_gate
261            .acquire(session_id)
262            .await
263            .map_err(|busy| RuntimeError::SessionBusy(busy.session_id))?;
264
265        // Admit the input before allocating any run resources.
266        let mut input_record = InputRecord::new(session_id, request.input.clone());
267        let admission_events = self
268            .input_admission
269            .admit(&mut input_record)
270            .map_err(|e| RuntimeError::InputAdmissionFailed(e.to_string()))?;
271        if input_record.state == super::input::InputState::Rejected {
272            let reason = input_record.rejection_reason.clone().unwrap_or_default();
273            return Err(RuntimeError::InputRejected {
274                input_id: input_record.id,
275                reason,
276            });
277        }
278        debug!(
279            input_id = %input_record.id,
280            events = admission_events.len(),
281            "input admitted"
282        );
283
284        // Push a Run-level tool scope. The RAII guard ensures cleanup
285        // on every exit path, including early returns and panics.
286        let _run_scope: ScopeGuard = self.tools.registry().push_scope_guarded();
287
288        let run_record = RunRecord::new(
289            run_id,
290            session_id,
291            request.provider.clone(),
292            request.model.clone(),
293            request.metadata.clone(),
294            request.client_request_id.clone(),
295        );
296        self.store.runs().create_run(run_record).await?;
297
298        // Create doom loop detector for this run.
299        let mut doom_detector = DoomLoopDetector::new(self.policy.doom_loop.clone());
300
301        // Create compaction circuit breaker for this run.
302        let mut compaction_breaker =
303            CompactionCircuitBreaker::new(self.policy.compaction.circuit_breaker_threshold);
304
305        self.emit(&AgentEvent::RunStarted(RunStarted {
306            run_id,
307            session_id,
308            provider: request.provider.clone(),
309            model: request.model.clone(),
310            timestamp: Utc::now(),
311        }));
312        self.update_status(run_id, RunStatus::SessionLoaded).await?;
313
314        let user_message = Message::user_text(&request.input);
315        let user_msg_id = self.store.append_message(session_id, &user_message).await?;
316        debug!(%run_id, %user_msg_id, "user message persisted");
317
318        let provider = self
319            .providers
320            .chat(&request.provider)
321            .ok_or_else(|| RuntimeError::ProviderNotFound(request.provider.to_string()))?;
322
323        let tool_specs = self.tools.registry().specs();
324        let has_tools = !tool_specs.is_empty();
325
326        self.run_loop(
327            run_id,
328            session_id,
329            provider,
330            request,
331            tool_specs,
332            has_tools,
333            0,
334            TokenUsage::new(0, 0),
335            None,
336            None,
337            None,
338            TurnState::CheckingPolicy,
339            &mut doom_detector,
340            &mut compaction_breaker,
341            0,
342        )
343        .await
344    }
345
346    /// Resumes a crashed or halted agent run from its last saved snapshot.
347    ///
348    /// Loads the snapshot by `run_id`, re-acquires the session lock, pushes
349    /// a Run-level tool scope, and restarts the turn loop from the saved state.
350    ///
351    /// # Errors
352    ///
353    /// Returns `RuntimeError` if the snapshot is not found, the session is busy,
354    /// or resuming fails.
355    pub async fn resume(&self, run_id: RunId) -> RuntimeResult<RunOutput> {
356        let snapshot_store = self.snapshot_store.as_ref().ok_or_else(|| {
357            RuntimeError::RecoveryFailed("snapshot store not configured".to_string())
358        })?;
359
360        let snapshot = snapshot_store
361            .load(run_id)
362            .await?
363            .ok_or_else(|| RuntimeError::RunNotFound(run_id))?;
364
365        // Re-acquire per-session lock
366        let _session_guard = self
367            .session_gate
368            .acquire(snapshot.session_id)
369            .await
370            .map_err(|busy| RuntimeError::SessionBusy(busy.session_id))?;
371
372        // Re-push a Run-level tool scope
373        let _run_scope: ScopeGuard = self.tools.registry().push_scope_guarded();
374
375        let provider = self
376            .providers
377            .chat(&snapshot.request.provider)
378            .ok_or_else(|| RuntimeError::ProviderNotFound(snapshot.request.provider.to_string()))?;
379
380        let tool_specs = self.tools.registry().specs();
381        let has_tools = !tool_specs.is_empty();
382
383        // Resume the run in the database/store status as well
384        self.update_status(run_id, TurnTransition::status_for(snapshot.current_state))
385            .await?;
386
387        let mut doom_detector = DoomLoopDetector::new(self.policy.doom_loop.clone());
388
389        let mut compaction_breaker =
390            CompactionCircuitBreaker::new(self.policy.compaction.circuit_breaker_threshold);
391
392        self.run_loop(
393            run_id,
394            snapshot.session_id,
395            provider,
396            snapshot.request,
397            tool_specs,
398            has_tools,
399            snapshot.iteration,
400            snapshot.total_usage,
401            snapshot.last_finish,
402            snapshot.assistant_message,
403            snapshot.assistant_msg_id,
404            snapshot.current_state,
405            &mut doom_detector,
406            &mut compaction_breaker,
407            snapshot.output_recovery_count,
408        )
409        .await
410    }
411
412    #[allow(clippy::too_many_arguments)]
413    /// Captures the current run state as a snapshot for crash recovery.
414    ///
415    /// A no-op when no snapshot store is configured.
416    ///
417    /// # Errors
418    ///
419    /// Returns [`RuntimeError::Storage`] on snapshot persistence failure.
420    pub(super) async fn save_snapshot_helper(
421        &self,
422        run_id: RunId,
423        session_id: Uuid,
424        iteration: usize,
425        state: TurnState,
426        total_usage: TokenUsage,
427        last_finish: Option<&FinishReason>,
428        assistant_message: Option<&Message>,
429        assistant_msg_id: Option<Uuid>,
430        request: &RunRequest,
431        output_recovery_count: u32,
432    ) -> RuntimeResult<()> {
433        if let Some(store) = &self.snapshot_store {
434            let snapshot = Snapshot {
435                run_id,
436                session_id,
437                status: TurnTransition::status_for(state),
438                iteration,
439                current_state: state,
440                total_usage,
441                last_finish: last_finish.cloned(),
442                assistant_message: assistant_message.cloned(),
443                assistant_msg_id,
444                request: request.clone(),
445                output_recovery_count,
446                timestamp: Utc::now(),
447            };
448            store.save(&snapshot).await?;
449        }
450        Ok(())
451    }
452
453    /// Deletes the snapshot for a completed or failed run.
454    ///
455    /// A no-op when no snapshot store is configured.
456    ///
457    /// # Errors
458    ///
459    /// Returns [`RuntimeError::Storage`] on snapshot deletion failure.
460    pub(super) async fn delete_snapshot_helper(&self, run_id: RunId) -> RuntimeResult<()> {
461        if let Some(store) = &self.snapshot_store {
462            store.delete(run_id).await?;
463        }
464        Ok(())
465    }
466
467    /// Emits an [`AgentEvent`] through all available delivery channels.
468    ///
469    /// Events are broadcast to local subscribers via the internal tokio
470    /// broadcast channel. If background jobs are configured, the event is
471    /// also scheduled for persistent storage. With the `queue` feature,
472    /// events are additionally forwarded to the external event publisher.
473    pub(super) fn emit(&self, event: &AgentEvent) {
474        if let Err(e) = self.event_tx.send(event.clone()) {
475            warn!(lag = ?e, "event channel full, consumer too slow — event dropped");
476        }
477        #[cfg(feature = "queue")]
478        if let Some(publisher) = &self.event_publisher {
479            let publisher = Arc::clone(publisher);
480            let event = event.clone();
481            tokio::spawn(async move {
482                if let Err(e) = publisher.publish(event).await {
483                    warn!(error = %e, "failed to publish runtime event");
484                }
485            });
486        }
487    }
488
489    /// Emits a [`CacheMetrics`](super::event::CacheMetrics) event when the
490    /// provider reported any cache-related token usage, and persists the
491    /// event to all registered runtime event stores for replay.
492    ///
493    /// No-op when all cache fields are `None` or `0`, or when no event
494    /// store is registered.
495    pub(super) async fn emit_cache_metrics(
496        &self,
497        run_id: RunId,
498        usage: &behest_core::message::TokenUsage,
499    ) {
500        let creation = usage.cache_creation_input_tokens.unwrap_or(0);
501        let read = usage.cache_read_input_tokens.unwrap_or(0);
502        let cached = usage.cached_input_tokens.unwrap_or(0);
503        if creation == 0 && read == 0 && cached == 0 {
504            return;
505        }
506        let event = super::event::CacheMetrics {
507            run_id,
508            cache_creation_input_tokens: creation,
509            cache_read_input_tokens: read,
510            cached_input_tokens: cached,
511            timestamp: chrono::Utc::now(),
512        };
513        self.emit(&AgentEvent::CacheMetrics(event.clone()));
514        for (_name, store) in self.extensions.runtime_event_stores.snapshot() {
515            if let Err(e) = store.append(AgentEvent::CacheMetrics(event.clone())).await {
516                warn!(error = %e, "failed to persist cache metrics to event store");
517            }
518        }
519    }
520
521    /// Persists a run status update to the underlying run store.
522    ///
523    /// # Errors
524    ///
525    /// Returns [`RuntimeError::Storage`] on persistence failure.
526    pub(super) async fn update_status(
527        &self,
528        run_id: RunId,
529        status: RunStatus,
530    ) -> RuntimeResult<()> {
531        self.store.runs().update_run_status(run_id, status).await
532    }
533
534    /// Marks a run as failed in the store and emits a terminal [`RunFailed`](super::event::RunFailed) event.
535    ///
536    /// The error message is logged. Status update errors are logged but
537    /// not propagated, making this safe to call from cleanup paths.
538    pub(super) async fn fail_run(&self, run_id: RunId, err: &RuntimeError) {
539        let error_msg = err.to_string();
540        error!(%run_id, error = %error_msg, "run failed");
541        let _ = self.update_status(run_id, RunStatus::Failed).await;
542        self.emit(&AgentEvent::RunFailed(super::event::RunFailed {
543            run_id,
544            error: error_msg,
545            timestamp: Utc::now(),
546        }));
547    }
548}
549
550/// Output of a completed agent run.
551#[derive(Debug, Clone)]
552pub struct RunOutput {
553    /// Run identifier.
554    pub run_id: RunId,
555    /// Session identifier.
556    pub session_id: Uuid,
557    /// Number of model call iterations.
558    pub iterations: usize,
559    /// Final finish reason.
560    pub finish_reason: FinishReason,
561    /// Aggregated token usage across all iterations.
562    pub total_usage: TokenUsage,
563}
564
565#[cfg(test)]
566#[allow(clippy::unwrap_used, clippy::expect_used)]
567mod tests {
568    use super::*;
569    use crate::memory::MemoryRunStore;
570    use crate::snapshot::{FileSnapshotStore, Snapshot};
571    use async_trait::async_trait;
572    use behest_provider::{
573        ChatProvider, ChatRequest, ChatResponse, ChatStream, ChatStreamEvent, ModelName,
574        ProviderCapabilities, ProviderId, ProviderResult, ToolCall,
575    };
576    use behest_store::memory::{MemoryExecutionStore, MemorySessionStore};
577    use behest_tool::{FunctionTool, ToolRegistry};
578    use futures_util::StreamExt as _;
579    use serde_json::json;
580    use std::time::Duration;
581
582    struct MockProvider {
583        responses: std::sync::Mutex<Vec<ChatResponse>>,
584    }
585
586    impl MockProvider {
587        fn new(responses: Vec<ChatResponse>) -> Self {
588            Self {
589                responses: std::sync::Mutex::new(responses),
590            }
591        }
592
593        fn text_response(text: &str) -> ChatResponse {
594            ChatResponse {
595                provider: ProviderId::new("mock"),
596                model: ModelName::new("test"),
597                message: Message::assistant_text(text),
598                finish_reason: FinishReason::Stop,
599                usage: Some(TokenUsage::new(10, 20)),
600                raw: None,
601            }
602        }
603
604        fn tool_call_response(
605            call_id: &str,
606            tool_name: &str,
607            args: serde_json::Value,
608        ) -> ChatResponse {
609            ChatResponse {
610                provider: ProviderId::new("mock"),
611                model: ModelName::new("test"),
612                message: Message::Assistant {
613                    content: vec![],
614                    tool_calls: vec![ToolCall::new(call_id, tool_name, args)],
615                },
616                finish_reason: FinishReason::ToolCalls,
617                usage: Some(TokenUsage::new(15, 25)),
618                raw: None,
619            }
620        }
621
622        fn length_response(text: &str) -> ChatResponse {
623            ChatResponse {
624                provider: ProviderId::new("mock"),
625                model: ModelName::new("test"),
626                message: Message::assistant_text(text),
627                finish_reason: FinishReason::Length,
628                usage: Some(TokenUsage::new(10, 20)),
629                raw: None,
630            }
631        }
632    }
633
634    struct IdleStreamProvider;
635
636    #[async_trait]
637    impl ChatProvider for IdleStreamProvider {
638        fn id(&self) -> ProviderId {
639            ProviderId::new("mock")
640        }
641
642        fn capabilities(&self) -> ProviderCapabilities {
643            ProviderCapabilities {
644                chat: true,
645                chat_stream: true,
646                ..ProviderCapabilities::empty()
647            }
648        }
649
650        async fn complete(&self, _request: ChatRequest) -> ProviderResult<ChatResponse> {
651            Ok(MockProvider::text_response("fallback"))
652        }
653
654        async fn stream(&self, request: ChatRequest) -> ProviderResult<ChatStream> {
655            let started = ChatStreamEvent::Started {
656                provider: ProviderId::new("mock"),
657                model: request.model,
658            };
659            let stream = futures_util::stream::once(async { Ok(started) })
660                .chain(futures_util::stream::pending());
661
662            Ok(Box::pin(stream))
663        }
664    }
665
666    #[async_trait]
667    impl ChatProvider for MockProvider {
668        fn id(&self) -> ProviderId {
669            ProviderId::new("mock")
670        }
671
672        fn capabilities(&self) -> ProviderCapabilities {
673            ProviderCapabilities::chat()
674        }
675
676        async fn complete(&self, _request: ChatRequest) -> ProviderResult<ChatResponse> {
677            let mut responses = self.responses.lock().unwrap();
678            if responses.is_empty() {
679                Ok(Self::text_response("no more responses"))
680            } else {
681                Ok(responses.remove(0))
682            }
683        }
684    }
685
686    fn make_runtime(provider: MockProvider, tools: ToolRegistry) -> AgentRuntime {
687        let exts = Extensions::new();
688        exts.chat_providers
689            .register_or_replace("mock", Arc::new(provider));
690
691        let sessions = MemorySessionStore::new();
692        let executions = MemoryExecutionStore::new();
693        let runs = MemoryRunStore::new();
694        exts.session_stores
695            .register_or_replace("default", Arc::new(sessions));
696        exts.execution_stores
697            .register_or_replace("default", Arc::new(executions));
698        exts.run_stores
699            .register_or_replace("default", Arc::new(runs));
700
701        let policy = RuntimePolicy::new().with_max_iterations(5);
702
703        AgentRuntime::new(Arc::new(exts), policy).with_tool_registry(tools)
704    }
705
706    fn make_runtime_from_provider(
707        provider: Arc<dyn ChatProvider>,
708        tools: ToolRegistry,
709        policy: RuntimePolicy,
710    ) -> AgentRuntime {
711        let exts = Extensions::new();
712        exts.chat_providers
713            .register_or_replace("mock", Arc::clone(&provider));
714
715        let sessions = MemorySessionStore::new();
716        let executions = MemoryExecutionStore::new();
717        let runs = MemoryRunStore::new();
718        exts.session_stores
719            .register_or_replace("default", Arc::new(sessions));
720        exts.execution_stores
721            .register_or_replace("default", Arc::new(executions));
722        exts.run_stores
723            .register_or_replace("default", Arc::new(runs));
724
725        AgentRuntime::new(Arc::new(exts), policy).with_tool_registry(tools)
726    }
727
728    #[tokio::test]
729    async fn run_should_complete_with_text_response() {
730        let provider = MockProvider::new(vec![MockProvider::text_response("Hello!")]);
731        let runtime = make_runtime(provider, ToolRegistry::new());
732
733        let request = RunRequest::new(ProviderId::new("mock"), ModelName::new("test"), "Hi there");
734
735        let output = runtime.run(request).await.unwrap();
736        assert_eq!(output.iterations, 1);
737        assert!(matches!(output.finish_reason, FinishReason::Stop));
738        assert_eq!(output.total_usage.input_tokens, 10);
739        assert_eq!(output.total_usage.output_tokens, 20);
740    }
741
742    #[tokio::test]
743    async fn run_should_execute_tools_and_loop() {
744        let provider = MockProvider::new(vec![
745            MockProvider::tool_call_response("call_1", "echo", json!({"message": "hello"})),
746            MockProvider::text_response("Done!"),
747        ]);
748
749        let tools = ToolRegistry::new();
750        tools.register(FunctionTool::new(
751            "echo",
752            "Echoes input",
753            json!({"type": "object", "properties": {"message": {"type": "string"}}}),
754            |args: serde_json::Value| -> std::pin::Pin<
755                Box<
756                    dyn std::future::Future<Output = behest_tool::ToolResult<serde_json::Value>>
757                        + Send,
758                >,
759            > {
760                Box::pin(async move {
761                    Ok(args
762                        .get("message")
763                        .cloned()
764                        .unwrap_or(serde_json::Value::Null))
765                })
766            },
767        ));
768
769        let runtime = make_runtime(provider, tools);
770
771        let request = RunRequest::new(
772            ProviderId::new("mock"),
773            ModelName::new("test"),
774            "Echo hello",
775        );
776
777        let output = runtime.run(request).await.unwrap();
778        assert_eq!(output.iterations, 2);
779        assert!(matches!(output.finish_reason, FinishReason::Stop));
780    }
781
782    #[tokio::test]
783    async fn run_should_respect_iteration_limit() {
784        let responses: Vec<ChatResponse> = (0..10)
785            .map(|i| {
786                MockProvider::tool_call_response(
787                    &format!("call_{i}"),
788                    "echo",
789                    json!({"message": format!("msg_{i}")}),
790                )
791            })
792            .collect();
793
794        let provider = MockProvider::new(responses);
795
796        let tools = ToolRegistry::new();
797        tools.register(FunctionTool::new(
798            "echo",
799            "Echoes",
800            json!({"type": "object"}),
801            |_args: serde_json::Value| -> std::pin::Pin<
802                Box<
803                    dyn std::future::Future<Output = behest_tool::ToolResult<serde_json::Value>>
804                        + Send,
805                >,
806            > { Box::pin(async move { Ok(json!("ok")) }) },
807        ));
808
809        let runtime = make_runtime(provider, tools);
810
811        let request = RunRequest::new(ProviderId::new("mock"), ModelName::new("test"), "loop");
812
813        let result = runtime.run(request).await;
814        assert!(result.is_err());
815        assert!(matches!(
816            result.unwrap_err(),
817            RuntimeError::IterationLimitExceeded(_)
818        ));
819    }
820
821    #[tokio::test]
822    async fn run_should_emit_events() {
823        let provider = MockProvider::new(vec![MockProvider::text_response("Hello!")]);
824        let runtime = make_runtime(provider, ToolRegistry::new());
825        let mut rx = runtime.subscribe();
826
827        let request = RunRequest::new(ProviderId::new("mock"), ModelName::new("test"), "Hi");
828
829        let _output = runtime.run(request).await.unwrap();
830
831        let mut events = Vec::new();
832        while let Ok(event) = rx.try_recv() {
833            events.push(event);
834        }
835
836        assert!(
837            events
838                .iter()
839                .any(|e| matches!(e, AgentEvent::RunStarted(_)))
840        );
841        assert!(
842            events
843                .iter()
844                .any(|e| matches!(e, AgentEvent::ContextBuilt(_)))
845        );
846        assert!(
847            events
848                .iter()
849                .any(|e| matches!(e, AgentEvent::ModelStarted(_)))
850        );
851        assert!(
852            events
853                .iter()
854                .any(|e| matches!(e, AgentEvent::RunCompleted(_)))
855        );
856    }
857
858    #[tokio::test]
859    async fn run_should_create_session_when_none_provided() {
860        let provider = MockProvider::new(vec![MockProvider::text_response("Hi")]);
861        let runtime = make_runtime(provider, ToolRegistry::new());
862
863        let request = RunRequest::new(ProviderId::new("mock"), ModelName::new("test"), "Hello");
864
865        let output = runtime.run(request).await.unwrap();
866        assert_ne!(output.session_id, Uuid::nil());
867    }
868
869    #[tokio::test]
870    async fn run_should_timeout_when_stream_stalls_between_events() {
871        let policy = RuntimePolicy::new()
872            .with_max_iterations(1)
873            .with_provider_timeout(Duration::from_millis(20));
874        let runtime =
875            make_runtime_from_provider(Arc::new(IdleStreamProvider), ToolRegistry::new(), policy);
876
877        let request = RunRequest::new(
878            ProviderId::new("mock"),
879            ModelName::new("test"),
880            "stall stream",
881        );
882        let result = tokio::time::timeout(Duration::from_millis(300), runtime.run(request))
883            .await
884            .expect("runtime should return provider timeout instead of hanging");
885
886        assert!(matches!(
887            result,
888            Err(RuntimeError::Provider(
889                behest_core::error::ProviderError::Timeout { .. }
890            ))
891        ));
892    }
893
894    #[tokio::test]
895    async fn run_should_fail_for_unknown_provider() {
896        let provider = MockProvider::new(vec![]);
897        let runtime = make_runtime(provider, ToolRegistry::new());
898
899        let request = RunRequest::new(
900            ProviderId::new("nonexistent"),
901            ModelName::new("test"),
902            "Hello",
903        );
904
905        let result = runtime.run(request).await;
906        assert!(result.is_err());
907        assert!(matches!(
908            result.unwrap_err(),
909            RuntimeError::ProviderNotFound(_)
910        ));
911    }
912
913    #[tokio::test]
914    async fn run_should_create_snapshots_and_resume_successfully() {
915        let temp_dir = tempfile::tempdir().unwrap();
916        let snapshot_store = Arc::new(FileSnapshotStore::new(temp_dir.path().to_path_buf()));
917
918        let provider = MockProvider::new(vec![
919            MockProvider::tool_call_response("call_rec", "echo", json!({"message": "rec"})),
920            MockProvider::text_response("Done after resume!"),
921        ]);
922
923        let tools = ToolRegistry::new();
924        tools.register(FunctionTool::new(
925            "echo",
926            "Echoes message",
927            json!({"type": "object"}),
928            |args: serde_json::Value| -> std::pin::Pin<
929                Box<
930                    dyn std::future::Future<Output = behest_tool::ToolResult<serde_json::Value>>
931                        + Send,
932                >,
933            > {
934                Box::pin(async move { Ok(args.get("message").cloned().unwrap_or_default()) })
935            },
936        ));
937
938        let runtime = make_runtime(provider, tools).with_snapshot_store(snapshot_store.clone());
939
940        let request = RunRequest::new(
941            ProviderId::new("mock"),
942            ModelName::new("test"),
943            "test snapshot and resume",
944        );
945
946        let run_id = RunId::new();
947        let session_id = runtime.store().ensure_session(None).await.unwrap();
948
949        // Real runs would already have a run record in the store before crashing/suspending.
950        let run_record = RunRecord::new(
951            run_id,
952            session_id,
953            ProviderId::new("mock"),
954            ModelName::new("test"),
955            serde_json::Value::Null,
956            None,
957        );
958        runtime.store().runs().create_run(run_record).await.unwrap();
959
960        let snapshot = Snapshot {
961            run_id,
962            session_id,
963            status: RunStatus::CallingModel,
964            iteration: 1,
965            current_state: TurnState::CallingModel,
966            total_usage: TokenUsage::new(5, 5),
967            last_finish: Some(FinishReason::ToolCalls),
968            assistant_message: Some(Message::Assistant {
969                content: vec![],
970                tool_calls: vec![ToolCall::new("call_rec", "echo", json!({"message": "rec"}))],
971            }),
972            assistant_msg_id: Some(Uuid::new_v4()),
973            request: request.clone(),
974            output_recovery_count: 0,
975            timestamp: Utc::now(),
976        };
977
978        snapshot_store.save(&snapshot).await.unwrap();
979
980        let output = runtime.resume(run_id).await.unwrap();
981
982        assert_eq!(output.run_id, run_id);
983        assert_eq!(output.session_id, session_id);
984        assert!(matches!(output.finish_reason, FinishReason::Stop));
985    }
986
987    #[tokio::test]
988    async fn run_should_recover_from_length_finish() {
989        let provider = MockProvider::new(vec![
990            MockProvider::length_response("First half..."),
991            MockProvider::length_response("Second half..."),
992            MockProvider::text_response("Complete response."),
993        ]);
994        let mut policy = RuntimePolicy::new();
995        policy.max_output_recovery_attempts = 2;
996        let runtime = make_runtime_with_policy(provider, ToolRegistry::new(), policy);
997
998        let request = RunRequest::new(
999            ProviderId::new("mock"),
1000            ModelName::new("test"),
1001            "Long story",
1002        );
1003        let output = runtime.run(request).await.unwrap();
1004
1005        assert_eq!(output.iterations, 3);
1006        assert!(matches!(output.finish_reason, FinishReason::Stop));
1007    }
1008
1009    #[tokio::test]
1010    async fn run_should_stop_recovery_after_max_attempts() {
1011        let provider = MockProvider::new(vec![
1012            MockProvider::length_response("Try 1..."),
1013            MockProvider::length_response("Try 2..."),
1014            MockProvider::length_response("Still truncated..."),
1015        ]);
1016        let mut policy = RuntimePolicy::new();
1017        policy.max_output_recovery_attempts = 2;
1018        let runtime = make_runtime_with_policy(provider, ToolRegistry::new(), policy);
1019
1020        let request = RunRequest::new(
1021            ProviderId::new("mock"),
1022            ModelName::new("test"),
1023            "Even longer story",
1024        );
1025        let output = runtime.run(request).await.unwrap();
1026
1027        assert_eq!(output.iterations, 3);
1028        assert!(matches!(output.finish_reason, FinishReason::Length));
1029    }
1030
1031    fn make_runtime_with_policy(
1032        provider: MockProvider,
1033        tools: ToolRegistry,
1034        policy: RuntimePolicy,
1035    ) -> AgentRuntime {
1036        let exts = Extensions::new();
1037        exts.chat_providers
1038            .register_or_replace("mock", Arc::new(provider));
1039
1040        let sessions = MemorySessionStore::new();
1041        let executions = MemoryExecutionStore::new();
1042        let runs = MemoryRunStore::new();
1043        exts.session_stores
1044            .register_or_replace("default", Arc::new(sessions));
1045        exts.execution_stores
1046            .register_or_replace("default", Arc::new(executions));
1047        exts.run_stores
1048            .register_or_replace("default", Arc::new(runs));
1049
1050        AgentRuntime::new(Arc::new(exts), policy).with_tool_registry(tools)
1051    }
1052}