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::{AgentEmitter, AgentEvent, 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
266            .opts
267            .convert_to_llm
268            .unwrap_or_else(default_convert_to_llm_fn);
269
270        let state = MutableAgentState {
271            system_prompt: initial.system_prompt.unwrap_or_default(),
272            model: model.clone(),
273            thinking_level,
274            tools: initial.tools.unwrap_or_default(),
275            messages: initial.messages.unwrap_or_default(),
276            is_streaming: false,
277            streaming_message: None,
278            pending_tool_calls: HashSet::new(),
279            error_message: None,
280        };
281
282        let (event_tx, _) = broadcast::channel(256);
283
284        let queue_mode = self.opts.queue_mode.unwrap_or_default();
285        let follow_up_mode = self.opts.follow_up_mode.unwrap_or_default();
286
287        let inner = Inner {
288            state: Arc::new(Mutex::new(state)),
289            convert_to_llm,
290            stream_fn,
291            steering_queue: Arc::new(Mutex::new(PendingMessageQueue::new(queue_mode))),
292            follow_up_queue: Arc::new(Mutex::new(PendingMessageQueue::new(follow_up_mode))),
293            session_id: self.opts.session_id,
294            tool_execution: self.opts.tool_execution.unwrap_or_default(),
295            event_tx,
296            active_run: Mutex::new(None),
297        };
298
299        Ok(Agent {
300            inner: Arc::new(inner),
301        })
302    }
303}
304
305impl Default for AgentBuilder {
306    fn default() -> Self {
307        Self::new()
308    }
309}
310
311impl Agent {
312    /// Subscribe to agent lifecycle events. Returns a `broadcast::Receiver`.
313    /// Mirrors TS `subscribe`.
314    pub fn subscribe(&self) -> broadcast::Receiver<AgentEvent> {
315        self.inner.event_tx.subscribe()
316    }
317
318    /// Current agent state snapshot. Clones tools/messages so callers can't
319    /// mutate internal state. Mirrors TS `get state()`.
320    pub fn state(&self) -> AgentState {
321        self.inner.state.lock().expect("state lock").snapshot()
322    }
323
324    /// Queue a steering message (injected after the current turn's tool batch).
325    pub fn steer(&self, message: AgentMessage) {
326        self.inner
327            .steering_queue
328            .lock()
329            .expect("steer lock")
330            .enqueue(message);
331    }
332
333    /// Queue a follow-up message (injected when the agent would otherwise stop).
334    pub fn follow_up(&self, message: AgentMessage) {
335        self.inner
336            .follow_up_queue
337            .lock()
338            .expect("followup lock")
339            .enqueue(message);
340    }
341
342    /// True when either queue has pending messages.
343    pub fn has_queued_messages(&self) -> bool {
344        let s = self.inner.steering_queue.lock().expect("steer lock");
345        if !s.is_empty() {
346            return true;
347        }
348        let f = self.inner.follow_up_queue.lock().expect("followup lock");
349        !f.is_empty()
350    }
351
352    /// Abort the current run, if one is active. No-op otherwise.
353    pub fn abort(&self) {
354        if let Some(run) = self.inner.active_run.lock().expect("run lock").as_ref() {
355            run.abort.cancel();
356        }
357    }
358
359    /// Resolve when the current run finishes (or immediately if idle).
360    pub async fn wait_for_idle(&self) {
361        let notify = {
362            let guard = self.inner.active_run.lock().expect("run lock");
363            guard.as_ref().map(|r| Arc::clone(&r.done))
364        };
365        if let Some(n) = notify {
366            n.notified().await;
367        }
368    }
369
370    /// Start a new prompt from text. Convenience for `prompt_message`.
371    pub async fn prompt(&self, text: impl Into<String>) -> Result<(), crate::AgentError> {
372        let message =
373            AgentMessage::User(UserMessage::new(UserContent::Text(text.into()), now_ms()));
374        self.prompt_messages(vec![message]).await
375    }
376
377    /// Start a new prompt from a single `AgentMessage`.
378    pub async fn prompt_message(&self, message: AgentMessage) -> Result<(), crate::AgentError> {
379        self.prompt_messages(vec![message]).await
380    }
381
382    /// Start a new prompt from a batch of `AgentMessage`s.
383    pub async fn prompt_messages(
384        &self,
385        messages: Vec<AgentMessage>,
386    ) -> Result<(), crate::AgentError> {
387        self.start_active_run()?;
388        let run_done = self.current_done();
389        let outcome = self.run_prompt(messages).await;
390        self.finish_run();
391        if let Some(n) = run_done {
392            n.notify_waiters();
393        }
394        outcome
395    }
396
397    /// Continue from the current transcript. Errors if the agent is busy or
398    /// the last message is an assistant message with no queued steering/follow-up.
399    pub async fn continue_run(&self) -> Result<(), crate::AgentError> {
400        self.start_active_run()?;
401        let run_done = self.current_done();
402        let outcome = self.run_continue().await;
403        self.finish_run();
404        if let Some(n) = run_done {
405            n.notify_waiters();
406        }
407        outcome
408    }
409
410    /// Clear transcript + runtime state + queues. Errors if a run is active.
411    pub fn reset(&self) -> Result<(), crate::AgentError> {
412        let mut state = self.inner.state.lock().expect("state lock");
413        if self.inner.active_run.lock().expect("run lock").is_some() {
414            return Err(crate::AgentError::State(
415                "Agent is already processing. Wait for completion before resetting.".into(),
416            ));
417        }
418        state.messages.clear();
419        state.is_streaming = false;
420        state.streaming_message = None;
421        state.pending_tool_calls.clear();
422        state.error_message = None;
423        drop(state);
424        let _ = self
425            .inner
426            .steering_queue
427            .lock()
428            .expect("steer lock")
429            .try_drain();
430        let _ = self
431            .inner
432            .follow_up_queue
433            .lock()
434            .expect("followup lock")
435            .try_drain();
436        Ok(())
437    }
438
439    // ---- internals --------------------------------------------------------
440
441    fn start_active_run(&self) -> Result<(), crate::AgentError> {
442        let mut guard = self.inner.active_run.lock().expect("run lock");
443        if guard.is_some() {
444            return Err(crate::AgentError::State(
445                "Agent is already processing a prompt. Use steer() or followUp() to queue messages, or wait for completion.".into(),
446            ));
447        }
448        let abort = CancellationToken::new();
449        let done = Arc::new(Notify::new());
450        *guard = Some(ActiveRun {
451            abort: abort.clone(),
452            done: Arc::clone(&done),
453        });
454
455        let mut state = self.inner.state.lock().expect("state lock");
456        state.is_streaming = true;
457        state.streaming_message = None;
458        state.error_message = None;
459        drop(state);
460        Ok(())
461    }
462
463    fn finish_run(&self) {
464        {
465            let mut state = self.inner.state.lock().expect("state lock");
466            state.is_streaming = false;
467            state.streaming_message = None;
468            state.pending_tool_calls.clear();
469        }
470        let mut guard = self.inner.active_run.lock().expect("run lock");
471        *guard = None;
472    }
473
474    fn current_done(&self) -> Option<Arc<Notify>> {
475        self.inner
476            .active_run
477            .lock()
478            .expect("run lock")
479            .as_ref()
480            .map(|r| Arc::clone(&r.done))
481    }
482
483    fn abort_token(&self) -> CancellationToken {
484        self.inner
485            .active_run
486            .lock()
487            .expect("run lock")
488            .as_ref()
489            .map(|r| r.abort.clone())
490            .unwrap_or_else(CancellationToken::new)
491    }
492
493    fn context_snapshot(&self) -> AgentContext {
494        let state = self.inner.state.lock().expect("state lock");
495        AgentContext {
496            system_prompt: state.system_prompt.clone(),
497            messages: state.messages.clone(),
498            tools: state.tools.clone(),
499        }
500    }
501
502    fn build_config(&self, signal: CancellationToken) -> AgentLoopConfig {
503        let state = self.inner.state.lock().expect("state lock");
504        let steering = Arc::clone(&self.inner.steering_queue);
505        let follow_up = Arc::clone(&self.inner.follow_up_queue);
506        AgentLoopConfig {
507            model: state.model.clone(),
508            convert_to_llm: Arc::clone(&self.inner.convert_to_llm),
509            transform_context: None,
510            get_api_key: None,
511            should_stop_after_turn: None,
512            prepare_next_turn: None,
513            after_tool_results: None,
514            get_steering_messages: Some(Arc::new(move || {
515                let q = Arc::clone(&steering);
516                Box::pin(async move { q.lock().expect("steer lock").try_drain() })
517            })),
518            get_follow_up_messages: Some(Arc::new(move || {
519                let q = Arc::clone(&follow_up);
520                Box::pin(async move { q.lock().expect("followup lock").try_drain() })
521            })),
522            before_tool_call: None,
523            after_tool_call: None,
524            tool_execution: self.inner.tool_execution,
525            thinking_level: state.thinking_level,
526            api_key: None,
527            timeout: None,
528            max_retries: None,
529            max_retry_delay: None,
530            cache_retention: rpi_ai::provider::CacheRetention::default(),
531            session_id: self.inner.session_id.clone(),
532            signal,
533        }
534    }
535
536    async fn run_prompt(&self, messages: Vec<AgentMessage>) -> Result<(), crate::AgentError> {
537        let signal = self.abort_token();
538        let context = self.context_snapshot();
539        let config = self.build_config(signal);
540        let emit: Arc<dyn AgentEmitter> = Arc::new(StatefulEmitter {
541            state: Arc::clone(&self.inner.state),
542            broadcast: BroadcastEmitter::from_sender(self.inner.event_tx.clone()),
543        });
544        let stream_fn = Arc::clone(&self.inner.stream_fn);
545        run_agent_loop(messages, context, config, emit, stream_fn).await?;
546        Ok(())
547    }
548
549    async fn run_continue(&self) -> Result<(), crate::AgentError> {
550        let signal = self.abort_token();
551        let context = self.context_snapshot();
552        if context.messages.is_empty() {
553            return Err(crate::AgentError::State(
554                "No messages to continue from".into(),
555            ));
556        }
557        if context.messages.last().unwrap().is_assistant() {
558            return Err(crate::AgentError::State(
559                "Cannot continue from message role: assistant".into(),
560            ));
561        }
562        let config = self.build_config(signal);
563        let emit: Arc<dyn AgentEmitter> = Arc::new(StatefulEmitter {
564            state: Arc::clone(&self.inner.state),
565            broadcast: BroadcastEmitter::from_sender(self.inner.event_tx.clone()),
566        });
567        let stream_fn = Arc::clone(&self.inner.stream_fn);
568        run_agent_loop_continue(context, config, emit, stream_fn).await?;
569        Ok(())
570    }
571}
572
573fn default_model() -> Model {
574    rpi_ai::model::Model::new(
575        "unknown",
576        "unknown",
577        rpi_ai::types::Api::Other("unknown".into()),
578        "unknown",
579        "",
580    )
581}
582
583fn now_ms() -> i64 {
584    use std::sync::atomic::{AtomicI64, Ordering};
585    static T: AtomicI64 = AtomicI64::new(1);
586    T.fetch_add(1, Ordering::Relaxed)
587}
588
589#[cfg(test)]
590mod tests {
591    use super::*;
592    use rpi_ai::event_stream::create_assistant_message_event_stream;
593    use rpi_ai::provider::Provider;
594    use rpi_ai::providers::faux::{FauxProvider, FauxScript};
595
596    fn faux_stream_fn(provider: Arc<FauxProvider>) -> StreamFn {
597        crate::stream_fn::stream_fn(move |model, ctx, opts| {
598            // The faux provider's stream_simple is async; StreamFn is sync-return.
599            // Bridge by spawning the producer ourselves.
600            let (mut prod, stream) = create_assistant_message_event_stream();
601            let p = Arc::clone(&provider);
602            let model = model.clone();
603            let ctx = ctx.clone();
604            let opts = opts.clone();
605            tokio::spawn(async move {
606                let mut s = p.stream_simple(&model, &ctx, &opts).await;
607                // Drain the real stream into our producer.
608                while let Some(ev) = s.next().await {
609                    if !prod.push(ev) {
610                        break;
611                    }
612                }
613            });
614            stream
615        })
616    }
617
618    #[tokio::test]
619    async fn builder_requires_stream_fn() {
620        let res = AgentBuilder::new().build();
621        assert!(res.is_err(), "build with no stream_fn should error");
622    }
623
624    #[tokio::test]
625    async fn prompt_with_faux_text_collects_events() {
626        let provider = FauxProvider::new(FauxScript::new().with_text("hello"));
627        let sf = faux_stream_fn(provider);
628        let agent = AgentBuilder::new().stream_fn(sf).build().unwrap();
629
630        let mut rx = agent.subscribe();
631        agent.prompt("hi").await.unwrap();
632
633        // Drain the broadcast until AgentEnd.
634        let mut saw_start = false;
635        let mut saw_end = false;
636        while let Ok(ev) = rx.try_recv() {
637            match ev {
638                AgentEvent::AgentStart => saw_start = true,
639                AgentEvent::AgentEnd { .. } => saw_end = true,
640                _ => {}
641            }
642        }
643        assert!(saw_start, "agent_start observed");
644        assert!(saw_end, "agent_end observed");
645        // State: 1 user prompt + 1 assistant reply.
646        assert_eq!(agent.state().messages.len(), 2);
647    }
648}