agent_base/engine/runtime/
mod.rs1use std::sync::Arc;
2
3use tokio_util::sync::CancellationToken;
4
5use crate::engine::AgentSession;
6use crate::engine::session_store::SessionStore;
7use crate::types::{
8 AgentConfig, AgentError, AgentResult, CheckpointData, MessageRole, RunOutcome, RuntimeEvent,
9 SessionId, TurnContext,
10};
11
12use super::approval::ApprovalHandler;
13use crate::tool::ToolPolicy;
14
15mod event_bus;
16pub(crate) use event_bus::EventBus;
17mod llm_engine;
18mod message_queue;
19mod plan_runner;
20mod react;
21mod session_manager;
22mod tool_engine;
23
24pub(super) const DEFAULT_MAX_TURNS: u32 = 160;
25
26pub use llm_engine::LlmEngine;
27pub use message_queue::QueueMode;
28pub(crate) use plan_runner::RuntimeCore;
29pub use session_manager::SessionManager;
30pub(crate) use tool_engine::ToolEngine;
31
32#[derive(Clone)]
33pub struct AgentRuntime {
34 pub(crate) runner: Arc<RuntimeCore>,
35}
36
37impl AgentRuntime {
38 pub async fn create_session(&self) -> SessionId {
39 let config = self.runner.config.read().await;
40 self.runner
41 .session_manager
42 .create_session(config.system_prompt.as_deref())
43 .await
44 }
45
46 pub async fn restore_session(&self, session_id: &SessionId) -> Option<AgentSession> {
47 self.runner
48 .session_manager
49 .restore_session(session_id)
50 .await
51 }
52
53 pub async fn session(&self, session_id: &SessionId) -> Option<AgentSession> {
54 self.runner.session_manager.session(session_id).await
55 }
56
57 pub async fn session_or_err(&self, session_id: &SessionId) -> AgentResult<AgentSession> {
58 self.runner.session_manager.session_or_err(session_id).await
59 }
60
61 pub async fn with_session_mut<F, R>(&self, session_id: &SessionId, f: F) -> AgentResult<R>
62 where
63 F: FnOnce(&mut AgentSession) -> R,
64 {
65 self.runner
66 .session_manager
67 .with_session_mut(session_id, f)
68 .await
69 }
70
71 pub fn emit_event(&self, event: RuntimeEvent) {
72 self.runner.event_bus.emit(event);
73 }
74
75 pub fn subscribe_runtime_events(&self) -> tokio::sync::broadcast::Receiver<RuntimeEvent> {
82 self.runner.event_bus.subscribe()
83 }
84
85 pub fn session_manager(&self) -> &SessionManager {
86 &self.runner.session_manager
87 }
88
89 pub fn llm_engine(&self) -> &LlmEngine {
90 &self.runner.llm_engine
91 }
92
93 pub fn provider(&self) -> Arc<dyn llm_trait::LlmProvider> {
94 self.runner.llm_engine.get_provider()
95 }
96
97 pub fn set_client(&mut self, provider: Arc<dyn llm_trait::LlmProvider>) {
100 self.runner.llm_engine.set_provider(provider);
101 }
102
103 pub fn get_model_override(&self) -> Option<String> {
105 self.runner.llm_engine.get_model_override()
106 }
107
108 pub fn set_model_override(&self, model: Option<String>) {
113 self.runner.llm_engine.set_model_override(model);
114 }
115
116 pub fn tools_mut(&self) -> Arc<tokio::sync::RwLock<crate::tool::ToolRegistry>> {
117 self.runner.tool_engine.tools_arc()
118 }
119
120 pub fn config(&self) -> tokio::sync::RwLockReadGuard<'_, AgentConfig> {
121 self.runner.config.blocking_read()
122 }
123
124 pub async fn system_prompt(&self) -> Option<String> {
127 self.runner.config.read().await.system_prompt.clone()
128 }
129
130 pub async fn set_reasoning_effort(&self, effort: crate::llm::ReasoningEffort) {
132 let mut config = self.runner.config.write().await;
133 let mut reasoning = config.reasoning.take().unwrap_or_default();
134 reasoning.effort = Some(effort);
135 config.reasoning = Some(reasoning);
136 }
137
138 pub fn set_reasoning_effort_sync(&self, effort: crate::llm::ReasoningEffort) {
140 let mut config = self.runner.config.blocking_write();
141 let mut reasoning = config.reasoning.take().unwrap_or_default();
142 reasoning.effort = Some(effort);
143 config.reasoning = Some(reasoning);
144 }
145
146 pub fn approval_handler(&self) -> Option<&Arc<dyn ApprovalHandler>> {
147 self.runner.tool_engine.approval_handler()
148 }
149
150 pub fn tool_policy(&self) -> Option<&Arc<dyn ToolPolicy>> {
151 self.runner.tool_engine.tool_policy()
152 }
153
154 pub async fn cached_approval(&self, session_id: &SessionId, action_key: &str) -> bool {
155 self.runner
156 .session_manager
157 .cached_approval(session_id, action_key)
158 .await
159 }
160
161 pub async fn cache_approval(&self, session_id: &SessionId, action_key: String) {
162 self.runner
163 .session_manager
164 .cache_approval(session_id, action_key)
165 .await
166 }
167
168 pub async fn save_checkpoint(
169 &self,
170 session_id: &SessionId,
171 checkpoint: CheckpointData,
172 ) -> AgentResult<()> {
173 self.emit_event(RuntimeEvent::Checkpoint {
174 session_id: session_id.clone(),
175 checkpoint,
176 agent_id: None,
177 trace_id: None,
178 });
179 Ok(())
180 }
181
182 pub async fn load_checkpoint(
183 &self,
184 _session_id: &SessionId,
185 _checkpoint: &CheckpointData,
186 ) -> AgentResult<Option<CheckpointData>> {
187 Ok(None)
188 }
189
190 pub async fn run<F>(&self, session_id: SessionId, on_event: F) -> AgentResult<RunOutcome>
191 where
192 F: FnMut(RuntimeEvent) -> AgentResult<()> + Send + 'static,
193 {
194 self.runner.run(session_id, on_event).await
195 }
196
197 pub async fn run_turn<F>(
198 &self,
199 session_id: SessionId,
200 user_input: &str,
201 on_event: F,
202 ) -> AgentResult<RunOutcome>
203 where
204 F: FnMut(RuntimeEvent) -> AgentResult<()> + Send + 'static,
205 {
206 self.runner.run_turn(session_id, user_input, on_event).await
207 }
208
209 pub async fn run_turn_ephemeral_input<F>(
212 &self,
213 session_id: SessionId,
214 user_input: &str,
215 on_event: F,
216 ) -> AgentResult<RunOutcome>
217 where
218 F: FnMut(RuntimeEvent) -> AgentResult<()> + Send + 'static,
219 {
220 self.runner
221 .run_turn_ephemeral_input(session_id, user_input, on_event)
222 .await
223 }
224
225 pub async fn run_turn_collect(
226 &self,
227 session_id: SessionId,
228 user_input: &str,
229 ) -> AgentResult<(Vec<RuntimeEvent>, RunOutcome)> {
230 self.runner.run_turn_collect(session_id, user_input).await
231 }
232
233 pub async fn add_user_message(
234 &self,
235 session_id: &SessionId,
236 text: impl Into<String>,
237 ) -> AgentResult<()> {
238 let text = text.into();
239 self.with_session_mut(session_id, |session| {
240 session.push_message(MessageRole::User, &text);
241 })
242 .await
243 }
244
245 pub async fn set_system_prompt(
248 &self,
249 session_id: &SessionId,
250 prompt: impl Into<String>,
251 ) -> AgentResult<()> {
252 let prompt = prompt.into();
253 self.with_session_mut(session_id, |session| {
254 session.set_system_prompt(prompt);
255 })
256 .await
257 }
258
259 pub async fn add_system_message(
260 &self,
261 session_id: &SessionId,
262 text: impl Into<String>,
263 ) -> AgentResult<()> {
264 let text = text.into();
265 self.with_session_mut(session_id, |session| {
266 session.push_message(MessageRole::System, &text);
267 })
268 .await
269 }
270
271 pub async fn add_tool_result(
272 &self,
273 session_id: &SessionId,
274 tool_call_id: &str,
275 summary: impl Into<String>,
276 ) -> AgentResult<()> {
277 let summary = summary.into();
278 self.with_session_mut(session_id, |session| {
279 session.push_tool_result(tool_call_id, summary.clone());
280 })
281 .await
282 }
283
284 pub async fn get_messages(
285 &self,
286 session_id: &SessionId,
287 ) -> AgentResult<Vec<crate::types::ChatMessage>> {
288 let session = self.session_or_err(session_id).await?;
289 Ok(session.chat_messages().to_vec())
290 }
291
292 pub async fn set_messages(
297 &self,
298 session_id: &SessionId,
299 messages: Vec<crate::types::ChatMessage>,
300 ) -> AgentResult<()> {
301 self.with_session_mut(session_id, |session| session.set_chat_messages(messages))
302 .await?
303 .map_err(AgentError::internal)
304 }
305
306 pub async fn validate_session(&self, session_id: &SessionId) -> AgentResult<()> {
307 if self
308 .runner
309 .session_manager
310 .session(session_id)
311 .await
312 .is_none()
313 {
314 return Err(AgentError::session_not_found(session_id.id));
315 }
316 Ok(())
317 }
318
319 pub fn session_store(&self) -> Arc<dyn SessionStore> {
320 self.runner.session_manager.session_store().clone()
321 }
322
323 pub fn on_turn_end<F>(&self, f: F)
330 where
331 F: Fn(&TurnContext) + Send + Sync + 'static,
332 {
333 self.runner
334 .turn_end_callbacks
335 .write()
336 .unwrap()
337 .push(Arc::new(f));
338 }
339
340 pub fn cancel(&self) {
345 self.runner.cancel();
346 }
347
348 pub fn reset_cancel(&self) {
350 self.runner.reset_cancel();
351 }
352
353 pub fn cancel_token(&self) -> CancellationToken {
355 self.runner.cancel_token()
356 }
357
358 pub fn is_cancelled(&self) -> bool {
360 self.runner.is_cancelled()
361 }
362
363 pub fn steer(&self, message: String) {
368 self.runner.message_queue.steer(message);
369 }
370
371 pub fn follow_up(&self, message: String) {
374 self.runner.message_queue.follow_up(message);
375 }
376
377 pub async fn run_managed<F>(
386 &self,
387 session_id: SessionId,
388 user_input: &str,
389 on_event: F,
390 ) -> AgentResult<RunOutcome>
391 where
392 F: FnMut(RuntimeEvent) -> AgentResult<()> + Send + 'static,
393 {
394 self.runner
395 .run_managed(session_id, user_input, on_event)
396 .await
397 }
398
399 pub fn set_queue_mode(&self, mode: crate::engine::runtime::message_queue::QueueMode) {
401 self.runner.message_queue.set_mode(mode);
402 }
403}
404
405#[cfg(test)]
406mod tests {
407 use super::*;
408 use crate::llm::ReasoningEffort;
409 use crate::types::{ChatMessage, RuntimeEvent, SessionId};
410 use async_trait::async_trait;
411 use llm_trait::{Capabilities, ChatRequest, ChatResponse, ChatStream, LlmError, ProviderInfo};
412
413 struct StubProvider;
414
415 #[async_trait]
416 impl llm_trait::LlmProvider for StubProvider {
417 async fn stream(&self, _request: ChatRequest) -> Result<ChatStream, LlmError> {
418 Ok(ChatStream::new(Box::pin(futures_util::stream::empty())))
419 }
420
421 async fn chat(&self, _request: ChatRequest) -> Result<ChatResponse, LlmError> {
422 Ok(ChatResponse {
423 content: String::new(),
424 reasoning_content: None,
425 tool_calls: vec![],
426 usage: Default::default(),
427 finish_reason: llm_trait::FinishReason::Stop,
428 raw: None,
429 thinking_signature: None,
430 })
431 }
432
433 fn capabilities(&self) -> Capabilities {
434 Capabilities::default()
435 }
436
437 fn info(&self) -> ProviderInfo {
438 ProviderInfo {
439 name: "stub".to_string(),
440 model: "stub".to_string(),
441 version: None,
442 }
443 }
444 }
445
446 fn runtime() -> AgentRuntime {
447 crate::engine::AgentBuilder::new(Arc::new(StubProvider))
448 .build()
449 .unwrap()
450 }
451
452 #[tokio::test]
453 async fn create_session_and_lookup() {
454 let rt = runtime();
455 let id = rt.create_session().await;
456 assert_eq!(id.id, 1);
457 assert!(rt.session(&id).await.is_some());
458 assert!(rt.session_or_err(&id).await.is_ok());
459 }
460
461 #[tokio::test]
462 async fn add_messages_and_get() {
463 let rt = runtime();
464 let id = rt.create_session().await;
465 rt.add_system_message(&id, "sys").await.unwrap();
466 rt.add_user_message(&id, "hello").await.unwrap();
467 let msgs = rt.get_messages(&id).await.unwrap();
468 assert_eq!(msgs.len(), 2);
469 assert!(matches!(msgs[0], ChatMessage::System { .. }));
470 assert!(matches!(msgs[1], ChatMessage::User { .. }));
471 }
472
473 #[tokio::test]
474 async fn add_tool_result_appends_tool_message() {
475 let rt = runtime();
476 let id = rt.create_session().await;
477 rt.add_tool_result(&id, "call_1", "done").await.unwrap();
478 let msgs = rt.get_messages(&id).await.unwrap();
479 assert_eq!(msgs.len(), 1);
480 assert!(matches!(msgs[0], ChatMessage::Tool { .. }));
481 }
482
483 #[tokio::test]
484 async fn set_messages_replaces_history() {
485 let rt = runtime();
486 let id = rt.create_session().await;
487 rt.add_user_message(&id, "old").await.unwrap();
488 rt.set_messages(
489 &id,
490 vec![ChatMessage::system("sys"), ChatMessage::user("new")],
491 )
492 .await
493 .unwrap();
494 let msgs = rt.get_messages(&id).await.unwrap();
495 assert_eq!(msgs.len(), 2);
496 }
497
498 #[tokio::test]
499 async fn validate_session_errors_for_unknown() {
500 let rt = runtime();
501 let id = rt.create_session().await;
502 assert!(rt.validate_session(&id).await.is_ok());
503 let err = rt.validate_session(&SessionId::new(999)).await.unwrap_err();
504 assert!(matches!(err, AgentError::SessionNotFound(_)));
505 }
506
507 #[test]
508 fn config_and_set_reasoning_effort() {
509 let rt = runtime();
510 assert!(rt.config().system_prompt.is_none());
511
512 rt.set_reasoning_effort_sync(ReasoningEffort::High);
513 let cfg = rt.config();
514 let effort = cfg.reasoning.as_ref().and_then(|r| r.effort.as_ref());
515 assert!(matches!(effort, Some(ReasoningEffort::High)));
516 }
517
518 #[tokio::test]
519 async fn session_store_is_available() {
520 let rt = runtime();
521 assert!(rt.session_store().list().await.unwrap().is_empty());
522 }
523
524 #[tokio::test]
525 async fn emit_and_subscribe_event() {
526 let rt = runtime();
527 let mut rx = rt.subscribe_runtime_events();
528 rt.emit_event(RuntimeEvent::TextDelta {
529 session_id: SessionId::new(1),
530 text: "hi".into(),
531 agent_id: None,
532 trace_id: None,
533 });
534 let ev = rx.recv().await.unwrap();
535 assert!(matches!(ev, RuntimeEvent::TextDelta { .. }));
536 }
537
538 #[tokio::test]
539 async fn cancel_reset_and_is_cancelled() {
540 let rt = runtime();
541 assert!(!rt.is_cancelled());
542 rt.cancel();
543 assert!(rt.is_cancelled());
544 rt.reset_cancel();
545 assert!(!rt.is_cancelled());
546 }
547}