1use 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
36pub 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 pub(super) extensions: Arc<Extensions>,
59}
60
61impl AgentRuntime {
62 #[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 #[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 #[cfg(feature = "queue")]
118 #[must_use]
119 pub fn with_event_publisher(mut self, publisher: Arc<dyn EventPublisher>) -> Self {
120 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 #[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 #[must_use]
150 pub fn extensions(&self) -> &Arc<Extensions> {
151 &self.extensions
152 }
153
154 #[must_use]
156 pub fn session_gate(&self) -> &SessionGate {
157 &self.session_gate
158 }
159
160 #[must_use]
165 pub fn subscribe(&self) -> broadcast::Receiver<AgentEvent> {
166 self.event_tx.subscribe()
167 }
168
169 #[must_use]
171 pub fn policy(&self) -> &RuntimePolicy {
172 &self.policy
173 }
174
175 #[must_use]
177 pub fn tools(&self) -> &ToolRuntime {
178 &self.tools
179 }
180
181 #[must_use]
183 pub fn providers(&self) -> &behest_provider::ProviderRegistry {
184 &self.providers
185 }
186
187 #[must_use]
189 pub fn context(&self) -> &ContextPipeline {
190 &self.context
191 }
192
193 #[must_use]
195 pub fn store(&self) -> &Arc<RuntimeStore> {
196 &self.store
197 }
198
199 #[must_use]
201 pub fn compaction(&self) -> &CompactionService {
202 &self.compaction
203 }
204
205 #[must_use]
207 pub fn snapshot_store(&self) -> Option<&Arc<dyn SnapshotStore>> {
208 self.snapshot_store.as_ref()
209 }
210
211 #[must_use]
213 pub fn sessions(&self) -> &dyn behest_store::SessionStore {
214 self.store.sessions()
215 }
216
217 #[must_use]
219 pub fn executions(&self) -> &dyn behest_store::ExecutionStore {
220 self.store.executions()
221 }
222
223 #[must_use]
225 pub fn runs(&self) -> &dyn RunStore {
226 self.store.runs()
227 }
228
229 #[must_use]
231 pub fn embeddings(&self) -> Option<&dyn behest_store::EmbeddingStore> {
232 self.store.embeddings()
233 }
234
235 #[must_use]
237 pub fn artifacts(&self) -> Option<&dyn behest_store::ArtifactStore> {
238 self.store.artifacts()
239 }
240
241 #[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 let _session_guard = self
260 .session_gate
261 .acquire(session_id)
262 .await
263 .map_err(|busy| RuntimeError::SessionBusy(busy.session_id))?;
264
265 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 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 let mut doom_detector = DoomLoopDetector::new(self.policy.doom_loop.clone());
300
301 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 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 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 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 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 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 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 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 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 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 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#[derive(Debug, Clone)]
552pub struct RunOutput {
553 pub run_id: RunId,
555 pub session_id: Uuid,
557 pub iterations: usize,
559 pub finish_reason: FinishReason,
561 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 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}