Skip to main content

rpi_agent/
agent.rs

1//! Mirrors `packages/agent/src/agent.ts` — the stateful wrapper around the
2//! low-level agent loop. `Agent` owns the current transcript, emits lifecycle
3//! events to subscribers, executes tools, and exposes queueing APIs for
4//! steering and follow-up messages.
5//!
6//! The TS `Agent` class drives the loop via `runWithLifecycle` +
7//! `processEvents`. Rust models the same shape:
8//! - internal `Mutex<MutableAgentState>` for the transcript + runtime flags;
9//! - a `broadcast::Sender<AgentEvent>` so multiple subscribers each get their
10//!   own copy;
11//! - a state-reducing emitter ([`StatefulEmitter`]) that folds each event into
12//!   `MutableAgentState` FIRST, then broadcasts — the direct mirror of TS
13//!   `processEvents` ("reduce internal state, then await listeners");
14//! - a per-run `ActiveRun` (abort handle + completion notify) so `prompt`
15//!   blocks until the run settles and `abort`/`wait_for_idle` can act on it.
16
17use crate::agent_loop::{run_agent_loop, run_agent_loop_continue};
18use crate::events::{AgentEvent, AgentEmitter, BroadcastEmitter};
19use crate::hooks::{default_convert_to_llm_fn, AgentLoopConfig, ConvertToLlm};
20use crate::message::AgentMessage;
21use crate::queue::PendingMessageQueue;
22use crate::stream_fn::{get_default_stream_fn, StreamFn};
23use crate::types::{AgentContext, AgentState, QueueMode, ToolExecutionMode};
24
25use rpi_ai::types::{UserContent, UserMessage};
26use rpi_ai::Model;
27use std::collections::HashSet;
28use std::sync::{Arc, Mutex};
29use tokio::sync::{broadcast, Notify};
30use tokio_util::sync::CancellationToken;
31
32/// Options for constructing an [`Agent`]. Mirrors TS `AgentOptions`. Only
33/// `stream_fn` is required (or a default must have been installed process-wide
34/// via [`crate::stream_fn::set_default_stream_fn`]).
35#[derive(Default)]
36pub struct AgentOptions {
37    pub initial_state: Option<InitialState>,
38    pub convert_to_llm: Option<ConvertToLlm>,
39    pub stream_fn: Option<StreamFn>,
40    pub thinking_level: Option<rpi_ai::types::ThinkingLevel>,
41    pub queue_mode: Option<QueueMode>,
42    pub follow_up_mode: Option<QueueMode>,
43    pub tool_execution: Option<ToolExecutionMode>,
44    pub session_id: Option<String>,
45}
46
47/// Subset of `AgentState` settable at construction.
48#[derive(Default)]
49pub struct InitialState {
50    pub system_prompt: Option<String>,
51    pub model: Option<Model>,
52    pub thinking_level: Option<rpi_ai::types::ThinkingLevel>,
53    pub tools: Option<Vec<Arc<dyn crate::agent_tool::AgentTool>>>,
54    pub messages: Option<Vec<AgentMessage>>,
55}
56
57/// The owned mutable state. Mirrors TS `MutableAgentState`. Guarded by a mutex.
58struct MutableAgentState {
59    system_prompt: String,
60    model: Model,
61    thinking_level: rpi_ai::types::ThinkingLevel,
62    tools: Vec<Arc<dyn crate::agent_tool::AgentTool>>,
63    messages: Vec<AgentMessage>,
64    is_streaming: bool,
65    streaming_message: Option<AgentMessage>,
66    pending_tool_calls: HashSet<String>,
67    error_message: Option<String>,
68}
69
70impl MutableAgentState {
71    fn snapshot(&self) -> AgentState {
72        AgentState {
73            system_prompt: self.system_prompt.clone(),
74            model: self.model.clone(),
75            thinking_level: self.thinking_level,
76            tools: self.tools.clone(),
77            messages: self.messages.clone(),
78            is_streaming: self.is_streaming,
79            streaming_message: self.streaming_message.clone(),
80            pending_tool_calls: self.pending_tool_calls.clone(),
81            error_message: self.error_message.clone(),
82        }
83    }
84
85    /// Reduce an event into state. Mirrors the switch in TS `processEvents`.
86    fn reduce(&mut self, event: &AgentEvent) {
87        match event {
88            AgentEvent::MessageStart { message } => {
89                self.streaming_message = Some(message.clone());
90            }
91            AgentEvent::MessageUpdate { message, .. } => {
92                self.streaming_message = Some(message.clone());
93            }
94            AgentEvent::MessageEnd { message } => {
95                self.streaming_message = None;
96                self.messages.push(message.clone());
97            }
98            AgentEvent::ToolExecutionStart { tool_call_id, .. } => {
99                self.pending_tool_calls.insert(tool_call_id.clone());
100            }
101            AgentEvent::ToolExecutionEnd { tool_call_id, .. } => {
102                self.pending_tool_calls.remove(tool_call_id);
103            }
104            AgentEvent::TurnEnd { message, .. } => {
105                if let Some(am) = message.as_assistant() {
106                    if am.error_message.is_some() {
107                        self.error_message = am.error_message.clone();
108                    }
109                }
110            }
111            AgentEvent::AgentEnd { .. } => {
112                self.streaming_message = None;
113            }
114            _ => {}
115        }
116    }
117}
118
119/// Per-run handle. Replaces TS `ActiveRun`.
120struct ActiveRun {
121    abort: CancellationToken,
122    done: Arc<Notify>,
123}
124
125/// A state-reducing emitter: folds each event into `MutableAgentState` FIRST,
126/// then broadcasts to subscribers. The direct mirror of TS `processEvents`.
127struct StatefulEmitter {
128    state: Arc<Mutex<MutableAgentState>>,
129    broadcast: BroadcastEmitter,
130}
131
132impl AgentEmitter for StatefulEmitter {
133    fn emit(&self, event: AgentEvent) -> futures::future::BoxFuture<'static, ()> {
134        self.state.lock().expect("state lock").reduce(&event);
135        self.broadcast.try_emit(event);
136        Box::pin(async {})
137    }
138    fn try_emit(&self, event: AgentEvent) {
139        self.state.lock().expect("state lock").reduce(&event);
140        self.broadcast.try_emit(event);
141    }
142}
143
144/// Shared queue storage: `Arc<Mutex<PendingMessageQueue>>` so the
145/// `get_steering_messages` / `get_follow_up_messages` hook closures (which must
146/// be `'static + Send + Sync`) can capture a clone.
147type SharedQueue = Arc<Mutex<PendingMessageQueue>>;
148
149/// A stateful agent. Clone shares the same inner state + event channel (like
150/// holding a second reference to the TS `Agent` instance).
151#[derive(Clone)]
152pub struct Agent {
153    inner: Arc<Inner>,
154}
155
156struct Inner {
157    state: Arc<Mutex<MutableAgentState>>,
158    convert_to_llm: ConvertToLlm,
159    stream_fn: StreamFn,
160    steering_queue: SharedQueue,
161    follow_up_queue: SharedQueue,
162    session_id: Option<String>,
163    tool_execution: ToolExecutionMode,
164    event_tx: broadcast::Sender<AgentEvent>,
165    active_run: Mutex<Option<ActiveRun>>,
166}
167
168/// Builder for [`Agent`] with a fluent API. Mirrors the TS `new Agent(opts)`.
169pub struct AgentBuilder {
170    opts: AgentOptions,
171}
172
173impl AgentBuilder {
174    pub fn new() -> Self {
175        Self {
176            opts: AgentOptions::default(),
177        }
178    }
179
180    pub fn model(mut self, model: Model) -> Self {
181        self.opts
182            .initial_state
183            .get_or_insert_with(InitialState::default)
184            .model = Some(model);
185        self
186    }
187
188    pub fn system_prompt(mut self, prompt: impl Into<String>) -> Self {
189        self.opts
190            .initial_state
191            .get_or_insert_with(InitialState::default)
192            .system_prompt = Some(prompt.into());
193        self
194    }
195
196    pub fn thinking_level(mut self, level: rpi_ai::types::ThinkingLevel) -> Self {
197        self.opts
198            .initial_state
199            .get_or_insert_with(InitialState::default)
200            .thinking_level = Some(level);
201        self
202    }
203
204    pub fn tools(mut self, tools: Vec<Arc<dyn crate::agent_tool::AgentTool>>) -> Self {
205        self.opts
206            .initial_state
207            .get_or_insert_with(InitialState::default)
208            .tools = Some(tools);
209        self
210    }
211
212    pub fn messages(mut self, messages: Vec<AgentMessage>) -> Self {
213        self.opts
214            .initial_state
215            .get_or_insert_with(InitialState::default)
216            .messages = Some(messages);
217        self
218    }
219
220    pub fn stream_fn(mut self, stream_fn: StreamFn) -> Self {
221        self.opts.stream_fn = Some(stream_fn);
222        self
223    }
224
225    pub fn convert_to_llm(mut self, f: ConvertToLlm) -> Self {
226        self.opts.convert_to_llm = Some(f);
227        self
228    }
229
230    pub fn queue_mode(mut self, mode: QueueMode) -> Self {
231        self.opts.queue_mode = Some(mode);
232        self
233    }
234
235    pub fn follow_up_mode(mut self, mode: QueueMode) -> Self {
236        self.opts.follow_up_mode = Some(mode);
237        self
238    }
239
240    pub fn tool_execution(mut self, mode: ToolExecutionMode) -> Self {
241        self.opts.tool_execution = Some(mode);
242        self
243    }
244
245    pub fn session_id(mut self, id: impl Into<String>) -> Self {
246        self.opts.session_id = Some(id.into());
247        self
248    }
249
250    /// Build the `Agent`. Errors if no `stream_fn` was supplied and no default
251    /// has been installed process-wide.
252    pub fn build(self) -> Result<Agent, crate::AgentError> {
253        let stream_fn = match self.opts.stream_fn {
254            Some(f) => f,
255            None => get_default_stream_fn()?,
256        };
257
258        let initial = self.opts.initial_state.unwrap_or_default();
259        let model = initial.model.unwrap_or_else(default_model);
260        let thinking_level = self
261            .opts
262            .thinking_level
263            .or(initial.thinking_level)
264            .unwrap_or(rpi_ai::types::ThinkingLevel::Off);
265        let convert_to_llm = self.opts.convert_to_llm.unwrap_or_else(default_convert_to_llm_fn);
266
267        let state = MutableAgentState {
268            system_prompt: initial.system_prompt.unwrap_or_default(),
269            model: model.clone(),
270            thinking_level,
271            tools: initial.tools.unwrap_or_default(),
272            messages: initial.messages.unwrap_or_default(),
273            is_streaming: false,
274            streaming_message: None,
275            pending_tool_calls: HashSet::new(),
276            error_message: None,
277        };
278
279        let (event_tx, _) = broadcast::channel(256);
280
281        let queue_mode = self.opts.queue_mode.unwrap_or_default();
282        let follow_up_mode = self.opts.follow_up_mode.unwrap_or_default();
283
284        let inner = Inner {
285            state: Arc::new(Mutex::new(state)),
286            convert_to_llm,
287            stream_fn,
288            steering_queue: Arc::new(Mutex::new(PendingMessageQueue::new(queue_mode))),
289            follow_up_queue: Arc::new(Mutex::new(PendingMessageQueue::new(follow_up_mode))),
290            session_id: self.opts.session_id,
291            tool_execution: self.opts.tool_execution.unwrap_or_default(),
292            event_tx,
293            active_run: Mutex::new(None),
294        };
295
296        Ok(Agent {
297            inner: Arc::new(inner),
298        })
299    }
300}
301
302impl Default for AgentBuilder {
303    fn default() -> Self {
304        Self::new()
305    }
306}
307
308impl Agent {
309    /// Subscribe to agent lifecycle events. Returns a `broadcast::Receiver`.
310    /// Mirrors TS `subscribe`.
311    pub fn subscribe(&self) -> broadcast::Receiver<AgentEvent> {
312        self.inner.event_tx.subscribe()
313    }
314
315    /// Current agent state snapshot. Clones tools/messages so callers can't
316    /// mutate internal state. Mirrors TS `get state()`.
317    pub fn state(&self) -> AgentState {
318        self.inner.state.lock().expect("state lock").snapshot()
319    }
320
321    /// Queue a steering message (injected after the current turn's tool batch).
322    pub fn steer(&self, message: AgentMessage) {
323        self.inner
324            .steering_queue
325            .lock()
326            .expect("steer lock")
327            .enqueue(message);
328    }
329
330    /// Queue a follow-up message (injected when the agent would otherwise stop).
331    pub fn follow_up(&self, message: AgentMessage) {
332        self.inner
333            .follow_up_queue
334            .lock()
335            .expect("followup lock")
336            .enqueue(message);
337    }
338
339    /// True when either queue has pending messages.
340    pub fn has_queued_messages(&self) -> bool {
341        let s = self.inner.steering_queue.lock().expect("steer lock");
342        if !s.is_empty() {
343            return true;
344        }
345        let f = self.inner.follow_up_queue.lock().expect("followup lock");
346        !f.is_empty()
347    }
348
349    /// Abort the current run, if one is active. No-op otherwise.
350    pub fn abort(&self) {
351        if let Some(run) = self.inner.active_run.lock().expect("run lock").as_ref() {
352            run.abort.cancel();
353        }
354    }
355
356    /// Resolve when the current run finishes (or immediately if idle).
357    pub async fn wait_for_idle(&self) {
358        let notify = {
359            let guard = self.inner.active_run.lock().expect("run lock");
360            guard.as_ref().map(|r| Arc::clone(&r.done))
361        };
362        if let Some(n) = notify {
363            n.notified().await;
364        }
365    }
366
367    /// Start a new prompt from text. Convenience for `prompt_message`.
368    pub async fn prompt(&self, text: impl Into<String>) -> Result<(), crate::AgentError> {
369        let message = AgentMessage::User(UserMessage::new(UserContent::Text(text.into()), now_ms()));
370        self.prompt_messages(vec![message]).await
371    }
372
373    /// Start a new prompt from a single `AgentMessage`.
374    pub async fn prompt_message(&self, message: AgentMessage) -> Result<(), crate::AgentError> {
375        self.prompt_messages(vec![message]).await
376    }
377
378    /// Start a new prompt from a batch of `AgentMessage`s.
379    pub async fn prompt_messages(
380        &self,
381        messages: Vec<AgentMessage>,
382    ) -> Result<(), crate::AgentError> {
383        self.start_active_run()?;
384        let run_done = self.current_done();
385        let outcome = self.run_prompt(messages).await;
386        self.finish_run();
387        if let Some(n) = run_done {
388            n.notify_waiters();
389        }
390        outcome
391    }
392
393    /// Continue from the current transcript. Errors if the agent is busy or
394    /// the last message is an assistant message with no queued steering/follow-up.
395    pub async fn continue_run(&self) -> Result<(), crate::AgentError> {
396        self.start_active_run()?;
397        let run_done = self.current_done();
398        let outcome = self.run_continue().await;
399        self.finish_run();
400        if let Some(n) = run_done {
401            n.notify_waiters();
402        }
403        outcome
404    }
405
406    /// Clear transcript + runtime state + queues. Errors if a run is active.
407    pub fn reset(&self) -> Result<(), crate::AgentError> {
408        let mut state = self.inner.state.lock().expect("state lock");
409        if self.inner.active_run.lock().expect("run lock").is_some() {
410            return Err(crate::AgentError::State(
411                "Agent is already processing. Wait for completion before resetting.".into(),
412            ));
413        }
414        state.messages.clear();
415        state.is_streaming = false;
416        state.streaming_message = None;
417        state.pending_tool_calls.clear();
418        state.error_message = None;
419        drop(state);
420        let _ = self.inner.steering_queue.lock().expect("steer lock").try_drain();
421        let _ = self
422            .inner
423            .follow_up_queue
424            .lock()
425            .expect("followup lock")
426            .try_drain();
427        Ok(())
428    }
429
430    // ---- internals --------------------------------------------------------
431
432    fn start_active_run(&self) -> Result<(), crate::AgentError> {
433        let mut guard = self.inner.active_run.lock().expect("run lock");
434        if guard.is_some() {
435            return Err(crate::AgentError::State(
436                "Agent is already processing a prompt. Use steer() or followUp() to queue messages, or wait for completion.".into(),
437            ));
438        }
439        let abort = CancellationToken::new();
440        let done = Arc::new(Notify::new());
441        *guard = Some(ActiveRun {
442            abort: abort.clone(),
443            done: Arc::clone(&done),
444        });
445
446        let mut state = self.inner.state.lock().expect("state lock");
447        state.is_streaming = true;
448        state.streaming_message = None;
449        state.error_message = None;
450        drop(state);
451        Ok(())
452    }
453
454    fn finish_run(&self) {
455        {
456            let mut state = self.inner.state.lock().expect("state lock");
457            state.is_streaming = false;
458            state.streaming_message = None;
459            state.pending_tool_calls.clear();
460        }
461        let mut guard = self.inner.active_run.lock().expect("run lock");
462        *guard = None;
463    }
464
465    fn current_done(&self) -> Option<Arc<Notify>> {
466        self.inner
467            .active_run
468            .lock()
469            .expect("run lock")
470            .as_ref()
471            .map(|r| Arc::clone(&r.done))
472    }
473
474    fn abort_token(&self) -> CancellationToken {
475        self.inner
476            .active_run
477            .lock()
478            .expect("run lock")
479            .as_ref()
480            .map(|r| r.abort.clone())
481            .unwrap_or_else(CancellationToken::new)
482    }
483
484    fn context_snapshot(&self) -> AgentContext {
485        let state = self.inner.state.lock().expect("state lock");
486        AgentContext {
487            system_prompt: state.system_prompt.clone(),
488            messages: state.messages.clone(),
489            tools: state.tools.clone(),
490        }
491    }
492
493    fn build_config(&self, signal: CancellationToken) -> AgentLoopConfig {
494        let state = self.inner.state.lock().expect("state lock");
495        let steering = Arc::clone(&self.inner.steering_queue);
496        let follow_up = Arc::clone(&self.inner.follow_up_queue);
497        AgentLoopConfig {
498            model: state.model.clone(),
499            convert_to_llm: Arc::clone(&self.inner.convert_to_llm),
500            transform_context: None,
501            get_api_key: None,
502            should_stop_after_turn: None,
503            prepare_next_turn: None,
504            get_steering_messages: Some(Arc::new(move || {
505                let q = Arc::clone(&steering);
506                Box::pin(async move { q.lock().expect("steer lock").try_drain() })
507            })),
508            get_follow_up_messages: Some(Arc::new(move || {
509                let q = Arc::clone(&follow_up);
510                Box::pin(async move { q.lock().expect("followup lock").try_drain() })
511            })),
512            before_tool_call: None,
513            after_tool_call: None,
514            tool_execution: self.inner.tool_execution,
515            thinking_level: state.thinking_level,
516            api_key: None,
517            timeout: None,
518            max_retries: None,
519            max_retry_delay: None,
520            cache_retention: rpi_ai::provider::CacheRetention::default(),
521            session_id: self.inner.session_id.clone(),
522            signal,
523        }
524    }
525
526    async fn run_prompt(&self, messages: Vec<AgentMessage>) -> Result<(), crate::AgentError> {
527        let signal = self.abort_token();
528        let context = self.context_snapshot();
529        let config = self.build_config(signal);
530        let emit: Arc<dyn AgentEmitter> = Arc::new(StatefulEmitter {
531            state: Arc::clone(&self.inner.state),
532            broadcast: BroadcastEmitter::from_sender(self.inner.event_tx.clone()),
533        });
534        let stream_fn = Arc::clone(&self.inner.stream_fn);
535        run_agent_loop(messages, context, config, emit, stream_fn).await?;
536        Ok(())
537    }
538
539    async fn run_continue(&self) -> Result<(), crate::AgentError> {
540        let signal = self.abort_token();
541        let context = self.context_snapshot();
542        if context.messages.is_empty() {
543            return Err(crate::AgentError::State("No messages to continue from".into()));
544        }
545        if context.messages.last().unwrap().is_assistant() {
546            return Err(crate::AgentError::State(
547                "Cannot continue from message role: assistant".into(),
548            ));
549        }
550        let config = self.build_config(signal);
551        let emit: Arc<dyn AgentEmitter> = Arc::new(StatefulEmitter {
552            state: Arc::clone(&self.inner.state),
553            broadcast: BroadcastEmitter::from_sender(self.inner.event_tx.clone()),
554        });
555        let stream_fn = Arc::clone(&self.inner.stream_fn);
556        run_agent_loop_continue(context, config, emit, stream_fn).await?;
557        Ok(())
558    }
559}
560
561fn default_model() -> Model {
562    rpi_ai::model::Model::new(
563        "unknown",
564        "unknown",
565        rpi_ai::types::Api::Other("unknown".into()),
566        "unknown",
567        "",
568    )
569}
570
571fn now_ms() -> i64 {
572    use std::sync::atomic::{AtomicI64, Ordering};
573    static T: AtomicI64 = AtomicI64::new(1);
574    T.fetch_add(1, Ordering::Relaxed)
575}
576
577#[cfg(test)]
578mod tests {
579    use super::*;
580    use rpi_ai::event_stream::create_assistant_message_event_stream;
581    use rpi_ai::providers::faux::{FauxScript, FauxProvider};
582    use rpi_ai::provider::Provider;
583
584    fn faux_stream_fn(provider: Arc<FauxProvider>) -> StreamFn {
585        crate::stream_fn::stream_fn(move |model, ctx, opts| {
586            // The faux provider's stream_simple is async; StreamFn is sync-return.
587            // Bridge by spawning the producer ourselves.
588            let (mut prod, stream) = create_assistant_message_event_stream();
589            let p = Arc::clone(&provider);
590            let model = model.clone();
591            let ctx = ctx.clone();
592            let opts = opts.clone();
593            tokio::spawn(async move {
594                let mut s = p.stream_simple(&model, &ctx, &opts).await;
595                // Drain the real stream into our producer.
596                while let Some(ev) = s.next().await {
597                    if !prod.push(ev) {
598                        break;
599                    }
600                }
601            });
602            stream
603        })
604    }
605
606    #[tokio::test]
607    async fn builder_requires_stream_fn() {
608        let res = AgentBuilder::new().build();
609        assert!(res.is_err(), "build with no stream_fn should error");
610    }
611
612    #[tokio::test]
613    async fn prompt_with_faux_text_collects_events() {
614        let provider = FauxProvider::new(FauxScript::new().with_text("hello"));
615        let sf = faux_stream_fn(provider);
616        let agent = AgentBuilder::new().stream_fn(sf).build().unwrap();
617
618        let mut rx = agent.subscribe();
619        agent.prompt("hi").await.unwrap();
620
621        // Drain the broadcast until AgentEnd.
622        let mut saw_start = false;
623        let mut saw_end = false;
624        while let Ok(ev) = rx.try_recv() {
625            match ev {
626                AgentEvent::AgentStart => saw_start = true,
627                AgentEvent::AgentEnd { .. } => saw_end = true,
628                _ => {}
629            }
630        }
631        assert!(saw_start, "agent_start observed");
632        assert!(saw_end, "agent_end observed");
633        // State: 1 user prompt + 1 assistant reply.
634        assert_eq!(agent.state().messages.len(), 2);
635    }
636}