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, is_self_echo_user_event};
17mod llm_engine;
18mod message_queue;
19mod notify;
20mod plan_runner;
21mod react;
22mod session_manager;
23mod tool_engine;
24
25pub(super) const DEFAULT_MAX_TURNS: u32 = 160;
26
27pub use llm_engine::LlmEngine;
28pub use message_queue::QueueMode;
29pub use notify::NoticeHandle;
30pub(crate) use plan_runner::RuntimeCore;
31pub use session_manager::SessionManager;
32pub(crate) use tool_engine::ToolEngine;
33
34#[derive(Clone)]
35pub struct AgentRuntime {
36    pub(crate) runner: Arc<RuntimeCore>,
37}
38
39impl AgentRuntime {
40    pub async fn create_session(&self) -> SessionId {
41        let config = self.runner.config.read().await;
42        self.runner
43            .session_manager
44            .create_session(config.system_prompt.as_deref())
45            .await
46    }
47
48    pub async fn restore_session(&self, session_id: &SessionId) -> Option<AgentSession> {
49        self.runner
50            .session_manager
51            .restore_session(session_id)
52            .await
53    }
54
55    pub async fn session(&self, session_id: &SessionId) -> Option<AgentSession> {
56        self.runner.session_manager.session(session_id).await
57    }
58
59    pub async fn session_or_err(&self, session_id: &SessionId) -> AgentResult<AgentSession> {
60        self.runner.session_manager.session_or_err(session_id).await
61    }
62
63    pub async fn with_session_mut<F, R>(&self, session_id: &SessionId, f: F) -> AgentResult<R>
64    where
65        F: FnOnce(&mut AgentSession) -> R,
66    {
67        self.runner
68            .session_manager
69            .with_session_mut(session_id, f)
70            .await
71    }
72
73    /// Publish an event on the internal event bus.
74    ///
75    /// Events land on the bus: external subscribers (`subscribe_runtime_events`)
76    /// receive them, and the run's own bus loopback delivers them to the
77    /// `on_event` callback — **except** `UserEvent`s with `agent_id: None`,
78    /// which are treated as self-echoes of events the run already rendered
79    /// (see [`is_self_echo_user_event`]). Run-local `UserEvent`s should be
80    /// published through the dual-write APIs (`PreLlmCtx::emit`,
81    /// `ToolContext::emit_user_event`, `NoticeHandle`/the notice pump) instead;
82    /// events bridged from a child agent (`with_agent_id`) reach the renderer
83    /// through this method and its loopback.
84    pub fn emit_event(&self, event: RuntimeEvent) {
85        self.runner.event_bus.emit(event);
86    }
87
88    /// Subscribe to runtime events from the internal broadcast channel.
89    ///
90    /// Events are delivered directly from the runtime's event bus (capacity 2048).
91    /// Slow consumers may receive `Lagged(n)` errors if they cannot keep up —
92    /// ensure the receiver loop processes events promptly or use a buffering
93    /// layer in the consumer if backpressure is a concern.
94    pub fn subscribe_runtime_events(&self) -> tokio::sync::broadcast::Receiver<RuntimeEvent> {
95        self.runner.event_bus.subscribe()
96    }
97
98    pub fn session_manager(&self) -> &SessionManager {
99        &self.runner.session_manager
100    }
101
102    pub fn llm_engine(&self) -> &LlmEngine {
103        &self.runner.llm_engine
104    }
105
106    pub fn provider(&self) -> Arc<dyn llm_trait::LlmProvider> {
107        self.runner.llm_engine.get_provider()
108    }
109
110    /// Replace the LLM provider at runtime (e.g., model switch).
111    /// Requires `&mut self` — obtain via `runtime.lock().await`.
112    pub fn set_client(&mut self, provider: Arc<dyn llm_trait::LlmProvider>) {
113        self.runner.llm_engine.set_provider(provider);
114    }
115
116    /// Get the model override, if set.
117    pub fn get_model_override(&self) -> Option<String> {
118        self.runner.llm_engine.get_model_override()
119    }
120
121    /// Set a model override for all requests.
122    ///
123    /// This is used by sub-agents to specify their model tier (e.g., "lite").
124    /// The override is applied to all ChatRequests before sending to the provider.
125    pub fn set_model_override(&self, model: Option<String>) {
126        self.runner.llm_engine.set_model_override(model);
127    }
128
129    pub fn tools_mut(&self) -> Arc<tokio::sync::RwLock<crate::tool::ToolRegistry>> {
130        self.runner.tool_engine.tools_arc()
131    }
132
133    pub fn config(&self) -> tokio::sync::RwLockReadGuard<'_, AgentConfig> {
134        self.runner.config.blocking_read()
135    }
136
137    /// Snapshot the configured system prompt (async — safe to call inside a
138    /// runtime, unlike [`Self::config`]).
139    pub async fn system_prompt(&self) -> Option<String> {
140        self.runner.config.read().await.system_prompt.clone()
141    }
142
143    /// 设置 reasoning effort(异步版本)
144    pub async fn set_reasoning_effort(&self, effort: crate::llm::ReasoningEffort) {
145        let mut config = self.runner.config.write().await;
146        let mut reasoning = config.reasoning.take().unwrap_or_default();
147        reasoning.effort = Some(effort);
148        config.reasoning = Some(reasoning);
149    }
150
151    /// 设置 reasoning effort(同步版本,只在同步上下文中使用)
152    pub fn set_reasoning_effort_sync(&self, effort: crate::llm::ReasoningEffort) {
153        let mut config = self.runner.config.blocking_write();
154        let mut reasoning = config.reasoning.take().unwrap_or_default();
155        reasoning.effort = Some(effort);
156        config.reasoning = Some(reasoning);
157    }
158
159    pub fn approval_handler(&self) -> Option<&Arc<dyn ApprovalHandler>> {
160        self.runner.tool_engine.approval_handler()
161    }
162
163    pub fn tool_policy(&self) -> Option<&Arc<dyn ToolPolicy>> {
164        self.runner.tool_engine.tool_policy()
165    }
166
167    pub async fn cached_approval(&self, session_id: &SessionId, action_key: &str) -> bool {
168        self.runner
169            .session_manager
170            .cached_approval(session_id, action_key)
171            .await
172    }
173
174    pub async fn cache_approval(&self, session_id: &SessionId, action_key: String) {
175        self.runner
176            .session_manager
177            .cache_approval(session_id, action_key)
178            .await
179    }
180
181    pub async fn save_checkpoint(
182        &self,
183        session_id: &SessionId,
184        checkpoint: CheckpointData,
185    ) -> AgentResult<()> {
186        self.emit_event(RuntimeEvent::Checkpoint {
187            session_id: session_id.clone(),
188            checkpoint,
189            agent_id: None,
190            trace_id: None,
191        });
192        Ok(())
193    }
194
195    pub async fn load_checkpoint(
196        &self,
197        _session_id: &SessionId,
198        _checkpoint: &CheckpointData,
199    ) -> AgentResult<Option<CheckpointData>> {
200        Ok(None)
201    }
202
203    pub async fn run<F>(&self, session_id: SessionId, on_event: F) -> AgentResult<RunOutcome>
204    where
205        F: FnMut(RuntimeEvent) -> AgentResult<()> + Send + 'static,
206    {
207        self.runner.run(session_id, on_event).await
208    }
209
210    pub async fn run_turn<F>(
211        &self,
212        session_id: SessionId,
213        user_input: &str,
214        on_event: F,
215    ) -> AgentResult<RunOutcome>
216    where
217        F: FnMut(RuntimeEvent) -> AgentResult<()> + Send + 'static,
218    {
219        self.runner.run_turn(session_id, user_input, on_event).await
220    }
221
222    /// Like `run_turn`, but the user input is pushed as an ephemeral message:
223    /// visible to the LLM for this turn only, auto-removed at turn end.
224    pub async fn run_turn_ephemeral_input<F>(
225        &self,
226        session_id: SessionId,
227        user_input: &str,
228        on_event: F,
229    ) -> AgentResult<RunOutcome>
230    where
231        F: FnMut(RuntimeEvent) -> AgentResult<()> + Send + 'static,
232    {
233        self.runner
234            .run_turn_ephemeral_input(session_id, user_input, on_event)
235            .await
236    }
237
238    pub async fn run_turn_collect(
239        &self,
240        session_id: SessionId,
241        user_input: &str,
242    ) -> AgentResult<(Vec<RuntimeEvent>, RunOutcome)> {
243        self.runner.run_turn_collect(session_id, user_input).await
244    }
245
246    pub async fn add_user_message(
247        &self,
248        session_id: &SessionId,
249        text: impl Into<String>,
250    ) -> AgentResult<()> {
251        let text = text.into();
252        self.with_session_mut(session_id, |session| {
253            session.push_message(MessageRole::User, &text);
254        })
255        .await
256    }
257
258    /// Replace the session's system prompt (the first non-ephemeral System
259    /// message). Hosts that own prompt composition re-bake through this.
260    pub async fn set_system_prompt(
261        &self,
262        session_id: &SessionId,
263        prompt: impl Into<String>,
264    ) -> AgentResult<()> {
265        let prompt = prompt.into();
266        self.with_session_mut(session_id, |session| {
267            session.set_system_prompt(prompt);
268        })
269        .await
270    }
271
272    pub async fn add_system_message(
273        &self,
274        session_id: &SessionId,
275        text: impl Into<String>,
276    ) -> AgentResult<()> {
277        let text = text.into();
278        self.with_session_mut(session_id, |session| {
279            session.push_message(MessageRole::System, &text);
280        })
281        .await
282    }
283
284    pub async fn add_tool_result(
285        &self,
286        session_id: &SessionId,
287        tool_call_id: &str,
288        summary: impl Into<String>,
289    ) -> AgentResult<()> {
290        let summary = summary.into();
291        self.with_session_mut(session_id, |session| {
292            session.push_tool_result(tool_call_id, summary.clone());
293        })
294        .await
295    }
296
297    pub async fn get_messages(
298        &self,
299        session_id: &SessionId,
300    ) -> AgentResult<Vec<crate::types::ChatMessage>> {
301        let session = self.session_or_err(session_id).await?;
302        Ok(session.chat_messages().to_vec())
303    }
304
305    /// Replace the chat messages for a session — only for persistence restore.
306    /// Validates message sequence before applying.
307    ///
308    /// 仅供持久化恢复使用。
309    pub async fn set_messages(
310        &self,
311        session_id: &SessionId,
312        messages: Vec<crate::types::ChatMessage>,
313    ) -> AgentResult<()> {
314        self.with_session_mut(session_id, |session| session.set_chat_messages(messages))
315            .await?
316            .map_err(AgentError::internal)
317    }
318
319    pub async fn validate_session(&self, session_id: &SessionId) -> AgentResult<()> {
320        if self
321            .runner
322            .session_manager
323            .session(session_id)
324            .await
325            .is_none()
326        {
327            return Err(AgentError::session_not_found(session_id.id));
328        }
329        Ok(())
330    }
331
332    pub fn session_store(&self) -> Arc<dyn SessionStore> {
333        self.runner.session_manager.session_store().clone()
334    }
335
336    // ── Observability hook ──
337
338    /// Register a turn-end callback. The callback receives a [`TurnContext`]
339    /// with raw data about the completed turn iteration. Consumers (e.g.
340    /// phi-telemetry) use this to build their own metrics without agent-base
341    /// knowing anything about metrics.
342    pub fn on_turn_end<F>(&self, f: F)
343    where
344        F: Fn(&TurnContext) + Send + Sync + 'static,
345    {
346        self.runner
347            .turn_end_callbacks
348            .write()
349            .unwrap()
350            .push(Arc::new(f));
351    }
352
353    // --- Cancellation support ---
354
355    /// Cancel the currently executing run_turn / run.
356    /// No-op if there is no current execution.
357    pub fn cancel(&self) {
358        self.runner.cancel();
359    }
360
361    /// Reset the cancel token (called automatically before each run_turn)
362    pub fn reset_cancel(&self) {
363        self.runner.reset_cancel();
364    }
365
366    /// Get a clone of the cancel token
367    pub fn cancel_token(&self) -> CancellationToken {
368        self.runner.cancel_token()
369    }
370
371    /// Check if cancellation has been requested
372    pub fn is_cancelled(&self) -> bool {
373        self.runner.is_cancelled()
374    }
375
376    // ── Message Queue (P2) ──
377
378    /// Push a steering message — will be processed at the start of the next turn
379    /// in the current `run_managed()` loop.
380    pub fn steer(&self, message: String) {
381        self.runner.message_queue.steer(message);
382    }
383
384    /// Push a follow-up message — will be processed after the inner turn loop
385    /// stops naturally (no tool calls or max turns).
386    pub fn follow_up(&self, message: String) {
387        self.runner.message_queue.follow_up(message);
388    }
389
390    /// Run the agent in managed mode with message queue support.
391    ///
392    /// This wraps `run_turn()` (or `run()`) in an outer loop: after the inner
393    /// turn loop completes, any follow-up messages are drained and a new inner
394    /// loop is started. Steering messages are drained automatically at each
395    /// iteration of the inner turn loop.
396    ///
397    /// The `on_event` callback receives all events from every inner run.
398    pub async fn run_managed<F>(
399        &self,
400        session_id: SessionId,
401        user_input: &str,
402        on_event: F,
403    ) -> AgentResult<RunOutcome>
404    where
405        F: FnMut(RuntimeEvent) -> AgentResult<()> + Send + 'static,
406    {
407        self.runner
408            .run_managed(session_id, user_input, on_event)
409            .await
410    }
411
412    /// Set the drain mode for the message queues.
413    pub fn set_queue_mode(&self, mode: crate::engine::runtime::message_queue::QueueMode) {
414        self.runner.message_queue.set_mode(mode);
415    }
416}
417
418#[cfg(test)]
419mod tests {
420    use super::*;
421    use crate::llm::ReasoningEffort;
422    use crate::types::{ChatMessage, RuntimeEvent, SessionId};
423    use async_trait::async_trait;
424    use llm_trait::{Capabilities, ChatRequest, ChatResponse, ChatStream, LlmError, ProviderInfo};
425
426    struct StubProvider;
427
428    #[async_trait]
429    impl llm_trait::LlmProvider for StubProvider {
430        async fn stream(&self, _request: ChatRequest) -> Result<ChatStream, LlmError> {
431            Ok(ChatStream::new(Box::pin(futures_util::stream::empty())))
432        }
433
434        async fn chat(&self, _request: ChatRequest) -> Result<ChatResponse, LlmError> {
435            Ok(ChatResponse {
436                content: String::new(),
437                reasoning_content: None,
438                tool_calls: vec![],
439                usage: Default::default(),
440                finish_reason: llm_trait::FinishReason::Stop,
441                raw: None,
442                thinking_signature: None,
443            })
444        }
445
446        fn capabilities(&self) -> Capabilities {
447            Capabilities::default()
448        }
449
450        fn info(&self) -> ProviderInfo {
451            ProviderInfo {
452                name: "stub".to_string(),
453                model: "stub".to_string(),
454                version: None,
455            }
456        }
457    }
458
459    fn runtime() -> AgentRuntime {
460        crate::engine::AgentBuilder::new(Arc::new(StubProvider))
461            .build()
462            .unwrap()
463    }
464
465    #[tokio::test]
466    async fn create_session_and_lookup() {
467        let rt = runtime();
468        let id = rt.create_session().await;
469        assert_eq!(id.id, 1);
470        assert!(rt.session(&id).await.is_some());
471        assert!(rt.session_or_err(&id).await.is_ok());
472    }
473
474    #[tokio::test]
475    async fn add_messages_and_get() {
476        let rt = runtime();
477        let id = rt.create_session().await;
478        rt.add_system_message(&id, "sys").await.unwrap();
479        rt.add_user_message(&id, "hello").await.unwrap();
480        let msgs = rt.get_messages(&id).await.unwrap();
481        assert_eq!(msgs.len(), 2);
482        assert!(matches!(msgs[0], ChatMessage::System { .. }));
483        assert!(matches!(msgs[1], ChatMessage::User { .. }));
484    }
485
486    #[tokio::test]
487    async fn add_tool_result_appends_tool_message() {
488        let rt = runtime();
489        let id = rt.create_session().await;
490        rt.add_tool_result(&id, "call_1", "done").await.unwrap();
491        let msgs = rt.get_messages(&id).await.unwrap();
492        assert_eq!(msgs.len(), 1);
493        assert!(matches!(msgs[0], ChatMessage::Tool { .. }));
494    }
495
496    #[tokio::test]
497    async fn set_messages_replaces_history() {
498        let rt = runtime();
499        let id = rt.create_session().await;
500        rt.add_user_message(&id, "old").await.unwrap();
501        rt.set_messages(
502            &id,
503            vec![ChatMessage::system("sys"), ChatMessage::user("new")],
504        )
505        .await
506        .unwrap();
507        let msgs = rt.get_messages(&id).await.unwrap();
508        assert_eq!(msgs.len(), 2);
509    }
510
511    #[tokio::test]
512    async fn validate_session_errors_for_unknown() {
513        let rt = runtime();
514        let id = rt.create_session().await;
515        assert!(rt.validate_session(&id).await.is_ok());
516        let err = rt.validate_session(&SessionId::new(999)).await.unwrap_err();
517        assert!(matches!(err, AgentError::SessionNotFound(_)));
518    }
519
520    #[test]
521    fn config_and_set_reasoning_effort() {
522        let rt = runtime();
523        assert!(rt.config().system_prompt.is_none());
524
525        rt.set_reasoning_effort_sync(ReasoningEffort::High);
526        let cfg = rt.config();
527        let effort = cfg.reasoning.as_ref().and_then(|r| r.effort.as_ref());
528        assert!(matches!(effort, Some(ReasoningEffort::High)));
529    }
530
531    #[tokio::test]
532    async fn session_store_is_available() {
533        let rt = runtime();
534        assert!(rt.session_store().list().await.unwrap().is_empty());
535    }
536
537    #[tokio::test]
538    async fn emit_and_subscribe_event() {
539        let rt = runtime();
540        let mut rx = rt.subscribe_runtime_events();
541        rt.emit_event(RuntimeEvent::TextDelta {
542            session_id: SessionId::new(1),
543            text: "hi".into(),
544            agent_id: None,
545            trace_id: None,
546        });
547        let ev = rx.recv().await.unwrap();
548        assert!(matches!(ev, RuntimeEvent::TextDelta { .. }));
549    }
550
551    #[tokio::test]
552    async fn cancel_reset_and_is_cancelled() {
553        let rt = runtime();
554        assert!(!rt.is_cancelled());
555        rt.cancel();
556        assert!(rt.is_cancelled());
557        rt.reset_cancel();
558        assert!(!rt.is_cancelled());
559    }
560}