Skip to main content

telos_agent/agent/runtime/
mod.rs

1//! Public agent runtime facade.
2
3mod pass;
4mod session;
5mod state;
6
7use std::pin::Pin;
8use std::sync::Arc;
9use std::sync::atomic::{AtomicBool, Ordering};
10use std::task::{Context, Poll};
11
12use futures_core::Stream;
13use tokio::sync::{Mutex, mpsc, oneshot};
14
15use crate::agent::context::Conversation;
16use crate::agent::policies::{PolicyContext, PolicyDecision, SessionMode};
17use crate::agent::turn::{TurnEvent, TurnInputSender, TurnResult, turn_input_channel};
18use crate::config::{AgentConfig, CancellationState};
19use crate::error::AgentError;
20use crate::model::message::Message;
21use crate::model::provider::{ErasedProvider, ModelProvider};
22use crate::tools::api::ToolRegistry;
23
24use self::pass::runner::run_turn;
25use self::session::SessionInfo;
26use self::state::RuntimeState;
27
28/// Provider and tool dependencies shared by agent sessions.
29#[derive(Clone)]
30pub struct AgentRuntime {
31    config: AgentConfig,
32    provider: Arc<dyn ModelProvider>,
33    tools: Arc<ToolRegistry>,
34}
35
36/// A concurrency-safe conversation managed by [`AgentRuntime`].
37#[derive(Clone)]
38pub struct AgentSession {
39    session_id: Arc<str>,
40    busy: Arc<AtomicBool>,
41    inner: Arc<Mutex<SessionData>>,
42}
43
44struct SessionData {
45    info: SessionInfo,
46    conversation: Conversation,
47    state: RuntimeState,
48}
49
50/// Live event stream and completion handle for one turn.
51pub struct TurnHandle {
52    events: mpsc::UnboundedReceiver<TurnEvent>,
53    result: Option<oneshot::Receiver<Result<TurnResult, AgentError>>>,
54    input: TurnInputSender,
55    cancellation: CancellationState,
56    completed: bool,
57}
58
59impl AgentRuntime {
60    pub fn new(
61        config: AgentConfig,
62        provider: Arc<dyn ModelProvider>,
63        tools: ToolRegistry,
64    ) -> Result<Self, AgentError> {
65        config.validate()?;
66        Ok(Self { config, provider, tools: Arc::new(tools) })
67    }
68
69    pub async fn create_session(&self) -> Result<AgentSession, AgentError> {
70        let info = SessionInfo::new(self.config.clone())?;
71        let session_id: Arc<str> = Arc::from(info.session_id());
72        let mut conversation = Conversation::new();
73        conversation.initial_messages(&self.config);
74        run_session_policies(&info, &mut conversation, SessionMode::Create).await?;
75        let state = RuntimeState::new();
76        session::persistence::save(
77            info.session_id(),
78            info.config(),
79            conversation.messages(),
80            state.metrics(),
81            state.read_file_state(),
82            info.next_turn_id(),
83        )
84        .await?;
85        Ok(AgentSession {
86            session_id,
87            busy: Arc::new(AtomicBool::new(false)),
88            inner: Arc::new(Mutex::new(SessionData { info, conversation, state })),
89        })
90    }
91
92    pub async fn resume_session(
93        &self,
94        session_id: impl Into<String>,
95    ) -> Result<AgentSession, AgentError> {
96        let storage =
97            self.config.storage.clone().ok_or_else(|| {
98                AgentError::Config("cannot resume without configured storage".into())
99            })?;
100        let (info, conversation, state) =
101            session::persistence::resume(session_id, self.config.clone(), storage).await?;
102        let mut conversation = conversation;
103        run_session_policies(&info, &mut conversation, SessionMode::Resume).await?;
104        session::persistence::save(
105            info.session_id(),
106            info.config(),
107            conversation.messages(),
108            state.metrics(),
109            state.read_file_state(),
110            info.next_turn_id(),
111        )
112        .await?;
113        let session_id: Arc<str> = Arc::from(info.session_id());
114        Ok(AgentSession {
115            session_id,
116            busy: Arc::new(AtomicBool::new(false)),
117            inner: Arc::new(Mutex::new(SessionData { info, conversation, state })),
118        })
119    }
120
121    pub fn start_turn(
122        &self,
123        session: &AgentSession,
124        input: impl Into<String>,
125    ) -> Result<TurnHandle, AgentError> {
126        tokio::runtime::Handle::try_current()
127            .map_err(|_| AgentError::Config("start_turn requires a Tokio runtime".into()))?;
128        if session.busy.compare_exchange(false, true, Ordering::AcqRel, Ordering::Acquire).is_err()
129        {
130            return Err(AgentError::SessionBusy);
131        }
132
133        let (event_tx, event_rx) = mpsc::unbounded_channel();
134        let (result_tx, result_rx) = oneshot::channel();
135        let (input_tx, input_rx) = turn_input_channel();
136        let cancellation = CancellationState::new();
137        let worker_cancellation = cancellation.clone();
138        let session = session.clone();
139        let provider = Arc::clone(&self.provider);
140        let tools = Arc::clone(&self.tools);
141        let input = input.into();
142
143        tokio::spawn(async move {
144            let _busy = BusyGuard(session.busy.clone());
145            let mut data = session.inner.lock().await;
146            data.info.config_mut().cancellation = worker_cancellation;
147            let event_log = Arc::new(std::sync::Mutex::new(Vec::new()));
148            data.info.turn_event_sender = Some(event_tx.clone());
149            data.info.turn_event_log = Some(Arc::clone(&event_log));
150
151            let snapshot = SessionSnapshot::capture(&data).await;
152            let SessionData { info, conversation, state } = &mut *data;
153            let erased = ErasedProvider(provider.as_ref());
154            let execution =
155                run_turn(info, conversation, state, &erased, tools.as_ref(), input, input_rx).await;
156            let execution = match execution {
157                Ok(result) => result,
158                Err(error) => {
159                    snapshot.restore(&mut data).await;
160                    let failed = TurnEvent::TurnFailed { error: error.to_string() };
161                    data.info.emit_turn_event(&failed);
162                    data.info.turn_event_sender = None;
163                    data.info.turn_event_log = None;
164                    let _ = result_tx.send(Err(error));
165                    return;
166                }
167            };
168
169            let events = event_log.lock().map(|log| log.clone()).unwrap_or_default();
170            data.info.turn_event_sender = None;
171            data.info.turn_event_log = None;
172            let _ = result_tx.send(Ok(TurnResult {
173                events,
174                final_message: execution.final_message,
175                stop_reason: execution.stop_reason,
176            }));
177        });
178
179        Ok(TurnHandle {
180            events: event_rx,
181            result: Some(result_rx),
182            input: input_tx,
183            cancellation,
184            completed: false,
185        })
186    }
187
188    pub async fn run_turn(
189        &self,
190        session: &AgentSession,
191        input: impl Into<String>,
192    ) -> Result<TurnResult, AgentError> {
193        self.start_turn(session, input)?.finish().await
194    }
195}
196
197async fn run_session_policies(
198    info: &SessionInfo,
199    conversation: &mut Conversation,
200    mode: SessionMode,
201) -> Result<(), AgentError> {
202    for policy in info.config().policies.session_start(mode) {
203        let outcome = policy
204            .evaluate(&PolicyContext::SessionStart {
205                session_id: info.session_id().to_string(),
206                mode,
207                message_count: conversation.messages().len(),
208            })
209            .await?;
210        for feedback in outcome.feedback {
211            conversation.push_message(Message::user(feedback));
212        }
213        if let PolicyDecision::Reject { reason } = outcome.decision {
214            return Err(AgentError::PermissionDenied(format!(
215                "policy `{}` rejected SessionStart: {reason}",
216                policy.name()
217            )));
218        }
219    }
220    Ok(())
221}
222
223impl AgentSession {
224    pub fn session_id(&self) -> &str {
225        &self.session_id
226    }
227
228    pub fn is_busy(&self) -> bool {
229        self.busy.load(Ordering::Acquire)
230    }
231
232    pub async fn messages(&self) -> Vec<Message> {
233        self.inner.lock().await.conversation.messages().to_vec()
234    }
235
236    pub async fn metrics(&self) -> crate::SessionMetrics {
237        self.inner.lock().await.state.metrics().clone()
238    }
239
240    pub async fn reset(&self) -> Result<(), AgentError> {
241        if self.busy.compare_exchange(false, true, Ordering::AcqRel, Ordering::Acquire).is_err() {
242            return Err(AgentError::SessionBusy);
243        }
244        let _busy = BusyGuard(self.busy.clone());
245        let mut data = self.inner.lock().await;
246        data.conversation.reset();
247        data.info.next_turn_id = 1;
248        data.state = RuntimeState::new();
249        Ok(())
250    }
251}
252
253impl TurnHandle {
254    pub fn input_sender(&self) -> TurnInputSender {
255        self.input.clone()
256    }
257
258    pub fn cancel(&self) {
259        self.cancellation.cancel();
260    }
261
262    pub async fn finish(mut self) -> Result<TurnResult, AgentError> {
263        while self.events.recv().await.is_some() {}
264        let result = self
265            .result
266            .take()
267            .expect("turn result receiver is present")
268            .await
269            .map_err(|_| AgentError::Cancelled)?;
270        self.completed = true;
271        result
272    }
273}
274
275impl Stream for TurnHandle {
276    type Item = TurnEvent;
277
278    fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
279        Pin::new(&mut self.events).poll_recv(cx)
280    }
281}
282
283impl Drop for TurnHandle {
284    fn drop(&mut self) {
285        if !self.completed {
286            self.cancellation.cancel();
287        }
288    }
289}
290
291struct BusyGuard(Arc<AtomicBool>);
292
293impl Drop for BusyGuard {
294    fn drop(&mut self) {
295        self.0.store(false, Ordering::Release);
296    }
297}
298
299struct SessionSnapshot {
300    config: AgentConfig,
301    next_turn_id: u64,
302    messages: Vec<Message>,
303    metrics: crate::metrics::MetricsCheckpoint,
304    read_file_state:
305        std::collections::HashMap<std::path::PathBuf, crate::tools::api::FileReadRecord>,
306    compaction_failures: usize,
307}
308
309impl SessionSnapshot {
310    async fn capture(data: &SessionData) -> Self {
311        Self {
312            config: data.info.config().clone(),
313            next_turn_id: data.info.next_turn_id(),
314            messages: data.conversation.messages().to_vec(),
315            metrics: data.state.metrics().checkpoint(),
316            read_file_state: data.state.read_file_state().lock().await.clone(),
317            compaction_failures: data.state.compaction_failures(),
318        }
319    }
320
321    async fn restore(self, data: &mut SessionData) {
322        *data.info.config_mut() = self.config;
323        data.info.next_turn_id = self.next_turn_id;
324        *data.conversation.messages_mut() = self.messages;
325        data.state.metrics().restore(&self.metrics);
326        *data.state.read_file_state().lock().await = self.read_file_state;
327        data.state.set_compaction_failures(self.compaction_failures);
328    }
329}
330
331#[cfg(test)]
332mod tests {
333    use super::*;
334    use crate::agent::policies::{
335        Policy, PolicyContext, PolicyDecision, PolicyEntry, PolicyOutcome, PolicyPoint,
336        PolicyRegistry,
337    };
338    use crate::model::message::{ContentBlock, Role, ToolCall};
339    use crate::model::mock::MockProvider;
340    use crate::model::provider::{
341        CompletionRequest, CompletionResponse, ProviderEvent, StopReason,
342    };
343    use crate::storage::Storage;
344    use crate::tools::api::{Tool, ToolContext, ToolDefinition, ToolOutput};
345    use async_trait::async_trait;
346    use futures_util::StreamExt;
347    use std::sync::atomic::AtomicUsize;
348
349    fn runtime_with_response(text: &str) -> AgentRuntime {
350        let provider = Arc::new(MockProvider::new(vec![CompletionResponse {
351            message: Message::assistant(text),
352            stop_reason: StopReason::EndTurn,
353            usage: None,
354            model: None,
355        }]));
356        AgentRuntime::new(AgentConfig::default(), provider, ToolRegistry::new()).unwrap()
357    }
358
359    struct SessionFeedback;
360
361    struct EchoTool;
362
363    #[async_trait]
364    impl Tool for EchoTool {
365        fn definition(&self) -> ToolDefinition {
366            ToolDefinition {
367                name: "Echo".into(),
368                description: "echo".into(),
369                input_schema: serde_json::json!({"type":"object"}),
370            }
371        }
372        async fn invoke(
373            &self,
374            _: serde_json::Value,
375            _: ToolContext,
376        ) -> Result<ToolOutput, AgentError> {
377            Ok(ToolOutput::text("ok"))
378        }
379    }
380
381    struct OneShotPassFeedback(AtomicBool);
382
383    #[async_trait]
384    impl Policy for OneShotPassFeedback {
385        fn name(&self) -> &str {
386            "pass-feedback"
387        }
388
389        async fn evaluate(&self, _: &PolicyContext) -> Result<PolicyOutcome, AgentError> {
390            if self.0.swap(true, Ordering::SeqCst) {
391                Ok(PolicyOutcome::continue_())
392            } else {
393                Ok(PolicyOutcome {
394                    decision: PolicyDecision::Continue,
395                    feedback: vec!["revise".into()],
396                })
397            }
398        }
399    }
400
401    #[async_trait]
402    impl Policy for SessionFeedback {
403        fn name(&self) -> &str {
404            "session-feedback"
405        }
406        async fn evaluate(&self, context: &PolicyContext) -> Result<PolicyOutcome, AgentError> {
407            assert!(matches!(
408                context,
409                PolicyContext::SessionStart { mode: SessionMode::Create, .. }
410            ));
411            Ok(PolicyOutcome {
412                decision: PolicyDecision::Continue,
413                feedback: vec!["session context".into()],
414            })
415        }
416    }
417
418    #[tokio::test]
419    async fn create_session_runs_session_start_policies() {
420        let mut registry = PolicyRegistry::new();
421        registry.register(PolicyEntry {
422            point: PolicyPoint::SessionStart { mode: Some(SessionMode::Create) },
423            policy: Arc::new(SessionFeedback),
424        });
425        let mut config = AgentConfig::default();
426        config.policies = Arc::new(registry);
427        let runtime =
428            AgentRuntime::new(config, Arc::new(MockProvider::new(Vec::new())), ToolRegistry::new())
429                .unwrap();
430        let session = runtime.create_session().await.unwrap();
431        assert_eq!(session.messages().await.last().unwrap().text_content(), "session context");
432    }
433
434    #[tokio::test]
435    async fn pass_feedback_triggers_another_model_iteration() {
436        let mut registry = PolicyRegistry::new();
437        registry.register(PolicyEntry {
438            point: PolicyPoint::TurnBeforeFinish,
439            policy: Arc::new(OneShotPassFeedback(AtomicBool::new(false))),
440        });
441        let mut config = AgentConfig::default();
442        config.policies = Arc::new(registry);
443        let provider = Arc::new(MockProvider::new(vec![
444            CompletionResponse {
445                message: Message::assistant("first"),
446                stop_reason: StopReason::EndTurn,
447                usage: None,
448                model: None,
449            },
450            CompletionResponse {
451                message: Message::assistant("revised"),
452                stop_reason: StopReason::EndTurn,
453                usage: None,
454                model: None,
455            },
456        ]));
457        let runtime = AgentRuntime::new(config, provider, ToolRegistry::new()).unwrap();
458        let session = runtime.create_session().await.unwrap();
459        let result = runtime.run_turn(&session, "hello").await.unwrap();
460        assert_eq!(result.final_message.text_content(), "revised");
461        assert!(session.messages().await.iter().any(|message| message.text_content() == "revise"));
462    }
463
464    #[tokio::test]
465    async fn model_policy_feedback_waits_until_tool_results_are_committed() {
466        let mut registry = PolicyRegistry::new();
467        registry.register(PolicyEntry {
468            point: PolicyPoint::ModelResponse,
469            policy: Arc::new(OneShotPassFeedback(AtomicBool::new(false))),
470        });
471        let mut config = AgentConfig::default();
472        config.policies = Arc::new(registry);
473        let provider = Arc::new(MockProvider::new(vec![
474            CompletionResponse {
475                message: Message {
476                    role: Role::Assistant,
477                    blocks: vec![ContentBlock::ToolCall(ToolCall {
478                        id: "call-1".into(),
479                        name: "Echo".into(),
480                        arguments: serde_json::json!({}),
481                    })],
482                },
483                stop_reason: StopReason::ToolUse,
484                usage: None,
485                model: None,
486            },
487            CompletionResponse {
488                message: Message::assistant("done"),
489                stop_reason: StopReason::EndTurn,
490                usage: None,
491                model: None,
492            },
493        ]));
494        let mut tools = ToolRegistry::new();
495        tools.register(EchoTool);
496        let runtime = AgentRuntime::new(config, provider, tools).unwrap();
497        let session = runtime.create_session().await.unwrap();
498        runtime.run_turn(&session, "hello").await.unwrap();
499        let messages = session.messages().await;
500        let tool_index = messages.iter().position(|message| message.role == Role::Tool).unwrap();
501        let feedback_index =
502            messages.iter().position(|message| message.text_content() == "revise").unwrap();
503        assert!(tool_index < feedback_index);
504    }
505
506    #[tokio::test]
507    async fn run_turn_commits_messages_and_returns_events() {
508        let runtime = runtime_with_response("done");
509        let session = runtime.create_session().await.unwrap();
510        let result = runtime.run_turn(&session, "hello").await.unwrap();
511        assert_eq!(result.final_message.text_content(), "done");
512        assert!(result.events.iter().any(|event| matches!(event, TurnEvent::TurnFinished { .. })));
513        assert!(matches!(result.events.last(), Some(TurnEvent::TurnFinished { .. })));
514        assert_eq!(session.messages().await.last().unwrap().text_content(), "done");
515    }
516
517    #[tokio::test]
518    async fn rejects_a_second_concurrent_turn() {
519        let runtime = runtime_with_response("done");
520        let session = runtime.create_session().await.unwrap();
521        let first = runtime.start_turn(&session, "first").unwrap();
522        assert!(matches!(runtime.start_turn(&session, "second"), Err(AgentError::SessionBusy)));
523        first.finish().await.unwrap();
524    }
525
526    struct ControlledProvider {
527        release: Arc<tokio::sync::Notify>,
528    }
529
530    #[async_trait]
531    impl ModelProvider for ControlledProvider {
532        async fn complete(
533            &self,
534            _request: CompletionRequest,
535        ) -> Result<CompletionResponse, AgentError> {
536            unreachable!("the controlled provider uses its stream implementation")
537        }
538
539        fn stream_complete<'a>(
540            &'a self,
541            _request: CompletionRequest,
542        ) -> Pin<Box<dyn Stream<Item = Result<ProviderEvent, AgentError>> + Send + 'a>> {
543            Box::pin(async_stream::try_stream! {
544                yield ProviderEvent::MessageStart;
545                yield ProviderEvent::TextDelta("partial".into());
546                self.release.notified().await;
547                yield ProviderEvent::MessageStop {
548                    stop_reason: StopReason::EndTurn,
549                    usage: None,
550                    model: None,
551                };
552            })
553        }
554    }
555
556    #[tokio::test]
557    async fn provider_delta_is_visible_before_provider_finishes() {
558        let release = Arc::new(tokio::sync::Notify::new());
559        let runtime = AgentRuntime::new(
560            AgentConfig::default(),
561            Arc::new(ControlledProvider { release: Arc::clone(&release) }),
562            ToolRegistry::new(),
563        )
564        .unwrap();
565        let session = runtime.create_session().await.unwrap();
566        let mut handle = runtime.start_turn(&session, "hello").unwrap();
567
568        loop {
569            let event = tokio::time::timeout(std::time::Duration::from_secs(1), handle.next())
570                .await
571                .expect("delta should arrive while provider is blocked")
572                .expect("event stream should remain open");
573            if matches!(event, TurnEvent::AssistantDelta { ref text } if text == "partial") {
574                break;
575            }
576        }
577        assert!(session.is_busy());
578        release.notify_waiters();
579        handle.finish().await.unwrap();
580    }
581
582    #[tokio::test]
583    async fn dropping_handle_cancels_and_rolls_back_session() {
584        let release = Arc::new(tokio::sync::Notify::new());
585        let runtime = AgentRuntime::new(
586            AgentConfig::default(),
587            Arc::new(ControlledProvider { release }),
588            ToolRegistry::new(),
589        )
590        .unwrap();
591        let session = runtime.create_session().await.unwrap();
592        let mut handle = runtime.start_turn(&session, "temporary").unwrap();
593        while let Some(event) = handle.next().await {
594            if matches!(event, TurnEvent::AssistantDelta { .. }) {
595                break;
596            }
597        }
598        drop(handle);
599
600        tokio::time::timeout(std::time::Duration::from_secs(1), async {
601            while session.is_busy() {
602                tokio::task::yield_now().await;
603            }
604        })
605        .await
606        .expect("cancelled worker should release the session");
607        assert!(session.messages().await.is_empty());
608    }
609
610    #[derive(Debug)]
611    struct FailingStorage(AtomicUsize);
612
613    #[async_trait]
614    impl Storage for FailingStorage {
615        async fn save_snapshot(
616            &self,
617            _session_id: &str,
618            _messages: &[Message],
619        ) -> Result<(), AgentError> {
620            if self.0.fetch_add(1, Ordering::SeqCst) == 0 {
621                Ok(())
622            } else {
623                Err(AgentError::Config("storage unavailable".into()))
624            }
625        }
626
627        async fn append(&self, _session_id: &str, _messages: &[Message]) -> Result<(), AgentError> {
628            Ok(())
629        }
630
631        async fn load(&self, _session_id: &str) -> Result<Vec<Message>, AgentError> {
632            Ok(Vec::new())
633        }
634    }
635
636    #[tokio::test]
637    async fn persistence_failure_rolls_back_turn() {
638        let mut config = AgentConfig::default();
639        config.storage = Some(Arc::new(FailingStorage(AtomicUsize::new(0))));
640        let provider = Arc::new(MockProvider::new(vec![CompletionResponse {
641            message: Message::assistant("done"),
642            stop_reason: StopReason::EndTurn,
643            usage: None,
644            model: None,
645        }]));
646        let runtime = AgentRuntime::new(config, provider, ToolRegistry::new()).unwrap();
647        let session = runtime.create_session().await.unwrap();
648        let result = runtime.run_turn(&session, "hello").await;
649        assert!(
650            matches!(result, Err(AgentError::Config(message)) if message.contains("storage unavailable"))
651        );
652        assert!(session.messages().await.is_empty());
653    }
654
655    #[tokio::test]
656    async fn resume_restores_persisted_conversation() {
657        let dir = tempfile::tempdir().unwrap();
658        let mut config = AgentConfig::default();
659        config.storage = Some(Arc::new(crate::storage::JsonlStorage::new(dir.path()).unwrap()));
660        let provider = Arc::new(MockProvider::new(vec![CompletionResponse {
661            message: Message::assistant("persisted"),
662            stop_reason: StopReason::EndTurn,
663            usage: None,
664            model: None,
665        }]));
666        let runtime = AgentRuntime::new(config, provider, ToolRegistry::new()).unwrap();
667        let session = runtime.create_session().await.unwrap();
668        runtime.run_turn(&session, "hello").await.unwrap();
669
670        let resumed = runtime.resume_session(session.session_id()).await.unwrap();
671        assert_eq!(resumed.messages().await.last().unwrap().text_content(), "persisted");
672    }
673}