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