gemini_adk_rs/
llm_agent.rs

1//! LlmAgent — concrete Agent implementation with builder pattern.
2//!
3//! The builder freezes tools at `build()` time (respecting Gemini Live's
4//! constraint that tools are fixed at session setup). Auto-registers
5//! `transfer_to_{name}` tools for each sub-agent.
6//!
7//! The event loop subscribes to SessionEvents, auto-dispatches tool calls,
8//! detects transfers via `__transfer_to` signal in tool results, and handles
9//! streaming/input-streaming tools.
10
11use 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
31/// Concrete Agent implementation that runs a Gemini Live event loop.
32///
33/// Tools are declared at build time and sent during session setup.
34/// The event loop subscribes to SessionEvents, auto-dispatches tool calls,
35/// detects transfers, and emits AgentEvents.
36pub 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    /// Start building a new LlmAgent.
46    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    /// Access the tool dispatcher (for testing/introspection).
57    pub fn dispatcher(&self) -> &ToolDispatcher {
58        &self.dispatcher
59    }
60
61    /// Access the middleware chain.
62    pub fn middleware(&self) -> &MiddlewareChain {
63        &self.middleware
64    }
65
66    /// Access the plugin manager.
67    pub fn plugins(&self) -> &PluginManager {
68        &self.plugins
69    }
70
71    /// Core event loop -- processes SessionEvents, dispatches tools, detects transfers.
72    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, // channel closed
82            };
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                        // Emit events + middleware hooks
91                        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                        // Plugin before_tool hook — can deny or short-circuit
98                        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                                        // Check for transfer signal
148                                        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                    // Send all responses back to Gemini
223                    ctx.agent_session.send_tool_response(responses).await?;
224
225                    // Handle transfer AFTER sending response
226                    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                    // Pass through all other events (TextDelta, AudioData, etc.)
254                    ctx.emit(AgentEvent::Session(other));
255                }
256            }
257        }
258        Ok(())
259    }
260
261    /// Spawn a streaming or input-streaming tool as a background task.
262    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(_) => {} // shouldn't reach here
339        }
340
341        // Spawn collector: reads yields and forwards as events + sends final FunctionResponse
342        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            // Send final response when tool completes
356            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
375/// Builder for LlmAgent -- fluent API for declaring tools, middleware, sub-agents.
376pub 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    /// Register a regular function tool.
386    pub fn tool(mut self, tool: impl ToolFunction + 'static) -> Self {
387        self.dispatcher.register_function(Arc::new(tool));
388        self
389    }
390
391    /// Register a typed tool with auto-generated JSON Schema.
392    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    /// Register a streaming tool.
401    pub fn streaming_tool(mut self, tool: impl StreamingTool + 'static) -> Self {
402        self.dispatcher.register_streaming(Arc::new(tool));
403        self
404    }
405
406    /// Register an input-streaming tool.
407    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    /// Add middleware to the agent.
413    pub fn middleware(mut self, mw: impl crate::middleware::Middleware + 'static) -> Self {
414        self.middleware.add(Arc::new(mw));
415        self
416    }
417
418    /// Add a plugin to the agent.
419    pub fn plugin(mut self, plugin: impl crate::plugin::Plugin + 'static) -> Self {
420        self.plugins.add(Arc::new(plugin));
421        self
422    }
423
424    /// Register a sub-agent (enables transfer_to_{name} tool).
425    pub fn sub_agent(mut self, agent: impl Agent + 'static) -> Self {
426        self.sub_agents.push(Arc::new(agent));
427        self
428    }
429
430    /// Set the default timeout for tool execution.
431    pub fn tool_timeout(mut self, timeout: Duration) -> Self {
432        self.dispatcher = self.dispatcher.with_timeout(timeout);
433        self
434    }
435
436    /// Build the LlmAgent, freezing all tool declarations.
437    ///
438    /// This:
439    /// 1. Auto-registers `transfer_to_{name}` SimpleTool for each sub_agent
440    /// 2. Prepends TelemetryMiddleware
441    /// 3. Returns the frozen LlmAgent
442    pub fn build(mut self) -> LlmAgent {
443        // Auto-register transfer tools for sub-agents
444        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        // Prepend TelemetryMiddleware so it runs first
463        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        // Telemetry + middleware + plugins
489        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        // Plugin before_agent hook
494        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        // Cleanup
510        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    /// Mock writer that accepts all commands without error.
553    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    /// Create a mock AgentSession backed by MockWriter, returning the session
594    /// and the event sender so tests can inject SessionEvents.
595    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        // my_tool is the only user tool (TelemetryMiddleware doesn't add tools)
616        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        // Should have transfer_to_billing auto-registered
627        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        // TelemetryMiddleware is auto-prepended
666        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        // TelemetryMiddleware (prepended) + LogMiddleware (user-added)
677        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    // ── Event loop tests ──────────────────────────────────────────────────
689
690    #[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        // Send TurnComplete after a short delay
697        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            // The tool response will be sent back; then end the turn.
739            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        // Check that we got ToolCallStarted and ToolCallCompleted events
747        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        // Check that we got a ToolCallFailed event
787        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        // Check that AgentTransfer event was emitted
829        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        // Check that we got AgentStarted, TextDelta passthrough, TurnComplete, AgentCompleted
857        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        // Check that the error event was passed through
898        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        // First event should be AgentStarted
928        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        // Last event should be AgentCompleted
935        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}