gemini_genai_rs/transport/connection/
mod.rs

1//! WebSocket connection lifecycle — connect, setup, full-duplex split, reconnection.
2
3mod message_handler;
4mod reconnect;
5mod session_loop;
6
7use std::sync::Arc;
8
9use tokio::sync::{broadcast, mpsc, watch};
10
11use crate::protocol::types::*;
12use crate::session::{SessionHandle, SessionPhase, SessionState};
13use crate::transport::TransportConfig;
14use crate::transport::codec::{Codec, JsonCodec};
15use crate::transport::ws::{Transport, TungsteniteTransport};
16
17/// Connect to the Gemini Multimodal Live API with the default transport and
18/// return a session handle.
19///
20/// Timeouts, reconnection, a custom transport or codec, or a wire recorder:
21/// [`ConnectBuilder`](crate::transport::ConnectBuilder), of which this is the
22/// zero-option form.
23pub async fn connect(config: SessionConfig) -> Result<SessionHandle, crate::session::SessionError> {
24    connect_with(
25        config,
26        TransportConfig::default(),
27        TungsteniteTransport::new(),
28        JsonCodec,
29    )
30    .await
31}
32
33/// Connect with an explicit transport config, transport, and codec — the
34/// one path every public entry point ends in.
35pub(crate) async fn connect_with<T, C>(
36    config: SessionConfig,
37    transport_config: TransportConfig,
38    transport: T,
39    codec: C,
40) -> Result<SessionHandle, crate::session::SessionError>
41where
42    T: Transport,
43    C: Codec,
44{
45    let (command_tx, command_rx) = mpsc::channel(transport_config.send_queue_depth);
46    let (event_tx, _) = broadcast::channel(transport_config.event_channel_capacity);
47    let (phase_tx, phase_rx) = watch::channel(SessionPhase::Disconnected);
48
49    let state = Arc::new(SessionState::with_events(phase_tx, event_tx.clone()));
50
51    let mut handle = SessionHandle::new(command_tx, event_tx.clone(), state.clone(), phase_rx);
52    if let Some(pacing) = config.audio_pacing.clone() {
53        handle = handle.with_audio_pacing(pacing);
54    }
55
56    // Honor a config-installed wire recorder by wrapping the codec. The
57    // boxed indirection keeps `connect_with` generic while letting the
58    // recorder be a runtime decision.
59    let task = if let Some(recorder) = config.wire_recorder.clone() {
60        let codec: Box<dyn Codec> = Box::new(crate::transport::recording::RecordingCodec::new(
61            codec,
62            recorder.recorder(),
63        ));
64        tokio::spawn(async move {
65            session_loop::generic_connection_loop(
66                config,
67                transport_config,
68                state,
69                command_rx,
70                event_tx,
71                transport,
72                codec,
73            )
74            .await;
75        })
76    } else {
77        tokio::spawn(async move {
78            session_loop::generic_connection_loop(
79                config,
80                transport_config,
81                state,
82                command_rx,
83                event_tx,
84                transport,
85                codec,
86            )
87            .await;
88        })
89    };
90    handle.set_task(task);
91
92    Ok(handle)
93}
94
95#[cfg(test)]
96mod tests {
97    use super::message_handler::{MessageAction, handle_server_msg};
98    use super::reconnect::reconnect_delay;
99    use super::*;
100
101    use std::time::Duration;
102
103    use crate::protocol::messages::ServerMessage;
104    use crate::session::{SessionEvent, SessionPhase, SessionState};
105    use crate::transport::codec::JsonCodec;
106    use crate::transport::ws::MockTransport;
107
108    /// TransportConfig that disables reconnection for mock tests.
109    fn no_reconnect_config() -> TransportConfig {
110        TransportConfig {
111            max_reconnect_attempts: 0,
112            connect_timeout_secs: 5,
113            setup_timeout_secs: 5,
114            ..TransportConfig::default()
115        }
116    }
117
118    #[tokio::test]
119    async fn connect_with_mock_transport() {
120        let mut transport = MockTransport::new();
121        // Script setupComplete response
122        transport.script_recv(br#"{"setupComplete":{}}"#.to_vec());
123        // Script a text response then turn complete
124        transport.script_recv(
125            br#"{"serverContent":{"modelTurn":{"parts":[{"text":"Hello!"}]},"turnComplete":true}}"#
126                .to_vec(),
127        );
128
129        let config = SessionConfig::new("test-key")
130            .model(ModelId::from_static("models/gemini-2.0-flash-live-001"));
131
132        let handle = connect_with(config, no_reconnect_config(), transport, JsonCodec)
133            .await
134            .unwrap();
135
136        // Should reach Active phase after setup completes
137        handle.wait_for_phase(SessionPhase::Active).await;
138        assert_eq!(handle.phase(), SessionPhase::Active);
139    }
140
141    #[tokio::test]
142    async fn connect_with_mock_receives_text_events() {
143        let mut transport = MockTransport::new();
144        transport.script_recv(br#"{"setupComplete":{}}"#.to_vec());
145        transport.script_recv(
146            br#"{"serverContent":{"modelTurn":{"parts":[{"text":"Hello from mock!"}]},"turnComplete":true}}"#
147                .to_vec(),
148        );
149
150        let config = SessionConfig::new("test-key")
151            .model(ModelId::from_static("models/gemini-2.0-flash-live-001"));
152        let handle = connect_with(config, no_reconnect_config(), transport, JsonCodec)
153            .await
154            .unwrap();
155
156        let mut events = handle.subscribe();
157
158        // Wait for the session to become active
159        handle.wait_for_phase(SessionPhase::Active).await;
160
161        // Collect events until TurnComplete
162        let mut got_text_delta = false;
163        let mut got_text_complete = false;
164        let mut got_turn_complete = false;
165
166        for _ in 0..20 {
167            match tokio::time::timeout(Duration::from_millis(100), events.recv()).await {
168                Ok(Ok(SessionEvent::TextDelta(t))) => {
169                    assert_eq!(t, "Hello from mock!");
170                    got_text_delta = true;
171                }
172                Ok(Ok(SessionEvent::TextComplete(t))) => {
173                    assert_eq!(t, "Hello from mock!");
174                    got_text_complete = true;
175                }
176                Ok(Ok(SessionEvent::TurnComplete)) => {
177                    got_turn_complete = true;
178                    break;
179                }
180                Ok(Ok(_)) => continue,
181                Ok(Err(_)) => break,
182                Err(_) => break,
183            }
184        }
185
186        assert!(got_text_delta, "should have received TextDelta");
187        assert!(got_text_complete, "should have received TextComplete");
188        assert!(got_turn_complete, "should have received TurnComplete");
189    }
190
191    #[tokio::test]
192    async fn connect_with_mock_tool_call() {
193        let mut transport = MockTransport::new();
194        transport.script_recv(br#"{"setupComplete":{}}"#.to_vec());
195        transport.script_recv(
196            br#"{"toolCall":{"functionCalls":[{"name":"get_weather","args":{"city":"London"},"id":"call-1"}]}}"#
197                .to_vec(),
198        );
199
200        let config = SessionConfig::new("test-key")
201            .model(ModelId::from_static("models/gemini-2.0-flash-live-001"));
202        let handle = connect_with(config, no_reconnect_config(), transport, JsonCodec)
203            .await
204            .unwrap();
205
206        let mut events = handle.subscribe();
207        handle.wait_for_phase(SessionPhase::Active).await;
208
209        // Look for the ToolCall event
210        let mut got_tool_call = false;
211        for _ in 0..20 {
212            match tokio::time::timeout(Duration::from_millis(100), events.recv()).await {
213                Ok(Ok(SessionEvent::ToolCall(calls))) => {
214                    assert_eq!(calls.len(), 1);
215                    assert_eq!(calls[0].name, "get_weather");
216                    got_tool_call = true;
217                    break;
218                }
219                Ok(Ok(_)) => continue,
220                Ok(Err(_)) => break,
221                Err(_) => break,
222            }
223        }
224
225        assert!(got_tool_call, "should have received ToolCall event");
226    }
227
228    #[tokio::test]
229    async fn connect_with_mock_graceful_disconnect() {
230        let mut transport = MockTransport::new();
231        transport.script_recv(br#"{"setupComplete":{}}"#.to_vec());
232        // Keep the connection alive with a message that arrives before disconnect
233        transport.script_recv(
234            br#"{"serverContent":{"modelTurn":{"parts":[{"text":"hi"}]},"turnComplete":true}}"#
235                .to_vec(),
236        );
237
238        let config = SessionConfig::new("test-key")
239            .model(ModelId::from_static("models/gemini-2.0-flash-live-001"));
240        let handle = connect_with(config, no_reconnect_config(), transport, JsonCodec)
241            .await
242            .unwrap();
243
244        handle.wait_for_phase(SessionPhase::Active).await;
245        // Small delay to let the background task process
246        tokio::time::sleep(Duration::from_millis(50)).await;
247
248        // Disconnect gracefully
249        handle.disconnect().await.unwrap();
250
251        // Wait for disconnected phase
252        handle.wait_for_phase(SessionPhase::Disconnected).await;
253        assert_eq!(handle.phase(), SessionPhase::Disconnected);
254    }
255
256    #[test]
257    fn handle_server_msg_preserves_interruption() {
258        let (phase_tx, _phase_rx) = watch::channel(SessionPhase::Active);
259        let (event_tx, mut event_rx) = broadcast::channel(16);
260        let state = Arc::new(SessionState::with_events(phase_tx, event_tx.clone()));
261
262        let json = r#"{"serverContent":{"interrupted":true}}"#;
263        let msg = ServerMessage::parse(json).unwrap();
264        let action = handle_server_msg(msg, &state, &event_tx);
265
266        assert!(matches!(action, MessageAction::Continue));
267        // Should have emitted Interrupted event
268        let mut found_interrupted = false;
269        while let Ok(evt) = event_rx.try_recv() {
270            if matches!(evt, SessionEvent::Interrupted) {
271                found_interrupted = true;
272            }
273        }
274        assert!(found_interrupted, "should emit Interrupted event");
275    }
276
277    #[test]
278    fn handle_server_msg_go_away() {
279        let (phase_tx, _phase_rx) = watch::channel(SessionPhase::Active);
280        let (event_tx, _event_rx) = broadcast::channel(16);
281        let state = Arc::new(SessionState::with_events(phase_tx, event_tx.clone()));
282
283        let json = r#"{"goAway":{"timeLeft":"30s"}}"#;
284        let msg = ServerMessage::parse(json).unwrap();
285        let action = handle_server_msg(msg, &state, &event_tx);
286
287        assert!(matches!(action, MessageAction::GoAway(Some(_))));
288    }
289
290    #[test]
291    fn handle_server_msg_unknown_is_continue() {
292        let (phase_tx, _phase_rx) = watch::channel(SessionPhase::Active);
293        let (event_tx, _event_rx) = broadcast::channel(16);
294        let state = Arc::new(SessionState::with_events(phase_tx, event_tx.clone()));
295
296        let json = r#"{"unknownField":{"data":"test"}}"#;
297        let msg = ServerMessage::parse(json).unwrap();
298        let action = handle_server_msg(msg, &state, &event_tx);
299
300        assert!(matches!(action, MessageAction::Continue));
301    }
302
303    #[tokio::test]
304    async fn session_handle_join_after_disconnect() {
305        let mut transport = MockTransport::new();
306        transport.script_recv(br#"{"setupComplete":{}}"#.to_vec());
307
308        let config = SessionConfig::new("test-key")
309            .model(ModelId::from_static("models/gemini-2.0-flash-live-001"));
310        let handle = connect_with(config, no_reconnect_config(), transport, JsonCodec)
311            .await
312            .unwrap();
313
314        handle.wait_for_phase(SessionPhase::Active).await;
315
316        // Disconnect to end the connection loop task
317        handle.disconnect().await.unwrap();
318        handle.wait_for_phase(SessionPhase::Disconnected).await;
319
320        // join() should return Ok after the task completes
321        let result = handle.join().await;
322        assert!(result.is_ok(), "join() should succeed after disconnect");
323    }
324
325    #[tokio::test]
326    async fn session_handle_join_after_command_channel_closed() {
327        let mut transport = MockTransport::new();
328        transport.script_recv(br#"{"setupComplete":{}}"#.to_vec());
329
330        let config = SessionConfig::new("test-key")
331            .model(ModelId::from_static("models/gemini-2.0-flash-live-001"));
332        let handle = connect_with(config, no_reconnect_config(), transport, JsonCodec)
333            .await
334            .unwrap();
335
336        handle.wait_for_phase(SessionPhase::Active).await;
337
338        // Drop all senders to close the command channel, which triggers disconnect
339        // We need to get the handle before dropping the original
340        let join_handle = handle.clone();
341
342        // Drop command_tx by dropping the handle — but we cloned it first.
343        // Instead, disconnect and then join.
344        handle.disconnect().await.unwrap();
345
346        let result = join_handle.join().await;
347        assert!(result.is_ok(), "join() should succeed after channel close");
348    }
349
350    #[test]
351    fn reconnect_delay_exponential_backoff() {
352        let config = TransportConfig::default();
353        let d1 = reconnect_delay(1, &config);
354        let d2 = reconnect_delay(2, &config);
355        let d3 = reconnect_delay(3, &config);
356        // Each step should roughly double (plus jitter)
357        assert!(d2 > d1);
358        assert!(d3 > d2);
359        // Should not exceed max
360        let d_large = reconnect_delay(100, &config);
361        let max_with_jitter = Duration::from_millis(
362            config.reconnect_max_delay_ms as u64 + config.reconnect_max_delay_ms as u64 / 4,
363        );
364        assert!(d_large <= max_with_jitter);
365    }
366}