1use std::sync::Arc;
12use std::time::Duration;
13
14use async_trait::async_trait;
15use serde_json::json;
16use tokio::sync::broadcast;
17
18use gemini_genai_rs::prelude::{FunctionResponse, Tool, recv_event};
19use gemini_genai_rs::session::SessionEvent;
20
21use crate::agent::Agent;
22use crate::context::{AgentEvent, InvocationContext};
23use crate::error::{AgentError, ToolError};
24use crate::middleware::MiddlewareChain;
25use crate::plugin::{PluginManager, PluginResult};
26use crate::tool::{
27 ActiveStreamingTool, InputStreamingTool, SimpleTool, StreamingTool, ToolClass, ToolDispatcher,
28 ToolFunction, ToolKind, TypedTool,
29};
30
31pub struct LlmAgent {
37 name: String,
38 dispatcher: ToolDispatcher,
39 middleware: MiddlewareChain,
40 plugins: PluginManager,
41 sub_agents: Vec<Arc<dyn Agent>>,
42}
43
44impl LlmAgent {
45 pub fn builder(name: impl Into<String>) -> LlmAgentBuilder {
47 LlmAgentBuilder {
48 name: name.into(),
49 dispatcher: ToolDispatcher::new(),
50 middleware: MiddlewareChain::new(),
51 plugins: PluginManager::new(),
52 sub_agents: Vec::new(),
53 }
54 }
55
56 pub fn dispatcher(&self) -> &ToolDispatcher {
58 &self.dispatcher
59 }
60
61 pub fn middleware(&self) -> &MiddlewareChain {
63 &self.middleware
64 }
65
66 pub fn plugins(&self) -> &PluginManager {
68 &self.plugins
69 }
70
71 async fn event_loop(
73 &self,
74 ctx: &mut InvocationContext,
75 events: &mut broadcast::Receiver<SessionEvent>,
76 agent_name: &str,
77 ) -> Result<(), AgentError> {
78 loop {
79 let event = match recv_event(events).await {
80 Some(e) => e,
81 None => break, };
83
84 match event {
85 SessionEvent::ToolCall(calls) => {
86 let mut responses = Vec::new();
87 let mut transfer_target = None;
88
89 for call in &calls {
90 ctx.emit(AgentEvent::ToolCallStarted {
92 name: call.name.clone(),
93 args: call.args.clone(),
94 });
95 let _ = ctx.middleware.run_before_tool(call).await;
96
97 let plugin_result = self.plugins.run_before_tool(call, ctx).await;
99 match &plugin_result {
100 PluginResult::Deny(reason) => {
101 ctx.emit(AgentEvent::ToolCallFailed {
102 name: call.name.clone(),
103 error: format!("Denied by plugin: {reason}"),
104 });
105 responses.push(ToolDispatcher::build_response(
106 call,
107 Err(ToolError::ExecutionFailed(format!(
108 "Denied by plugin: {reason}"
109 ))),
110 ));
111 continue;
112 }
113 PluginResult::ShortCircuit(value) => {
114 let _ = ctx.middleware.run_after_tool(call, value).await;
115 ctx.emit(AgentEvent::ToolCallCompleted {
116 name: call.name.clone(),
117 result: value.clone(),
118 duration: std::time::Duration::ZERO,
119 });
120 responses
121 .push(ToolDispatcher::build_response(call, Ok(value.clone())));
122 continue;
123 }
124 PluginResult::Continue => {}
125 }
126
127 let tool_start = std::time::Instant::now();
128 let tool_class = self.dispatcher.classify(&call.name);
129
130 match tool_class {
131 Some(ToolClass::Regular) => {
132 crate::telemetry::logging::log_tool_dispatch(
133 agent_name, &call.name, "function",
134 );
135 crate::telemetry::metrics::record_agent_tool_dispatched(
136 agent_name, &call.name,
137 );
138
139 let result = self
140 .dispatcher
141 .call_function(&call.name, call.args.clone())
142 .await;
143 let elapsed = tool_start.elapsed();
144
145 match &result {
146 Ok(value) => {
147 if let Some(target) =
149 value.get("__transfer_to").and_then(|v| v.as_str())
150 {
151 transfer_target = Some(target.to_string());
152 }
153
154 let _ = ctx.middleware.run_after_tool(call, value).await;
155 let _ = self.plugins.run_after_tool(call, value, ctx).await;
156 ctx.emit(AgentEvent::ToolCallCompleted {
157 name: call.name.clone(),
158 result: value.clone(),
159 duration: elapsed,
160 });
161 crate::telemetry::logging::log_tool_result(
162 agent_name,
163 &call.name,
164 true,
165 elapsed.as_millis() as f64,
166 );
167 crate::telemetry::metrics::record_agent_tool_duration(
168 agent_name,
169 &call.name,
170 elapsed.as_millis() as f64,
171 );
172 }
173 Err(e) => {
174 let _ = ctx.middleware.run_on_tool_error(call, e).await;
175 ctx.emit(AgentEvent::ToolCallFailed {
176 name: call.name.clone(),
177 error: e.to_string(),
178 });
179 crate::telemetry::logging::log_tool_result(
180 agent_name,
181 &call.name,
182 false,
183 elapsed.as_millis() as f64,
184 );
185 }
186 }
187
188 responses.push(ToolDispatcher::build_response(call, result));
189 }
190 Some(ToolClass::Streaming) | Some(ToolClass::InputStream) => {
191 let class_str = if tool_class == Some(ToolClass::Streaming) {
192 "streaming"
193 } else {
194 "input_stream"
195 };
196 crate::telemetry::logging::log_tool_dispatch(
197 agent_name, &call.name, class_str,
198 );
199
200 self.spawn_streaming_tool(call, ctx, agent_name).await;
201
202 responses.push(FunctionResponse {
203 name: call.name.clone(),
204 response: json!({"status": "streaming"}),
205 id: call.id.clone(),
206 scheduling: None,
207 });
208 }
209 None => {
210 ctx.emit(AgentEvent::ToolCallFailed {
211 name: call.name.clone(),
212 error: format!("Tool not found: {}", call.name),
213 });
214 responses.push(ToolDispatcher::build_response(
215 call,
216 Err(ToolError::NotFound(call.name.clone())),
217 ));
218 }
219 }
220 }
221
222 ctx.agent_session.send_tool_response(responses).await?;
224
225 if let Some(target) = transfer_target {
227 ctx.emit(AgentEvent::AgentTransfer {
228 from: agent_name.to_string(),
229 to: target.clone(),
230 });
231 crate::telemetry::metrics::record_agent_transfer(agent_name, &target);
232 crate::telemetry::logging::log_agent_transfer(agent_name, &target);
233 return Err(AgentError::TransferRequested(target));
234 }
235 }
236 SessionEvent::ToolCallCancelled(ids) => {
237 self.dispatcher.cancel_by_ids(&ids).await;
238 }
239 SessionEvent::TurnComplete => {
240 ctx.emit(AgentEvent::Session(SessionEvent::TurnComplete));
241 break;
242 }
243 SessionEvent::Disconnected(reason) => {
244 ctx.emit(AgentEvent::Session(SessionEvent::Disconnected(reason)));
245 break;
246 }
247 SessionEvent::Error(ref e) => {
248 ctx.emit(AgentEvent::Session(event.clone()));
249 crate::telemetry::metrics::record_agent_error(agent_name, "session_error");
250 crate::telemetry::logging::log_agent_error(agent_name, &e.to_string());
251 }
252 other => {
253 ctx.emit(AgentEvent::Session(other));
255 }
256 }
257 }
258 Ok(())
259 }
260
261 async fn spawn_streaming_tool(
263 &self,
264 call: &gemini_genai_rs::prelude::FunctionCall,
265 ctx: &InvocationContext,
266 _agent_name: &str,
267 ) {
268 let tool_kind = match self.dispatcher.get_tool(&call.name) {
269 Some(kind) => kind,
270 None => return,
271 };
272
273 let (yield_tx, mut yield_rx) = tokio::sync::mpsc::channel::<serde_json::Value>(32);
274 let cancel = tokio_util::sync::CancellationToken::new();
275
276 let tool_name = call.name.clone();
277 let call_id = call.id.clone();
278 let args = call.args.clone();
279 let event_tx = ctx.event_tx.clone();
280 let agent_session = ctx.agent_session.clone();
281
282 match tool_kind {
283 ToolKind::Streaming(tool) => {
284 let tool = tool.clone();
285 let cancel_clone = cancel.clone();
286 let tool_name_err = tool_name.clone();
287 let event_tx_err = event_tx.clone();
288
289 let tool_task = tokio::spawn(async move {
290 tokio::select! {
291 result = tool.run(args, yield_tx) => {
292 if let Err(e) = result {
293 let _ = event_tx_err.send(AgentEvent::ToolCallFailed {
294 name: tool_name_err,
295 error: e.to_string(),
296 });
297 }
298 }
299 _ = cancel_clone.cancelled() => {}
300 }
301 });
302
303 let active = ActiveStreamingTool {
304 task: tool_task,
305 cancel,
306 };
307 let id = call_id.clone().unwrap_or_else(|| tool_name.clone());
308 self.dispatcher.store_active(id, active).await;
309 }
310 ToolKind::InputStream(tool) => {
311 let tool = tool.clone();
312 let input_rx = ctx.agent_session.subscribe_input();
313 let cancel_clone = cancel.clone();
314 let tool_name_err = tool_name.clone();
315 let event_tx_err = event_tx.clone();
316
317 let tool_task = tokio::spawn(async move {
318 tokio::select! {
319 result = tool.run(args, input_rx, yield_tx) => {
320 if let Err(e) = result {
321 let _ = event_tx_err.send(AgentEvent::ToolCallFailed {
322 name: tool_name_err,
323 error: e.to_string(),
324 });
325 }
326 }
327 _ = cancel_clone.cancelled() => {}
328 }
329 });
330
331 let active = ActiveStreamingTool {
332 task: tool_task,
333 cancel,
334 };
335 let id = call_id.clone().unwrap_or_else(|| tool_name.clone());
336 self.dispatcher.store_active(id, active).await;
337 }
338 ToolKind::Function(_) => {} }
340
341 let yield_tool_name = call.name.clone();
343 let yield_call_id = call.id.clone();
344
345 tokio::spawn(async move {
346 let mut all_yields = Vec::new();
347 while let Some(value) = yield_rx.recv().await {
348 let _ = event_tx.send(AgentEvent::StreamingToolYield {
349 name: yield_tool_name.clone(),
350 value: value.clone(),
351 });
352 all_yields.push(value);
353 }
354
355 let final_response = if all_yields.is_empty() {
357 json!({"status": "completed"})
358 } else if all_yields.len() == 1 {
359 all_yields.into_iter().next().unwrap()
360 } else {
361 json!({"results": all_yields})
362 };
363
364 let resp = FunctionResponse {
365 name: yield_tool_name,
366 response: final_response,
367 id: yield_call_id,
368 scheduling: None,
369 };
370 let _ = agent_session.send_tool_response(vec![resp]).await;
371 });
372 }
373}
374
375pub struct LlmAgentBuilder {
377 name: String,
378 dispatcher: ToolDispatcher,
379 middleware: MiddlewareChain,
380 plugins: PluginManager,
381 sub_agents: Vec<Arc<dyn Agent>>,
382}
383
384impl LlmAgentBuilder {
385 pub fn tool(mut self, tool: impl ToolFunction + 'static) -> Self {
387 self.dispatcher.register_function(Arc::new(tool));
388 self
389 }
390
391 pub fn typed_tool<T>(mut self, tool: TypedTool<T>) -> Self
393 where
394 T: serde::de::DeserializeOwned + schemars::JsonSchema + Send + Sync + 'static,
395 {
396 self.dispatcher.register_function(Arc::new(tool));
397 self
398 }
399
400 pub fn streaming_tool(mut self, tool: impl StreamingTool + 'static) -> Self {
402 self.dispatcher.register_streaming(Arc::new(tool));
403 self
404 }
405
406 pub fn input_streaming_tool(mut self, tool: impl InputStreamingTool + 'static) -> Self {
408 self.dispatcher.register_input_streaming(Arc::new(tool));
409 self
410 }
411
412 pub fn middleware(mut self, mw: impl crate::middleware::Middleware + 'static) -> Self {
414 self.middleware.add(Arc::new(mw));
415 self
416 }
417
418 pub fn plugin(mut self, plugin: impl crate::plugin::Plugin + 'static) -> Self {
420 self.plugins.add(Arc::new(plugin));
421 self
422 }
423
424 pub fn sub_agent(mut self, agent: impl Agent + 'static) -> Self {
426 self.sub_agents.push(Arc::new(agent));
427 self
428 }
429
430 pub fn tool_timeout(mut self, timeout: Duration) -> Self {
432 self.dispatcher = self.dispatcher.with_timeout(timeout);
433 self
434 }
435
436 pub fn build(mut self) -> LlmAgent {
443 for sub in &self.sub_agents {
445 let target_name = sub.name().to_string();
446 let tool_name = format!("transfer_to_{target_name}");
447 let transfer_tool = SimpleTool::new(
448 tool_name,
449 format!("Transfer conversation to the {target_name} agent"),
450 Some(json!({
451 "type": "object",
452 "properties": {},
453 })),
454 move |_args| {
455 let name = target_name.clone();
456 async move { Ok(json!({"__transfer_to": name})) }
457 },
458 );
459 self.dispatcher.register_function(Arc::new(transfer_tool));
460 }
461
462 self.middleware
464 .prepend(Arc::new(crate::telemetry::TelemetryMiddleware::new(
465 &self.name,
466 )));
467
468 LlmAgent {
469 name: self.name,
470 dispatcher: self.dispatcher,
471 middleware: self.middleware,
472 plugins: self.plugins,
473 sub_agents: self.sub_agents,
474 }
475 }
476}
477
478#[async_trait]
479impl Agent for LlmAgent {
480 fn name(&self) -> &str {
481 &self.name
482 }
483
484 async fn run_live(&self, ctx: &mut InvocationContext) -> Result<(), AgentError> {
485 let agent_name = self.name.clone();
486 let start = std::time::Instant::now();
487
488 crate::telemetry::logging::log_agent_started(&agent_name, self.dispatcher.len());
490 crate::telemetry::metrics::record_agent_started(&agent_name);
491 ctx.middleware.run_before_agent(ctx).await?;
492
493 let plugin_result = self.plugins.run_before_agent(ctx).await;
495 if let PluginResult::Deny(reason) = plugin_result {
496 return Err(AgentError::Other(format!(
497 "Agent denied by plugin: {reason}"
498 )));
499 }
500
501 ctx.emit(AgentEvent::AgentStarted {
502 name: agent_name.clone(),
503 });
504
505 let mut events = ctx.agent_session.subscribe_events();
506
507 let result = self.event_loop(ctx, &mut events, &agent_name).await;
508
509 let elapsed = start.elapsed();
511 ctx.middleware.run_after_agent(ctx).await?;
512 let _ = self.plugins.run_after_agent(ctx).await;
513 ctx.emit(AgentEvent::AgentCompleted {
514 name: agent_name.clone(),
515 });
516 crate::telemetry::logging::log_agent_completed(&agent_name, elapsed.as_millis() as f64);
517 crate::telemetry::metrics::record_agent_completed(&agent_name, elapsed.as_millis() as f64);
518
519 result
520 }
521
522 fn tools(&self) -> Vec<Tool> {
523 self.dispatcher.to_tool_declarations()
524 }
525
526 fn sub_agents(&self) -> Vec<Arc<dyn Agent>> {
527 self.sub_agents.clone()
528 }
529}
530
531#[cfg(test)]
532mod tests {
533 use super::*;
534 use gemini_genai_rs::prelude::FunctionCall;
535 use gemini_genai_rs::session::{SessionError, SessionWriter};
536 use serde_json::json;
537
538 struct NoopAgent {
539 name: String,
540 }
541
542 #[async_trait]
543 impl Agent for NoopAgent {
544 fn name(&self) -> &str {
545 &self.name
546 }
547 async fn run_live(&self, _ctx: &mut InvocationContext) -> Result<(), AgentError> {
548 Ok(())
549 }
550 }
551
552 struct MockWriter;
554
555 #[async_trait]
556 impl SessionWriter for MockWriter {
557 async fn send_audio(&self, _data: bytes::Bytes) -> Result<(), SessionError> {
558 Ok(())
559 }
560 async fn send_text(&self, _text: String) -> Result<(), SessionError> {
561 Ok(())
562 }
563 async fn send_tool_response(
564 &self,
565 _responses: Vec<FunctionResponse>,
566 ) -> Result<(), SessionError> {
567 Ok(())
568 }
569 async fn send_client_content(
570 &self,
571 _turns: Vec<gemini_genai_rs::prelude::Content>,
572 _turn_complete: bool,
573 ) -> Result<(), SessionError> {
574 Ok(())
575 }
576 async fn send_video(&self, _jpeg_data: bytes::Bytes) -> Result<(), SessionError> {
577 Ok(())
578 }
579 async fn update_instruction(&self, _instruction: String) -> Result<(), SessionError> {
580 Ok(())
581 }
582 async fn signal_activity_start(&self) -> Result<(), SessionError> {
583 Ok(())
584 }
585 async fn signal_activity_end(&self) -> Result<(), SessionError> {
586 Ok(())
587 }
588 async fn disconnect(&self) -> Result<(), SessionError> {
589 Ok(())
590 }
591 }
592
593 fn mock_agent_session() -> (
596 crate::agent_session::AgentSession,
597 broadcast::Sender<SessionEvent>,
598 ) {
599 let (evt_tx, _) = broadcast::channel(64);
600 let writer: Arc<dyn SessionWriter> = Arc::new(MockWriter);
601 let session = crate::agent_session::AgentSession::from_writer(writer, evt_tx.clone());
602 (session, evt_tx)
603 }
604
605 #[test]
606 fn builder_creates_agent_with_name() {
607 let agent = LlmAgent::builder("test_agent").build();
608 assert_eq!(agent.name(), "test_agent");
609 }
610
611 #[test]
612 fn builder_registers_tools() {
613 let tool = SimpleTool::new("my_tool", "desc", None, |_| async { Ok(json!({})) });
614 let agent = LlmAgent::builder("test").tool(tool).build();
615 assert_eq!(agent.dispatcher().len(), 1);
617 }
618
619 #[test]
620 fn builder_auto_registers_transfer_tools() {
621 let sub = NoopAgent {
622 name: "billing".to_string(),
623 };
624 let agent = LlmAgent::builder("root").sub_agent(sub).build();
625
626 assert!(agent.dispatcher().classify("transfer_to_billing").is_some());
628 }
629
630 #[test]
631 fn builder_with_multiple_sub_agents() {
632 let sub1 = NoopAgent {
633 name: "billing".to_string(),
634 };
635 let sub2 = NoopAgent {
636 name: "tech".to_string(),
637 };
638 let agent = LlmAgent::builder("root")
639 .sub_agent(sub1)
640 .sub_agent(sub2)
641 .build();
642
643 assert!(agent.dispatcher().classify("transfer_to_billing").is_some());
644 assert!(agent.dispatcher().classify("transfer_to_tech").is_some());
645 assert_eq!(agent.sub_agents().len(), 2);
646 }
647
648 #[test]
649 fn tools_returns_declarations() {
650 let tool = SimpleTool::new("my_tool", "desc", None, |_| async { Ok(json!({})) });
651 let agent = LlmAgent::builder("test").tool(tool).build();
652 let tools = agent.tools();
653 assert!(!tools.is_empty());
654 }
655
656 #[test]
657 fn transfer_requested_error() {
658 let err = AgentError::TransferRequested("billing".to_string());
659 assert!(err.to_string().contains("billing"));
660 }
661
662 #[test]
663 fn builder_prepends_telemetry_middleware() {
664 let agent = LlmAgent::builder("test").build();
665 assert_eq!(agent.middleware().len(), 1);
667 }
668
669 #[test]
670 fn builder_with_user_middleware_and_telemetry() {
671 use crate::middleware::LogMiddleware;
672
673 let agent = LlmAgent::builder("test")
674 .middleware(LogMiddleware::new())
675 .build();
676 assert_eq!(agent.middleware().len(), 2);
678 }
679
680 #[test]
681 fn get_tool_returns_tool_kind() {
682 let tool = SimpleTool::new("lookup", "desc", None, |_| async { Ok(json!({})) });
683 let agent = LlmAgent::builder("test").tool(tool).build();
684 assert!(agent.dispatcher().get_tool("lookup").is_some());
685 assert!(agent.dispatcher().get_tool("nonexistent").is_none());
686 }
687
688 #[tokio::test]
691 async fn event_loop_breaks_on_turn_complete() {
692 let agent = LlmAgent::builder("test").build();
693 let (session, evt_tx) = mock_agent_session();
694 let mut ctx = InvocationContext::new(session);
695
696 tokio::spawn(async move {
698 tokio::time::sleep(std::time::Duration::from_millis(10)).await;
699 let _ = evt_tx.send(SessionEvent::TurnComplete);
700 });
701
702 let result = agent.run_live(&mut ctx).await;
703 assert!(result.is_ok());
704 }
705
706 #[tokio::test]
707 async fn event_loop_breaks_on_disconnect() {
708 let agent = LlmAgent::builder("test").build();
709 let (session, evt_tx) = mock_agent_session();
710 let mut ctx = InvocationContext::new(session);
711
712 tokio::spawn(async move {
713 tokio::time::sleep(std::time::Duration::from_millis(10)).await;
714 let _ = evt_tx.send(SessionEvent::Disconnected(Some("bye".to_string())));
715 });
716
717 let result = agent.run_live(&mut ctx).await;
718 assert!(result.is_ok());
719 }
720
721 #[tokio::test]
722 async fn event_loop_dispatches_tool_call() {
723 let tool = SimpleTool::new("get_weather", "Get weather", None, |_| async {
724 Ok(json!({"temp": 22}))
725 });
726 let agent = LlmAgent::builder("test").tool(tool).build();
727 let (session, evt_tx) = mock_agent_session();
728 let mut ctx = InvocationContext::new(session);
729 let mut agent_events = ctx.subscribe();
730
731 tokio::spawn(async move {
732 tokio::time::sleep(std::time::Duration::from_millis(10)).await;
733 let _ = evt_tx.send(SessionEvent::ToolCall(vec![FunctionCall {
734 name: "get_weather".to_string(),
735 args: json!({"city": "London"}),
736 id: Some("call-1".to_string()),
737 }]));
738 tokio::time::sleep(std::time::Duration::from_millis(50)).await;
740 let _ = evt_tx.send(SessionEvent::TurnComplete);
741 });
742
743 let result = agent.run_live(&mut ctx).await;
744 assert!(result.is_ok());
745
746 let mut saw_tool_started = false;
748 let mut saw_tool_completed = false;
749 while let Ok(event) = agent_events.try_recv() {
750 match event {
751 AgentEvent::ToolCallStarted { name, .. } if name == "get_weather" => {
752 saw_tool_started = true;
753 }
754 AgentEvent::ToolCallCompleted { name, result, .. } if name == "get_weather" => {
755 assert_eq!(result["temp"], 22);
756 saw_tool_completed = true;
757 }
758 _ => {}
759 }
760 }
761 assert!(saw_tool_started, "should have emitted ToolCallStarted");
762 assert!(saw_tool_completed, "should have emitted ToolCallCompleted");
763 }
764
765 #[tokio::test]
766 async fn event_loop_handles_unknown_tool() {
767 let agent = LlmAgent::builder("test").build();
768 let (session, evt_tx) = mock_agent_session();
769 let mut ctx = InvocationContext::new(session);
770 let mut agent_events = ctx.subscribe();
771
772 tokio::spawn(async move {
773 tokio::time::sleep(std::time::Duration::from_millis(10)).await;
774 let _ = evt_tx.send(SessionEvent::ToolCall(vec![FunctionCall {
775 name: "nonexistent_tool".to_string(),
776 args: json!({}),
777 id: Some("call-1".to_string()),
778 }]));
779 tokio::time::sleep(std::time::Duration::from_millis(50)).await;
780 let _ = evt_tx.send(SessionEvent::TurnComplete);
781 });
782
783 let result = agent.run_live(&mut ctx).await;
784 assert!(result.is_ok());
785
786 let mut saw_tool_failed = false;
788 while let Ok(event) = agent_events.try_recv() {
789 if let AgentEvent::ToolCallFailed { name, error } = event
790 && name == "nonexistent_tool"
791 {
792 assert!(error.contains("not found") || error.contains("Not found"));
793 saw_tool_failed = true;
794 }
795 }
796 assert!(
797 saw_tool_failed,
798 "should have emitted ToolCallFailed for unknown tool"
799 );
800 }
801
802 #[tokio::test]
803 async fn event_loop_detects_transfer() {
804 let sub = NoopAgent {
805 name: "billing".to_string(),
806 };
807 let agent = LlmAgent::builder("root").sub_agent(sub).build();
808
809 let (session, evt_tx) = mock_agent_session();
810 let mut ctx = InvocationContext::new(session);
811 let mut agent_events = ctx.subscribe();
812
813 tokio::spawn(async move {
814 tokio::time::sleep(std::time::Duration::from_millis(10)).await;
815 let _ = evt_tx.send(SessionEvent::ToolCall(vec![FunctionCall {
816 name: "transfer_to_billing".to_string(),
817 args: json!({}),
818 id: Some("call-1".to_string()),
819 }]));
820 });
821
822 let result = agent.run_live(&mut ctx).await;
823 match result {
824 Err(AgentError::TransferRequested(target)) => assert_eq!(target, "billing"),
825 other => panic!("expected TransferRequested, got: {other:?}"),
826 }
827
828 let mut saw_transfer = false;
830 while let Ok(event) = agent_events.try_recv() {
831 if let AgentEvent::AgentTransfer { from, to } = event {
832 assert_eq!(from, "root");
833 assert_eq!(to, "billing");
834 saw_transfer = true;
835 }
836 }
837 assert!(saw_transfer, "should have emitted AgentTransfer event");
838 }
839
840 #[tokio::test]
841 async fn event_loop_passes_through_events() {
842 let agent = LlmAgent::builder("test").build();
843 let (session, evt_tx) = mock_agent_session();
844 let mut ctx = InvocationContext::new(session);
845 let mut agent_events = ctx.subscribe();
846
847 tokio::spawn(async move {
848 tokio::time::sleep(std::time::Duration::from_millis(10)).await;
849 let _ = evt_tx.send(SessionEvent::TextDelta("hello".to_string()));
850 tokio::time::sleep(std::time::Duration::from_millis(10)).await;
851 let _ = evt_tx.send(SessionEvent::TurnComplete);
852 });
853
854 agent.run_live(&mut ctx).await.unwrap();
855
856 let mut saw_text_delta = false;
858 let mut saw_started = false;
859 let mut saw_completed = false;
860 while let Ok(event) = agent_events.try_recv() {
861 match event {
862 AgentEvent::AgentStarted { .. } => saw_started = true,
863 AgentEvent::AgentCompleted { .. } => saw_completed = true,
864 AgentEvent::Session(SessionEvent::TextDelta(t)) if t == "hello" => {
865 saw_text_delta = true;
866 }
867 _ => {}
868 }
869 }
870 assert!(saw_started, "should have emitted AgentStarted");
871 assert!(saw_text_delta, "should have passed through TextDelta");
872 assert!(saw_completed, "should have emitted AgentCompleted");
873 }
874
875 #[tokio::test]
876 async fn event_loop_handles_error_event() {
877 let agent = LlmAgent::builder("test").build();
878 let (session, evt_tx) = mock_agent_session();
879 let mut ctx = InvocationContext::new(session);
880 let mut agent_events = ctx.subscribe();
881
882 tokio::spawn(async move {
883 tokio::time::sleep(std::time::Duration::from_millis(10)).await;
884 let _ = evt_tx.send(SessionEvent::Error(
885 gemini_genai_rs::session::SessionError::WebSocket(
886 gemini_genai_rs::session::WebSocketError::ProtocolError(
887 "something broke".to_string(),
888 ),
889 ),
890 ));
891 tokio::time::sleep(std::time::Duration::from_millis(10)).await;
892 let _ = evt_tx.send(SessionEvent::TurnComplete);
893 });
894
895 agent.run_live(&mut ctx).await.unwrap();
896
897 let mut saw_error = false;
899 while let Ok(event) = agent_events.try_recv() {
900 if let AgentEvent::Session(SessionEvent::Error(e)) = event {
901 assert!(e.to_string().contains("something broke"), "{e}");
902 saw_error = true;
903 }
904 }
905 assert!(saw_error, "should have passed through Error event");
906 }
907
908 #[tokio::test]
909 async fn event_loop_emits_lifecycle_events() {
910 let agent = LlmAgent::builder("lifecycle_test").build();
911 let (session, evt_tx) = mock_agent_session();
912 let mut ctx = InvocationContext::new(session);
913 let mut agent_events = ctx.subscribe();
914
915 tokio::spawn(async move {
916 tokio::time::sleep(std::time::Duration::from_millis(10)).await;
917 let _ = evt_tx.send(SessionEvent::TurnComplete);
918 });
919
920 agent.run_live(&mut ctx).await.unwrap();
921
922 let mut events = Vec::new();
923 while let Ok(event) = agent_events.try_recv() {
924 events.push(event);
925 }
926
927 assert!(
929 matches!(&events[0], AgentEvent::AgentStarted { name } if name == "lifecycle_test"),
930 "first event should be AgentStarted, got: {:?}",
931 events[0]
932 );
933
934 let last = events.last().unwrap();
936 assert!(
937 matches!(last, AgentEvent::AgentCompleted { name } if name == "lifecycle_test"),
938 "last event should be AgentCompleted, got: {last:?}"
939 );
940 }
941
942 #[tokio::test]
943 async fn event_loop_tool_failure_emits_failed_event() {
944 let tool = SimpleTool::new("failing_tool", "Always fails", None, |_| async {
945 Err(ToolError::ExecutionFailed("kaboom".to_string()))
946 });
947 let agent = LlmAgent::builder("test").tool(tool).build();
948 let (session, evt_tx) = mock_agent_session();
949 let mut ctx = InvocationContext::new(session);
950 let mut agent_events = ctx.subscribe();
951
952 tokio::spawn(async move {
953 tokio::time::sleep(std::time::Duration::from_millis(10)).await;
954 let _ = evt_tx.send(SessionEvent::ToolCall(vec![FunctionCall {
955 name: "failing_tool".to_string(),
956 args: json!({}),
957 id: Some("call-1".to_string()),
958 }]));
959 tokio::time::sleep(std::time::Duration::from_millis(50)).await;
960 let _ = evt_tx.send(SessionEvent::TurnComplete);
961 });
962
963 agent.run_live(&mut ctx).await.unwrap();
964
965 let mut saw_tool_failed = false;
966 while let Ok(event) = agent_events.try_recv() {
967 if let AgentEvent::ToolCallFailed { name, error } = event
968 && name == "failing_tool"
969 {
970 assert!(error.contains("kaboom"));
971 saw_tool_failed = true;
972 }
973 }
974 assert!(saw_tool_failed, "should have emitted ToolCallFailed");
975 }
976}