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 react = runtime.get_react_config().await;
284 let mut config = ReActLoopConfig::new(
285 react.channel_buffer,
286 react.max_iterations,
287 react.env_watcher,
288 env_state_tx.clone(),
289 turn_highway_handle,
290 );
291 #[cfg(feature = "tool")]
292 {
293 config = config.with_tool_bus(react.tool_bus);
294 }
295
296 let event_sender = build_event_sender(self.callbacks.clone(), self.event_tx.clone());
297
298 let result = session
299 .react_loop::<P, AgentEvent>(
300 init_msg,
301 config,
302 env_state_tx.clone(),
303 middleware_opt(runtime),
304 Some(event_sender),
305 )
306 .await;
307
308 let _ = env_state_tx.send(EnvStateEvent::SessionClosed);
309 aggregate_response(&mut event_rx, result).await
310 }
311
312 pub async fn fire_stream<P: ChatProvider, S>(
317 &self,
318 msg: impl Into<String>,
319 runtime: &AgentRuntime<P, S>,
320 ) -> Result<FireStreamHandle, OrchestrateError> {
321 let text = msg.into();
322 let event_rx = self.subscribe_events();
323
324 let (relay_tx, stream_rx) = mpsc::channel(256);
326 let relay_event_rx = self.subscribe_events();
327 tokio::spawn(async move {
328 relay_broadcast_to_mpsc(relay_event_rx, relay_tx).await;
329 });
330
331 let (env_state_bus, turn_highway_handle) = EnvStateBus::new();
332 let env_state_tx = env_state_bus.env_state_tx.clone();
333 let env_state_rx = env_state_bus.subscribe();
334 env_state_bus.start_turn_highway();
335
336 let _dispatcher = CallbackDispatcher::new(
337 env_state_rx,
338 self.event_tx.clone(),
339 self.raw_event_tx.clone(),
340 );
341 let _ = env_state_tx.send(EnvStateEvent::SessionStart);
342
343 let session_tx = funera_core::chat::session::spawn_session_actor();
344 let session = FuneraSession::new(session_tx);
345 if let Some(ref sys) = self.system_prompt {
346 session.push_message(FuneraMessage::new(
347 Role::System,
348 MsgVariant::Text(TextMessage {
349 text: sys.clone().into(),
350 reasoning_content: None,
351 }),
352 ));
353 }
354 let init_msg = FuneraMessage::new(
355 Role::User,
356 MsgVariant::Text(TextMessage {
357 text: text.into(),
358 reasoning_content: None,
359 }),
360 );
361 let react = runtime.get_react_config().await;
362 let mut config = ReActLoopConfig::new(
363 react.channel_buffer,
364 react.max_iterations,
365 react.env_watcher,
366 env_state_tx.clone(),
367 turn_highway_handle,
368 );
369 #[cfg(feature = "tool")]
370 {
371 config = config.with_tool_bus(react.tool_bus);
372 }
373 let event_sender = build_event_sender(self.callbacks.clone(), self.event_tx.clone());
374
375 let mw = middleware_opt(runtime);
377 let env_tx = env_state_tx.clone();
378 let handle = tokio::spawn(async move {
379 session
380 .react_loop::<P, AgentEvent>(init_msg, config, env_tx, mw, Some(event_sender))
381 .await
382 });
383
384 Ok(FireStreamHandle {
385 handle,
386 event_rx,
387 stream_rx,
388 env_state_tx,
389 })
390 }
391
392 pub async fn send<P: ChatProvider>(
399 &self,
400 msg: impl Into<String>,
401 runtime: AgentRuntime<P, Idle>,
402 ) -> Result<SendHandle<P>, OrchestrateError> {
403 let text = msg.into();
404 let event_rx = self.subscribe_events();
405
406 let (env_state_bus, turn_highway_handle) = EnvStateBus::new();
407 let env_state_tx = env_state_bus.env_state_tx.clone();
408 let env_state_rx = env_state_bus.subscribe();
409 env_state_bus.start_turn_highway();
410
411 let _dispatcher = CallbackDispatcher::new(
412 env_state_rx,
413 self.event_tx.clone(),
414 self.raw_event_tx.clone(),
415 );
416 let _ = env_state_tx.send(EnvStateEvent::SessionStart);
417
418 let session = FuneraSession::new(runtime.session_tx());
419 if let Some(ref sys) = self.system_prompt {
420 let msgs = session.session_context().await;
421 if msgs.is_empty() {
422 session.push_message(FuneraMessage::new(
423 Role::System,
424 MsgVariant::Text(TextMessage {
425 text: sys.clone().into(),
426 reasoning_content: None,
427 }),
428 ));
429 }
430 }
431 let init_msg = FuneraMessage::new(
432 Role::User,
433 MsgVariant::Text(TextMessage {
434 text: text.into(),
435 reasoning_content: None,
436 }),
437 );
438 let react = runtime.get_react_config().await;
439 let mut config = ReActLoopConfig::new(
440 react.channel_buffer,
441 react.max_iterations,
442 react.env_watcher,
443 env_state_tx.clone(),
444 turn_highway_handle,
445 );
446 #[cfg(feature = "tool")]
447 {
448 config = config.with_tool_bus(react.tool_bus);
449 }
450 let event_sender = build_event_sender(self.callbacks.clone(), self.event_tx.clone());
451
452 let env_tx = env_state_tx.clone();
453 let mw = middleware_opt(&runtime);
454 let handle = tokio::spawn(async move {
455 session
456 .react_loop::<P, AgentEvent>(init_msg, config, env_tx, mw, Some(event_sender))
457 .await
458 });
459
460 Ok(SendHandle {
461 runtime: runtime.into_acquired(),
462 handle,
463 event_rx,
464 env_state_tx,
465 })
466 }
467
468 pub async fn send_stream<P: ChatProvider>(
470 &self,
471 msg: impl Into<String>,
472 runtime: AgentRuntime<P, Idle>,
473 ) -> Result<SendStreamHandle<P>, OrchestrateError> {
474 let text = msg.into();
475 let event_rx = self.subscribe_events();
476
477 let (relay_tx, stream_rx) = mpsc::channel(256);
479 let relay_event_rx = self.subscribe_events();
480 tokio::spawn(async move {
481 relay_broadcast_to_mpsc(relay_event_rx, relay_tx).await;
482 });
483
484 let (env_state_bus, turn_highway_handle) = EnvStateBus::new();
485 let env_state_tx = env_state_bus.env_state_tx.clone();
486 let env_state_rx = env_state_bus.subscribe();
487 env_state_bus.start_turn_highway();
488
489 let _dispatcher = CallbackDispatcher::new(
490 env_state_rx,
491 self.event_tx.clone(),
492 self.raw_event_tx.clone(),
493 );
494 let _ = env_state_tx.send(EnvStateEvent::SessionStart);
495
496 let session = FuneraSession::new(runtime.session_tx());
497 if let Some(ref sys) = self.system_prompt {
498 let msgs = session.session_context().await;
499 if msgs.is_empty() {
500 session.push_message(FuneraMessage::new(
501 Role::System,
502 MsgVariant::Text(TextMessage {
503 text: sys.clone().into(),
504 reasoning_content: None,
505 }),
506 ));
507 }
508 }
509 let init_msg = FuneraMessage::new(
510 Role::User,
511 MsgVariant::Text(TextMessage {
512 text: text.into(),
513 reasoning_content: None,
514 }),
515 );
516 let react = runtime.get_react_config().await;
517 let mut config = ReActLoopConfig::new(
518 react.channel_buffer,
519 react.max_iterations,
520 react.env_watcher,
521 env_state_tx.clone(),
522 turn_highway_handle,
523 );
524 #[cfg(feature = "tool")]
525 {
526 config = config.with_tool_bus(react.tool_bus);
527 }
528 let event_sender = build_event_sender(self.callbacks.clone(), self.event_tx.clone());
529
530 let env_tx = env_state_tx.clone();
531 let mw = middleware_opt(&runtime);
532 let handle = tokio::spawn(async move {
533 session
534 .react_loop::<P, AgentEvent>(init_msg, config, env_tx, mw, Some(event_sender))
535 .await
536 });
537
538 Ok(SendStreamHandle {
539 runtime: runtime.into_acquired(),
540 handle,
541 event_rx,
542 stream_rx,
543 env_state_tx,
544 })
545 }
546}
547
548fn build_event_sender(
554 callbacks: Arc<CallbackRegistry>,
555 event_tx: broadcast::Sender<AgentEvent>,
556) -> EventSenderFn<AgentEvent> {
557 Box::new(move |event: AgentEvent| {
558 callbacks.dispatch(event.clone());
559 let _ = event_tx.send(event);
560 })
561}
562
563#[cfg(feature = "middleware")]
565fn middleware_opt<P: ChatProvider, S>(
566 runtime: &AgentRuntime<P, S>,
567) -> Option<Arc<MiddlewareChain<AgentEvent, ErrorsEnabled>>> {
568 Some(runtime.middleware_chain())
569}
570
571#[cfg(not(feature = "middleware"))]
572fn middleware_opt<P: ChatProvider, S>(
573 _runtime: &AgentRuntime<P, S>,
574) -> Option<
575 Arc<
576 funera_core::middleware::MiddlewareChain<
577 AgentEvent,
578 funera_core::middleware::ErrorsEnabled,
579 >,
580 >,
581> {
582 None
583}
584
585async fn relay_broadcast_to_mpsc(
587 mut event_rx: broadcast::Receiver<AgentEvent>,
588 relay_tx: mpsc::Sender<AgentEvent>,
589) {
590 while let Ok(event) = event_rx.recv().await {
591 let is_done = matches!(event, AgentEvent::Done);
592 if relay_tx.send(event).await.is_err() {
593 break;
594 }
595 if is_done {
596 break;
597 }
598 }
599}
600
601async fn aggregate_response(
603 event_rx: &mut broadcast::Receiver<AgentEvent>,
604 react_result: Result<(), anyhow::Error>,
605) -> Result<ChatResponse, OrchestrateError> {
606 react_result.map_err(OrchestrateError::Session)?;
607
608 let mut content = String::new();
609 let mut tool_calls = Vec::new();
610 let mut iterations = 0usize;
611 let mut finish_reason: Option<String> = None;
612
613 let mut pending_requests: Vec<(Arc<str>, String, serde_json::Value)> = Vec::new();
615
616 loop {
617 match event_rx.recv().await {
618 Ok(AgentEvent::Text(t)) => {
619 content = t;
620 }
621 Ok(AgentEvent::ToolCallRequest {
622 call_id,
623 name,
624 args,
625 ..
626 }) => {
627 pending_requests.push((call_id, name, args));
628 }
629 Ok(AgentEvent::ToolCallResult {
630 call_id,
631 name: _,
632 result,
633 }) => {
634 if let Some(pos) = pending_requests
635 .iter()
636 .position(|(id, _, _)| *id == call_id)
637 {
638 let (_, name, args) = pending_requests.remove(pos);
639 tool_calls.push(ToolCallInfo { name, args, result });
640 }
641 }
642 Ok(AgentEvent::TurnStart) => iterations += 1,
643 Ok(AgentEvent::TurnEnd { finish_reason: fr }) => finish_reason = fr,
644 Ok(AgentEvent::Done) => break,
645 Err(broadcast::error::RecvError::Closed) => break,
646 Err(broadcast::error::RecvError::Lagged(_)) => continue,
647 _ => {}
648 }
649 }
650
651 Ok(ChatResponse {
652 content,
653 tool_calls,
654 iterations,
655 finish_reason,
656 })
657}
658
659#[cfg(test)]
660mod tests {
661 use super::*;
662 use std::sync::atomic::{AtomicUsize, Ordering};
663
664 #[test]
667 fn builder_minimal_build_succeeds() {
668 let agent = AgentBuilder::new().build();
669 assert!(agent.system_prompt.is_none());
670 assert!(agent.callbacks.is_empty());
671 }
672
673 #[test]
674 fn builder_system_prompt() {
675 let agent = AgentBuilder::new()
676 .system_prompt("You are helpful.")
677 .build();
678 assert_eq!(agent.system_prompt, Some("You are helpful.".into()));
679 }
680
681 #[test]
682 fn builder_multiple_builds_independent() {
683 let a1 = AgentBuilder::new().system_prompt("P1").build();
684 let a2 = AgentBuilder::new().system_prompt("P2").build();
685 assert_eq!(a1.system_prompt, Some("P1".into()));
686 assert_eq!(a2.system_prompt, Some("P2".into()));
687 }
688
689 #[test]
692 fn builder_on_token_registers() {
693 let agent = AgentBuilder::new().on_token(|_| {}).build();
694 assert!(!agent.callbacks.is_empty());
695 }
696
697 #[test]
698 fn builder_on_tool_call_registers() {
699 let agent = AgentBuilder::new().on_tool_call(|_, _| {}).build();
700 assert!(!agent.callbacks.is_empty());
701 }
702
703 #[test]
704 fn builder_on_tool_result_registers() {
705 let agent = AgentBuilder::new().on_tool_result(|_, _| {}).build();
706 assert!(!agent.callbacks.is_empty());
707 }
708
709 #[test]
710 fn builder_on_turn_start_registers() {
711 let agent = AgentBuilder::new().on_turn_start(|| {}).build();
712 assert!(!agent.callbacks.is_empty());
713 }
714
715 #[test]
716 fn builder_on_turn_end_registers() {
717 let agent = AgentBuilder::new().on_turn_end(|| {}).build();
718 assert!(!agent.callbacks.is_empty());
719 }
720
721 #[test]
722 fn builder_on_event_registers() {
723 let agent = AgentBuilder::new().on_event(|_| {}).build();
724 assert!(!agent.callbacks.is_empty());
725 }
726
727 #[test]
728 fn builder_all_callbacks_stacked() {
729 let agent = AgentBuilder::new()
730 .on_token(|_| {})
731 .on_tool_call(|_, _| {})
732 .on_event(|_| {})
733 .build();
734 assert!(agent.callbacks.len() >= 3);
736 }
737
738 #[tokio::test]
741 async fn subscribe_events_receives_token() {
742 let agent = AgentBuilder::new().build();
743 let mut rx = agent.subscribe_events();
744 agent
745 .event_tx
746 .send(AgentEvent::Text("hello".into()))
747 .unwrap();
748 let got = tokio::time::timeout(std::time::Duration::from_secs(1), rx.recv()).await;
749 assert!(matches!(got, Ok(Ok(AgentEvent::Text(t))) if t == "hello"));
750 }
751
752 #[tokio::test]
753 async fn subscribe_events_receives_tool_call() {
754 let agent = AgentBuilder::new().build();
755 let mut rx = agent.subscribe_events();
756 agent
757 .event_tx
758 .send(AgentEvent::ToolCallRequest {
759 index: 0,
760 call_id: "call_abc".into(),
761 name: "test".into(),
762 args: serde_json::json!({}),
763 })
764 .unwrap();
765 let got = tokio::time::timeout(std::time::Duration::from_secs(1), rx.recv()).await;
766 assert!(matches!(got, Ok(Ok(AgentEvent::ToolCallRequest { .. }))));
767 }
768
769 #[tokio::test]
770 async fn subscribe_events_multiple_receivers() {
771 let agent = AgentBuilder::new().build();
772 let mut rx1 = agent.subscribe_events();
773 let mut rx2 = agent.subscribe_events();
774 agent.event_tx.send(AgentEvent::Done).unwrap();
775
776 let r1 = tokio::time::timeout(std::time::Duration::from_secs(1), rx1.recv()).await;
777 let r2 = tokio::time::timeout(std::time::Duration::from_secs(1), rx2.recv()).await;
778 assert!(r1.is_ok());
779 assert!(r2.is_ok());
780 }
781
782 #[tokio::test]
783 async fn subscribe_raw_events_receives_raw_token() {
784 use funera_core::event_bus::token_bus::TokenEvent;
785 let agent = AgentBuilder::new().build();
786 let mut rx = agent.subscribe_raw_events();
787 agent
788 .raw_event_tx
789 .send(RawAgentEvent::Token(TokenEvent::Text("raw".into())))
790 .unwrap();
791 let got = tokio::time::timeout(std::time::Duration::from_secs(1), rx.recv()).await;
792 assert!(matches!(
793 got,
794 Ok(Ok(RawAgentEvent::Token(TokenEvent::Text(t)))) if t == "raw"
795 ));
796 }
797
798 #[test]
801 fn callbacks_fire_on_dispatch() {
802 let counter = Arc::new(AtomicUsize::new(0));
803 let agent = AgentBuilder::new()
804 .on_event({
805 let c = counter.clone();
806 move |_| {
807 c.fetch_add(1, Ordering::SeqCst);
808 }
809 })
810 .build();
811 agent.callbacks.dispatch(AgentEvent::Done);
812 assert_eq!(counter.load(Ordering::SeqCst), 1);
813 }
814
815 #[test]
816 fn callbacks_only_fire_matching_event() {
817 let token_hits = Arc::new(AtomicUsize::new(0));
818 let tool_hits = Arc::new(AtomicUsize::new(0));
819
820 let agent = AgentBuilder::new()
821 .on_token({
822 let c = token_hits.clone();
823 move |_| {
824 c.fetch_add(1, Ordering::SeqCst);
825 }
826 })
827 .on_tool_call({
828 let c = tool_hits.clone();
829 move |_, _| {
830 c.fetch_add(1, Ordering::SeqCst);
831 }
832 })
833 .build();
834
835 agent.callbacks.dispatch(AgentEvent::Text("x".into()));
836 assert_eq!(token_hits.load(Ordering::SeqCst), 1);
837 assert_eq!(tool_hits.load(Ordering::SeqCst), 0);
838 }
839}