1mod 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#[derive(Clone)]
30pub struct AgentRuntime {
31 config: AgentConfig,
32 provider: Arc<dyn ModelProvider>,
33 tools: Arc<ToolRegistry>,
34}
35
36#[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
50pub 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}