1use 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#[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#[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
57struct 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 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
119struct ActiveRun {
121 abort: CancellationToken,
122 done: Arc<Notify>,
123}
124
125struct 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
144type SharedQueue = Arc<Mutex<PendingMessageQueue>>;
148
149#[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
168pub 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 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 pub fn subscribe(&self) -> broadcast::Receiver<AgentEvent> {
312 self.inner.event_tx.subscribe()
313 }
314
315 pub fn state(&self) -> AgentState {
318 self.inner.state.lock().expect("state lock").snapshot()
319 }
320
321 pub fn steer(&self, message: AgentMessage) {
323 self.inner
324 .steering_queue
325 .lock()
326 .expect("steer lock")
327 .enqueue(message);
328 }
329
330 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 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 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 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 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 pub async fn prompt_message(&self, message: AgentMessage) -> Result<(), crate::AgentError> {
375 self.prompt_messages(vec![message]).await
376 }
377
378 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 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 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 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 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 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 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 assert_eq!(agent.state().messages.len(), 2);
635 }
636}