1use 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#[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
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 pub fn subscribe(&self) -> broadcast::Receiver<AgentEvent> {
315 self.inner.event_tx.subscribe()
316 }
317
318 pub fn state(&self) -> AgentState {
321 self.inner.state.lock().expect("state lock").snapshot()
322 }
323
324 pub fn steer(&self, message: AgentMessage) {
326 self.inner
327 .steering_queue
328 .lock()
329 .expect("steer lock")
330 .enqueue(message);
331 }
332
333 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 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 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 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 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 pub async fn prompt_message(&self, message: AgentMessage) -> Result<(), crate::AgentError> {
379 self.prompt_messages(vec![message]).await
380 }
381
382 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 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 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 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 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 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 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 assert_eq!(agent.state().messages.len(), 2);
647 }
648}