1use std::sync::Arc;
2
3use tokio::sync::{broadcast, mpsc};
4
5use funera_core::chat::message::{FuneraMessage, MsgVariant, Role, TextMessage};
6use funera_core::chat::session::FuneraSession;
7use funera_core::event_bus::env_state_bus::{EnvStateBus, EnvStateEvent};
8use funera_core::middleware::EventSenderFn;
9#[cfg(feature = "middleware")]
10use funera_core::middleware::{ErrorsEnabled, MiddlewareChain};
11use funera_core::provider::ChatProvider;
12use funera_core::re_act::ReActLoopConfig;
13
14use crate::dispatcher::{CallbackDispatcher, CallbackRegistry};
15use crate::error::OrchestrateError;
16use crate::event::{AgentEvent, RawAgentEvent};
17use crate::response::{ChatResponse, ToolCallInfo};
18use crate::runtime::{AgentRuntime, Idle};
19use crate::send_handle::{FireStreamHandle, SendHandle, SendStreamHandle};
20
21pub struct AgentBuilder {
41 system_prompt: Option<String>,
42 callbacks: CallbackRegistry,
43}
44
45impl Default for AgentBuilder {
46 fn default() -> Self {
47 Self::new()
48 }
49}
50
51impl AgentBuilder {
52 pub fn new() -> Self {
53 Self {
54 system_prompt: None,
55 callbacks: CallbackRegistry::new(),
56 }
57 }
58
59 pub fn system_prompt(mut self, prompt: impl Into<String>) -> Self {
61 self.system_prompt = Some(prompt.into());
62 self
63 }
64
65 pub fn on_token<F>(mut self, f: F) -> Self
67 where
68 F: Fn(String) + Send + Sync + 'static,
69 {
70 self.callbacks.add(Arc::new(move |event| {
71 if let AgentEvent::Text(t) = event {
72 f(t);
73 }
74 }));
75 self
76 }
77
78 pub fn on_tool_call<F>(mut self, f: F) -> Self
80 where
81 F: Fn(String, serde_json::Value) + Send + Sync + 'static,
82 {
83 self.callbacks.add(Arc::new(move |event| {
84 if let AgentEvent::ToolCallRequest { name, args, .. } = event {
85 f(name, args);
86 }
87 }));
88 self
89 }
90
91 pub fn on_tool_result<F>(mut self, f: F) -> Self
93 where
94 F: Fn(String, Result<String, String>) + Send + Sync + 'static,
95 {
96 self.callbacks.add(Arc::new(move |event| {
97 if let AgentEvent::ToolCallResult { name, result, .. } = event {
98 f(name, result);
99 }
100 }));
101 self
102 }
103
104 pub fn on_turn_start<F>(mut self, f: F) -> Self
106 where
107 F: Fn() + Send + Sync + 'static,
108 {
109 self.callbacks.add(Arc::new(move |event| {
110 if matches!(event, AgentEvent::TurnStart) {
111 f();
112 }
113 }));
114 self
115 }
116
117 pub fn on_turn_end<F>(mut self, f: F) -> Self
119 where
120 F: Fn() + Send + Sync + 'static,
121 {
122 self.callbacks.add(Arc::new(move |event| {
123 if matches!(event, AgentEvent::TurnEnd { .. }) {
124 f();
125 }
126 }));
127 self
128 }
129
130 pub fn on_event<F>(mut self, f: F) -> Self
132 where
133 F: Fn(AgentEvent) + Send + Sync + 'static,
134 {
135 self.callbacks.add(Arc::new(f));
136 self
137 }
138
139 pub fn build(self) -> Agent {
141 let (event_tx, _) = broadcast::channel(256);
142 let (raw_event_tx, _) = broadcast::channel(256);
143 Agent {
144 system_prompt: self.system_prompt,
145 callbacks: Arc::new(self.callbacks),
146 event_tx,
147 raw_event_tx,
148 }
149 }
150}
151
152pub struct Agent {
202 pub(crate) system_prompt: Option<String>,
203 pub(crate) callbacks: Arc<CallbackRegistry>,
204 pub(crate) event_tx: broadcast::Sender<AgentEvent>,
205 pub(crate) raw_event_tx: broadcast::Sender<RawAgentEvent>,
206}
207
208impl Agent {
209 pub fn builder() -> AgentBuilder {
211 AgentBuilder::new()
212 }
213
214 pub fn subscribe_events(&self) -> broadcast::Receiver<AgentEvent> {
219 self.event_tx.subscribe()
220 }
221
222 pub fn subscribe_raw_events(&self) -> broadcast::Receiver<RawAgentEvent> {
234 self.raw_event_tx.subscribe()
235 }
236
237 pub async fn fire<P: ChatProvider, S>(
242 &self,
243 msg: impl Into<String>,
244 runtime: &AgentRuntime<P, S>,
245 ) -> Result<ChatResponse, OrchestrateError> {
246 let text = msg.into();
247 let mut event_rx = self.subscribe_events();
248
249 let (env_state_bus, turn_highway_handle) = EnvStateBus::new();
250 let env_state_tx = env_state_bus.env_state_tx.clone();
251 let env_state_rx = env_state_bus.subscribe();
252 env_state_bus.start_turn_highway();
253
254 let _dispatcher = CallbackDispatcher::new(
255 env_state_rx,
256 self.event_tx.clone(),
257 self.raw_event_tx.clone(),
258 );
259
260 let _ = env_state_tx.send(EnvStateEvent::SessionStart);
261
262 let session_tx = funera_core::chat::session::spawn_session_actor();
264 let session = FuneraSession::new(session_tx);
265 if let Some(ref sys) = self.system_prompt {
266 session.push_message(FuneraMessage::new(
267 Role::System,
268 MsgVariant::Text(TextMessage {
269 text: sys.clone().into(),
270 reasoning_content: None,
271 }),
272 ));
273 }
274
275 let init_msg = FuneraMessage::new(
276 Role::User,
277 MsgVariant::Text(TextMessage {
278 text: text.into(),
279 reasoning_content: None,
280 }),
281 );
282
283 let mut config = ReActLoopConfig::new(
284 runtime.channel_buffer(),
285 runtime.max_iterations(),
286 runtime.env_watcher(),
287 env_state_tx.clone(),
288 turn_highway_handle,
289 );
290 #[cfg(feature = "tool")]
291 {
292 config = config.with_tool_bus(runtime.tool_bus.clone());
293 }
294
295 let event_sender = build_event_sender(self.callbacks.clone(), self.event_tx.clone());
296
297 let result = session
298 .react_loop::<P, AgentEvent>(
299 init_msg,
300 config,
301 env_state_tx.clone(),
302 middleware_opt(runtime),
303 Some(event_sender),
304 )
305 .await;
306
307 let _ = env_state_tx.send(EnvStateEvent::SessionClosed);
308 aggregate_response(&mut event_rx, result).await
309 }
310
311 pub async fn fire_stream<P: ChatProvider, S>(
316 &self,
317 msg: impl Into<String>,
318 runtime: &AgentRuntime<P, S>,
319 ) -> Result<FireStreamHandle, OrchestrateError> {
320 let text = msg.into();
321 let event_rx = self.subscribe_events();
322
323 let (relay_tx, stream_rx) = mpsc::channel(256);
325 let relay_event_rx = self.subscribe_events();
326 tokio::spawn(async move {
327 relay_broadcast_to_mpsc(relay_event_rx, relay_tx).await;
328 });
329
330 let (env_state_bus, turn_highway_handle) = EnvStateBus::new();
331 let env_state_tx = env_state_bus.env_state_tx.clone();
332 let env_state_rx = env_state_bus.subscribe();
333 env_state_bus.start_turn_highway();
334
335 let _dispatcher = CallbackDispatcher::new(
336 env_state_rx,
337 self.event_tx.clone(),
338 self.raw_event_tx.clone(),
339 );
340 let _ = env_state_tx.send(EnvStateEvent::SessionStart);
341
342 let session_tx = funera_core::chat::session::spawn_session_actor();
343 let session = FuneraSession::new(session_tx);
344 if let Some(ref sys) = self.system_prompt {
345 session.push_message(FuneraMessage::new(
346 Role::System,
347 MsgVariant::Text(TextMessage {
348 text: sys.clone().into(),
349 reasoning_content: None,
350 }),
351 ));
352 }
353 let init_msg = FuneraMessage::new(
354 Role::User,
355 MsgVariant::Text(TextMessage {
356 text: text.into(),
357 reasoning_content: None,
358 }),
359 );
360 let mut config = ReActLoopConfig::new(
361 runtime.channel_buffer(),
362 runtime.max_iterations(),
363 runtime.env_watcher(),
364 env_state_tx.clone(),
365 turn_highway_handle,
366 );
367 #[cfg(feature = "tool")]
368 {
369 config = config.with_tool_bus(runtime.tool_bus.clone());
370 }
371 let event_sender = build_event_sender(self.callbacks.clone(), self.event_tx.clone());
372
373 let mw = middleware_opt(runtime);
375 let env_tx = env_state_tx.clone();
376 let handle = tokio::spawn(async move {
377 session
378 .react_loop::<P, AgentEvent>(init_msg, config, env_tx, mw, Some(event_sender))
379 .await
380 });
381
382 Ok(FireStreamHandle {
383 handle,
384 event_rx,
385 stream_rx,
386 env_state_tx,
387 })
388 }
389
390 pub async fn send<P: ChatProvider>(
397 &self,
398 msg: impl Into<String>,
399 runtime: AgentRuntime<P, Idle>,
400 ) -> Result<SendHandle<P>, OrchestrateError> {
401 let text = msg.into();
402 let event_rx = self.subscribe_events();
403
404 let (env_state_bus, turn_highway_handle) = EnvStateBus::new();
405 let env_state_tx = env_state_bus.env_state_tx.clone();
406 let env_state_rx = env_state_bus.subscribe();
407 env_state_bus.start_turn_highway();
408
409 let _dispatcher = CallbackDispatcher::new(
410 env_state_rx,
411 self.event_tx.clone(),
412 self.raw_event_tx.clone(),
413 );
414 let _ = env_state_tx.send(EnvStateEvent::SessionStart);
415
416 let session = FuneraSession::new(runtime.session_tx());
417 if let Some(ref sys) = self.system_prompt {
418 let msgs = session.session_context().await;
419 if msgs.is_empty() {
420 session.push_message(FuneraMessage::new(
421 Role::System,
422 MsgVariant::Text(TextMessage {
423 text: sys.clone().into(),
424 reasoning_content: None,
425 }),
426 ));
427 }
428 }
429 let init_msg = FuneraMessage::new(
430 Role::User,
431 MsgVariant::Text(TextMessage {
432 text: text.into(),
433 reasoning_content: None,
434 }),
435 );
436 let mut config = ReActLoopConfig::new(
437 runtime.channel_buffer(),
438 runtime.max_iterations(),
439 runtime.env_watcher(),
440 env_state_tx.clone(),
441 turn_highway_handle,
442 );
443 #[cfg(feature = "tool")]
444 {
445 config = config.with_tool_bus(runtime.tool_bus.clone());
446 }
447 let event_sender = build_event_sender(self.callbacks.clone(), self.event_tx.clone());
448
449 let env_tx = env_state_tx.clone();
450 let mw = middleware_opt(&runtime);
451 let handle = tokio::spawn(async move {
452 session
453 .react_loop::<P, AgentEvent>(init_msg, config, env_tx, mw, Some(event_sender))
454 .await
455 });
456
457 Ok(SendHandle {
458 runtime: runtime.into_acquired(),
459 handle,
460 event_rx,
461 env_state_tx,
462 })
463 }
464
465 pub async fn send_stream<P: ChatProvider>(
467 &self,
468 msg: impl Into<String>,
469 runtime: AgentRuntime<P, Idle>,
470 ) -> Result<SendStreamHandle<P>, OrchestrateError> {
471 let text = msg.into();
472 let event_rx = self.subscribe_events();
473
474 let (relay_tx, stream_rx) = mpsc::channel(256);
476 let relay_event_rx = self.subscribe_events();
477 tokio::spawn(async move {
478 relay_broadcast_to_mpsc(relay_event_rx, relay_tx).await;
479 });
480
481 let (env_state_bus, turn_highway_handle) = EnvStateBus::new();
482 let env_state_tx = env_state_bus.env_state_tx.clone();
483 let env_state_rx = env_state_bus.subscribe();
484 env_state_bus.start_turn_highway();
485
486 let _dispatcher = CallbackDispatcher::new(
487 env_state_rx,
488 self.event_tx.clone(),
489 self.raw_event_tx.clone(),
490 );
491 let _ = env_state_tx.send(EnvStateEvent::SessionStart);
492
493 let session = FuneraSession::new(runtime.session_tx());
494 if let Some(ref sys) = self.system_prompt {
495 let msgs = session.session_context().await;
496 if msgs.is_empty() {
497 session.push_message(FuneraMessage::new(
498 Role::System,
499 MsgVariant::Text(TextMessage {
500 text: sys.clone().into(),
501 reasoning_content: None,
502 }),
503 ));
504 }
505 }
506 let init_msg = FuneraMessage::new(
507 Role::User,
508 MsgVariant::Text(TextMessage {
509 text: text.into(),
510 reasoning_content: None,
511 }),
512 );
513 let mut config = ReActLoopConfig::new(
514 runtime.channel_buffer(),
515 runtime.max_iterations(),
516 runtime.env_watcher(),
517 env_state_tx.clone(),
518 turn_highway_handle,
519 );
520 #[cfg(feature = "tool")]
521 {
522 config = config.with_tool_bus(runtime.tool_bus.clone());
523 }
524 let event_sender = build_event_sender(self.callbacks.clone(), self.event_tx.clone());
525
526 let env_tx = env_state_tx.clone();
527 let mw = middleware_opt(&runtime);
528 let handle = tokio::spawn(async move {
529 session
530 .react_loop::<P, AgentEvent>(init_msg, config, env_tx, mw, Some(event_sender))
531 .await
532 });
533
534 Ok(SendStreamHandle {
535 runtime: runtime.into_acquired(),
536 handle,
537 event_rx,
538 stream_rx,
539 env_state_tx,
540 })
541 }
542}
543
544fn build_event_sender(
550 callbacks: Arc<CallbackRegistry>,
551 event_tx: broadcast::Sender<AgentEvent>,
552) -> EventSenderFn<AgentEvent> {
553 Box::new(move |event: AgentEvent| {
554 callbacks.dispatch(event.clone());
555 let _ = event_tx.send(event);
556 })
557}
558
559#[cfg(feature = "middleware")]
561fn middleware_opt<P: ChatProvider, S>(
562 runtime: &AgentRuntime<P, S>,
563) -> Option<Arc<MiddlewareChain<AgentEvent, ErrorsEnabled>>> {
564 Some(runtime.middleware_chain())
565}
566
567#[cfg(not(feature = "middleware"))]
568fn middleware_opt<P: ChatProvider, S>(
569 _runtime: &AgentRuntime<P, S>,
570) -> Option<
571 Arc<
572 funera_core::middleware::MiddlewareChain<
573 AgentEvent,
574 funera_core::middleware::ErrorsEnabled,
575 >,
576 >,
577> {
578 None
579}
580
581async fn relay_broadcast_to_mpsc(
583 mut event_rx: broadcast::Receiver<AgentEvent>,
584 relay_tx: mpsc::Sender<AgentEvent>,
585) {
586 while let Ok(event) = event_rx.recv().await {
587 let is_done = matches!(event, AgentEvent::Done);
588 if relay_tx.send(event).await.is_err() {
589 break;
590 }
591 if is_done {
592 break;
593 }
594 }
595}
596
597async fn aggregate_response(
599 event_rx: &mut broadcast::Receiver<AgentEvent>,
600 react_result: Result<(), anyhow::Error>,
601) -> Result<ChatResponse, OrchestrateError> {
602 react_result.map_err(OrchestrateError::Session)?;
603
604 let mut content = String::new();
605 let mut tool_calls = Vec::new();
606 let mut iterations = 0usize;
607 let mut finish_reason: Option<String> = None;
608
609 let mut pending_requests: Vec<(Arc<str>, String, serde_json::Value)> = Vec::new();
611
612 loop {
613 match event_rx.recv().await {
614 Ok(AgentEvent::Text(t)) => {
615 content = t;
616 }
617 Ok(AgentEvent::ToolCallRequest {
618 call_id,
619 name,
620 args,
621 ..
622 }) => {
623 pending_requests.push((call_id, name, args));
624 }
625 Ok(AgentEvent::ToolCallResult {
626 call_id,
627 name: _,
628 result,
629 }) => {
630 if let Some(pos) = pending_requests
631 .iter()
632 .position(|(id, _, _)| *id == call_id)
633 {
634 let (_, name, args) = pending_requests.remove(pos);
635 tool_calls.push(ToolCallInfo { name, args, result });
636 }
637 }
638 Ok(AgentEvent::TurnStart) => iterations += 1,
639 Ok(AgentEvent::TurnEnd { finish_reason: fr }) => finish_reason = fr,
640 Ok(AgentEvent::Done) => break,
641 Err(broadcast::error::RecvError::Closed) => break,
642 Err(broadcast::error::RecvError::Lagged(_)) => continue,
643 _ => {}
644 }
645 }
646
647 Ok(ChatResponse {
648 content,
649 tool_calls,
650 iterations,
651 finish_reason,
652 })
653}
654
655#[cfg(test)]
656mod tests {
657 use super::*;
658 use std::sync::atomic::{AtomicUsize, Ordering};
659
660 #[test]
663 fn builder_minimal_build_succeeds() {
664 let agent = AgentBuilder::new().build();
665 assert!(agent.system_prompt.is_none());
666 assert!(agent.callbacks.is_empty());
667 }
668
669 #[test]
670 fn builder_system_prompt() {
671 let agent = AgentBuilder::new()
672 .system_prompt("You are helpful.")
673 .build();
674 assert_eq!(agent.system_prompt, Some("You are helpful.".into()));
675 }
676
677 #[test]
678 fn builder_multiple_builds_independent() {
679 let a1 = AgentBuilder::new().system_prompt("P1").build();
680 let a2 = AgentBuilder::new().system_prompt("P2").build();
681 assert_eq!(a1.system_prompt, Some("P1".into()));
682 assert_eq!(a2.system_prompt, Some("P2".into()));
683 }
684
685 #[test]
688 fn builder_on_token_registers() {
689 let agent = AgentBuilder::new().on_token(|_| {}).build();
690 assert!(!agent.callbacks.is_empty());
691 }
692
693 #[test]
694 fn builder_on_tool_call_registers() {
695 let agent = AgentBuilder::new().on_tool_call(|_, _| {}).build();
696 assert!(!agent.callbacks.is_empty());
697 }
698
699 #[test]
700 fn builder_on_tool_result_registers() {
701 let agent = AgentBuilder::new().on_tool_result(|_, _| {}).build();
702 assert!(!agent.callbacks.is_empty());
703 }
704
705 #[test]
706 fn builder_on_turn_start_registers() {
707 let agent = AgentBuilder::new().on_turn_start(|| {}).build();
708 assert!(!agent.callbacks.is_empty());
709 }
710
711 #[test]
712 fn builder_on_turn_end_registers() {
713 let agent = AgentBuilder::new().on_turn_end(|| {}).build();
714 assert!(!agent.callbacks.is_empty());
715 }
716
717 #[test]
718 fn builder_on_event_registers() {
719 let agent = AgentBuilder::new().on_event(|_| {}).build();
720 assert!(!agent.callbacks.is_empty());
721 }
722
723 #[test]
724 fn builder_all_callbacks_stacked() {
725 let agent = AgentBuilder::new()
726 .on_token(|_| {})
727 .on_tool_call(|_, _| {})
728 .on_event(|_| {})
729 .build();
730 assert!(agent.callbacks.len() >= 3);
732 }
733
734 #[tokio::test]
737 async fn subscribe_events_receives_token() {
738 let agent = AgentBuilder::new().build();
739 let mut rx = agent.subscribe_events();
740 agent
741 .event_tx
742 .send(AgentEvent::Text("hello".into()))
743 .unwrap();
744 let got = tokio::time::timeout(std::time::Duration::from_secs(1), rx.recv()).await;
745 assert!(matches!(got, Ok(Ok(AgentEvent::Text(t))) if t == "hello"));
746 }
747
748 #[tokio::test]
749 async fn subscribe_events_receives_tool_call() {
750 let agent = AgentBuilder::new().build();
751 let mut rx = agent.subscribe_events();
752 agent
753 .event_tx
754 .send(AgentEvent::ToolCallRequest {
755 index: 0,
756 call_id: "call_abc".into(),
757 name: "test".into(),
758 args: serde_json::json!({}),
759 })
760 .unwrap();
761 let got = tokio::time::timeout(std::time::Duration::from_secs(1), rx.recv()).await;
762 assert!(matches!(got, Ok(Ok(AgentEvent::ToolCallRequest { .. }))));
763 }
764
765 #[tokio::test]
766 async fn subscribe_events_multiple_receivers() {
767 let agent = AgentBuilder::new().build();
768 let mut rx1 = agent.subscribe_events();
769 let mut rx2 = agent.subscribe_events();
770 agent.event_tx.send(AgentEvent::Done).unwrap();
771
772 let r1 = tokio::time::timeout(std::time::Duration::from_secs(1), rx1.recv()).await;
773 let r2 = tokio::time::timeout(std::time::Duration::from_secs(1), rx2.recv()).await;
774 assert!(r1.is_ok());
775 assert!(r2.is_ok());
776 }
777
778 #[tokio::test]
779 async fn subscribe_raw_events_receives_raw_token() {
780 use funera_core::event_bus::token_bus::TokenEvent;
781 let agent = AgentBuilder::new().build();
782 let mut rx = agent.subscribe_raw_events();
783 agent
784 .raw_event_tx
785 .send(RawAgentEvent::Token(TokenEvent::Text("raw".into())))
786 .unwrap();
787 let got = tokio::time::timeout(std::time::Duration::from_secs(1), rx.recv()).await;
788 assert!(matches!(
789 got,
790 Ok(Ok(RawAgentEvent::Token(TokenEvent::Text(t)))) if t == "raw"
791 ));
792 }
793
794 #[test]
797 fn callbacks_fire_on_dispatch() {
798 let counter = Arc::new(AtomicUsize::new(0));
799 let agent = AgentBuilder::new()
800 .on_event({
801 let c = counter.clone();
802 move |_| {
803 c.fetch_add(1, Ordering::SeqCst);
804 }
805 })
806 .build();
807 agent.callbacks.dispatch(AgentEvent::Done);
808 assert_eq!(counter.load(Ordering::SeqCst), 1);
809 }
810
811 #[test]
812 fn callbacks_only_fire_matching_event() {
813 let token_hits = Arc::new(AtomicUsize::new(0));
814 let tool_hits = Arc::new(AtomicUsize::new(0));
815
816 let agent = AgentBuilder::new()
817 .on_token({
818 let c = token_hits.clone();
819 move |_| {
820 c.fetch_add(1, Ordering::SeqCst);
821 }
822 })
823 .on_tool_call({
824 let c = tool_hits.clone();
825 move |_, _| {
826 c.fetch_add(1, Ordering::SeqCst);
827 }
828 })
829 .build();
830
831 agent.callbacks.dispatch(AgentEvent::Text("x".into()));
832 assert_eq!(token_hits.load(Ordering::SeqCst), 1);
833 assert_eq!(tool_hits.load(Ordering::SeqCst), 0);
834 }
835}