Skip to main content

agent_base/engine/runtime/
mod.rs

1use 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    /// Subscribe to runtime events from the internal broadcast channel.
76    ///
77    /// Events are delivered directly from the runtime's event bus (capacity 2048).
78    /// Slow consumers may receive `Lagged(n)` errors if they cannot keep up —
79    /// ensure the receiver loop processes events promptly or use a buffering
80    /// layer in the consumer if backpressure is a concern.
81    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    /// Replace the LLM provider at runtime (e.g., model switch).
98    /// Requires `&mut self` — obtain via `runtime.lock().await`.
99    pub fn set_client(&mut self, provider: Arc<dyn llm_trait::LlmProvider>) {
100        self.runner.llm_engine.set_provider(provider);
101    }
102
103    /// Get the model override, if set.
104    pub fn get_model_override(&self) -> Option<String> {
105        self.runner.llm_engine.get_model_override()
106    }
107
108    /// Set a model override for all requests.
109    ///
110    /// This is used by sub-agents to specify their model tier (e.g., "lite").
111    /// The override is applied to all ChatRequests before sending to the provider.
112    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    /// Snapshot the configured system prompt (async — safe to call inside a
125    /// runtime, unlike [`Self::config`]).
126    pub async fn system_prompt(&self) -> Option<String> {
127        self.runner.config.read().await.system_prompt.clone()
128    }
129
130    /// 设置 reasoning effort(异步版本)
131    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    /// 设置 reasoning effort(同步版本,只在同步上下文中使用)
139    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    /// Like `run_turn`, but the user input is pushed as an ephemeral message:
210    /// visible to the LLM for this turn only, auto-removed at turn end.
211    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    /// Replace the session's system prompt (the first non-ephemeral System
246    /// message). Hosts that own prompt composition re-bake through this.
247    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    /// Replace the chat messages for a session — only for persistence restore.
293    /// Validates message sequence before applying.
294    ///
295    /// 仅供持久化恢复使用。
296    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    // ── Observability hook ──
324
325    /// Register a turn-end callback. The callback receives a [`TurnContext`]
326    /// with raw data about the completed turn iteration. Consumers (e.g.
327    /// phi-telemetry) use this to build their own metrics without agent-base
328    /// knowing anything about metrics.
329    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    // --- Cancellation support ---
341
342    /// Cancel the currently executing run_turn / run.
343    /// No-op if there is no current execution.
344    pub fn cancel(&self) {
345        self.runner.cancel();
346    }
347
348    /// Reset the cancel token (called automatically before each run_turn)
349    pub fn reset_cancel(&self) {
350        self.runner.reset_cancel();
351    }
352
353    /// Get a clone of the cancel token
354    pub fn cancel_token(&self) -> CancellationToken {
355        self.runner.cancel_token()
356    }
357
358    /// Check if cancellation has been requested
359    pub fn is_cancelled(&self) -> bool {
360        self.runner.is_cancelled()
361    }
362
363    // ── Message Queue (P2) ──
364
365    /// Push a steering message — will be processed at the start of the next turn
366    /// in the current `run_managed()` loop.
367    pub fn steer(&self, message: String) {
368        self.runner.message_queue.steer(message);
369    }
370
371    /// Push a follow-up message — will be processed after the inner turn loop
372    /// stops naturally (no tool calls or max turns).
373    pub fn follow_up(&self, message: String) {
374        self.runner.message_queue.follow_up(message);
375    }
376
377    /// Run the agent in managed mode with message queue support.
378    ///
379    /// This wraps `run_turn()` (or `run()`) in an outer loop: after the inner
380    /// turn loop completes, any follow-up messages are drained and a new inner
381    /// loop is started. Steering messages are drained automatically at each
382    /// iteration of the inner turn loop.
383    ///
384    /// The `on_event` callback receives all events from every inner run.
385    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    /// Set the drain mode for the message queues.
400    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}