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    /// A transport that plays one script per connection and records every
119    /// setup message it is sent.
120    struct Reconnecting {
121        scripts: std::collections::VecDeque<Vec<Vec<u8>>>,
122        current: std::collections::VecDeque<Vec<u8>>,
123        setups: std::sync::Arc<parking_lot::Mutex<Vec<serde_json::Value>>>,
124    }
125
126    #[async_trait::async_trait]
127    impl crate::transport::ws::Transport for Reconnecting {
128        type Error = std::io::Error;
129
130        async fn connect(
131            &mut self,
132            _url: &str,
133            _headers: Vec<(String, String)>,
134        ) -> Result<(), Self::Error> {
135            self.current = self.scripts.pop_front().unwrap_or_default().into();
136            Ok(())
137        }
138
139        async fn send(&mut self, data: Vec<u8>) -> Result<(), Self::Error> {
140            if let Ok(message) = serde_json::from_slice::<serde_json::Value>(&data)
141                && message.get("setup").is_some()
142            {
143                self.setups.lock().push(message);
144            }
145            Ok(())
146        }
147
148        async fn recv(&mut self) -> Result<Option<Vec<u8>>, Self::Error> {
149            match self.current.pop_front() {
150                Some(frame) => Ok(Some(frame)),
151                // Out of script: stay open and quiet.
152                None => std::future::pending().await,
153            }
154        }
155
156        async fn close(&mut self) -> Result<(), Self::Error> {
157            Ok(())
158        }
159    }
160
161    /// A server that refuses every setup, closing with a status and reason.
162    struct Refusing {
163        connects: std::sync::Arc<std::sync::atomic::AtomicUsize>,
164        reason: &'static str,
165    }
166
167    #[async_trait::async_trait]
168    impl crate::transport::ws::Transport for Refusing {
169        type Error = std::io::Error;
170
171        async fn connect(
172            &mut self,
173            _url: &str,
174            _headers: Vec<(String, String)>,
175        ) -> Result<(), Self::Error> {
176            self.connects
177                .fetch_add(1, std::sync::atomic::Ordering::SeqCst);
178            Ok(())
179        }
180
181        async fn send(&mut self, _data: Vec<u8>) -> Result<(), Self::Error> {
182            Ok(())
183        }
184
185        async fn recv(&mut self) -> Result<Option<Vec<u8>>, Self::Error> {
186            // Give the test time to subscribe before the refusal.
187            tokio::time::sleep(Duration::from_millis(50)).await;
188            Ok(None)
189        }
190
191        async fn close(&mut self) -> Result<(), Self::Error> {
192            Ok(())
193        }
194
195        fn close_reason(&self) -> Option<String> {
196            Some(self.reason.to_string())
197        }
198    }
199
200    #[tokio::test]
201    async fn an_invalid_setup_is_reported_with_the_servers_reason_and_not_retried() {
202        let connects = std::sync::Arc::new(std::sync::atomic::AtomicUsize::new(0));
203        let transport = Refusing {
204            connects: connects.clone(),
205            reason: "server closed the connection (1007): The requested combination of response modalities (TEXT) is not supported by the model.",
206        };
207        let transport_config = TransportConfig {
208            max_reconnect_attempts: 3,
209            reconnect_base_delay_ms: 10,
210            reconnect_max_delay_ms: 10,
211            ..no_reconnect_config()
212        };
213        let handle = connect_with(
214            SessionConfig::new("test-key"),
215            transport_config,
216            transport,
217            JsonCodec,
218        )
219        .await
220        .unwrap();
221        let mut events = handle.subscribe();
222        let mut error = None;
223        let mut disconnected = None;
224        let deadline = tokio::time::Instant::now() + Duration::from_secs(5);
225        while disconnected.is_none() && tokio::time::Instant::now() < deadline {
226            match tokio::time::timeout(Duration::from_millis(200), events.recv()).await {
227                Ok(Ok(SessionEvent::Error(e))) => error = Some(e.to_string()),
228                Ok(Ok(SessionEvent::Disconnected(reason))) => disconnected = Some(reason),
229                _ => {}
230            }
231        }
232        let error = error.expect("the refusal is reported");
233        assert!(error.contains("response modalities (TEXT)"), "{error}");
234        assert!(disconnected.flatten().unwrap_or_default().contains("1007"));
235        assert_eq!(
236            connects.load(std::sync::atomic::Ordering::SeqCst),
237            1,
238            "no retry"
239        );
240    }
241
242    #[test]
243    fn close_codes_are_read_from_the_reason() {
244        assert_eq!(
245            super::session_loop::close_code("server closed the connection (1007): x"),
246            Some(1007)
247        );
248        assert_eq!(
249            super::session_loop::close_code("server closed the connection (1011)"),
250            Some(1011)
251        );
252        assert_eq!(super::session_loop::close_code("no code"), None);
253    }
254
255    #[tokio::test]
256    async fn a_reconnect_after_go_away_resumes_with_the_latest_handle() {
257        let setups = std::sync::Arc::new(parking_lot::Mutex::new(Vec::new()));
258        let transport = Reconnecting {
259            scripts: vec![
260                vec![
261                    br#"{"setupComplete":{}}"#.to_vec(),
262                    br#"{"sessionResumptionUpdate":{"newHandle":"h-1","resumable":true}}"#.to_vec(),
263                    br#"{"goAway":{"timeLeft":"0s"}}"#.to_vec(),
264                ],
265                vec![br#"{"setupComplete":{}}"#.to_vec()],
266            ]
267            .into(),
268            current: Default::default(),
269            setups: setups.clone(),
270        };
271        let mut config = SessionConfig::new("test-key")
272            .model(ModelId::from_static("models/gemini-2.0-flash-live-001"));
273        config.session_resumption = Some(SessionResumptionConfig {
274            handle: None,
275            transparent: None,
276        });
277        let transport_config = TransportConfig {
278            max_reconnect_attempts: 1,
279            reconnect_base_delay_ms: 10,
280            reconnect_max_delay_ms: 10,
281            ..no_reconnect_config()
282        };
283        let _handle = connect_with(config, transport_config, transport, JsonCodec)
284            .await
285            .unwrap();
286
287        let deadline = tokio::time::Instant::now() + Duration::from_secs(5);
288        while setups.lock().len() < 2 && tokio::time::Instant::now() < deadline {
289            tokio::time::sleep(Duration::from_millis(10)).await;
290        }
291        let setups = setups.lock();
292        assert_eq!(setups.len(), 2, "one setup per connection");
293        let resumption = |i: usize| setups[i].pointer("/setup/sessionResumption").cloned();
294        assert_eq!(
295            resumption(0),
296            Some(serde_json::json!({})),
297            "first connect: no handle yet"
298        );
299        assert_eq!(
300            resumption(1),
301            Some(serde_json::json!({ "handle": "h-1" })),
302            "the reconnect presents the handle the server issued"
303        );
304    }
305
306    #[tokio::test]
307    async fn connect_with_mock_transport() {
308        let mut transport = MockTransport::new();
309        // Script setupComplete response
310        transport.script_recv(br#"{"setupComplete":{}}"#.to_vec());
311        // Script a text response then turn complete
312        transport.script_recv(
313            br#"{"serverContent":{"modelTurn":{"parts":[{"text":"Hello!"}]},"turnComplete":true}}"#
314                .to_vec(),
315        );
316
317        let config = SessionConfig::new("test-key")
318            .model(ModelId::from_static("models/gemini-2.0-flash-live-001"));
319
320        let handle = connect_with(config, no_reconnect_config(), transport, JsonCodec)
321            .await
322            .unwrap();
323
324        // Should reach Active phase after setup completes
325        handle.wait_for_phase(SessionPhase::Active).await;
326        assert_eq!(handle.phase(), SessionPhase::Active);
327    }
328
329    #[tokio::test]
330    async fn connect_with_mock_receives_text_events() {
331        let mut transport = MockTransport::new();
332        transport.script_recv(br#"{"setupComplete":{}}"#.to_vec());
333        transport.script_recv(
334            br#"{"serverContent":{"modelTurn":{"parts":[{"text":"Hello from mock!"}]},"turnComplete":true}}"#
335                .to_vec(),
336        );
337
338        let config = SessionConfig::new("test-key")
339            .model(ModelId::from_static("models/gemini-2.0-flash-live-001"));
340        let handle = connect_with(config, no_reconnect_config(), transport, JsonCodec)
341            .await
342            .unwrap();
343
344        let mut events = handle.subscribe();
345
346        // Wait for the session to become active
347        handle.wait_for_phase(SessionPhase::Active).await;
348
349        // Collect events until TurnComplete
350        let mut got_text_delta = false;
351        let mut got_text_complete = false;
352        let mut got_turn_complete = false;
353
354        for _ in 0..20 {
355            match tokio::time::timeout(Duration::from_millis(100), events.recv()).await {
356                Ok(Ok(SessionEvent::TextDelta(t))) => {
357                    assert_eq!(t, "Hello from mock!");
358                    got_text_delta = true;
359                }
360                Ok(Ok(SessionEvent::TextComplete(t))) => {
361                    assert_eq!(t, "Hello from mock!");
362                    got_text_complete = true;
363                }
364                Ok(Ok(SessionEvent::TurnComplete)) => {
365                    got_turn_complete = true;
366                    break;
367                }
368                Ok(Ok(_)) => continue,
369                Ok(Err(_)) => break,
370                Err(_) => break,
371            }
372        }
373
374        assert!(got_text_delta, "should have received TextDelta");
375        assert!(got_text_complete, "should have received TextComplete");
376        assert!(got_turn_complete, "should have received TurnComplete");
377    }
378
379    #[tokio::test]
380    async fn connect_with_mock_tool_call() {
381        let mut transport = MockTransport::new();
382        transport.script_recv(br#"{"setupComplete":{}}"#.to_vec());
383        transport.script_recv(
384            br#"{"toolCall":{"functionCalls":[{"name":"get_weather","args":{"city":"London"},"id":"call-1"}]}}"#
385                .to_vec(),
386        );
387
388        let config = SessionConfig::new("test-key")
389            .model(ModelId::from_static("models/gemini-2.0-flash-live-001"));
390        let handle = connect_with(config, no_reconnect_config(), transport, JsonCodec)
391            .await
392            .unwrap();
393
394        let mut events = handle.subscribe();
395        handle.wait_for_phase(SessionPhase::Active).await;
396
397        // Look for the ToolCall event
398        let mut got_tool_call = false;
399        for _ in 0..20 {
400            match tokio::time::timeout(Duration::from_millis(100), events.recv()).await {
401                Ok(Ok(SessionEvent::ToolCall(calls))) => {
402                    assert_eq!(calls.len(), 1);
403                    assert_eq!(calls[0].name, "get_weather");
404                    got_tool_call = true;
405                    break;
406                }
407                Ok(Ok(_)) => continue,
408                Ok(Err(_)) => break,
409                Err(_) => break,
410            }
411        }
412
413        assert!(got_tool_call, "should have received ToolCall event");
414    }
415
416    #[tokio::test]
417    async fn connect_with_mock_graceful_disconnect() {
418        let mut transport = MockTransport::new();
419        transport.script_recv(br#"{"setupComplete":{}}"#.to_vec());
420        // Keep the connection alive with a message that arrives before disconnect
421        transport.script_recv(
422            br#"{"serverContent":{"modelTurn":{"parts":[{"text":"hi"}]},"turnComplete":true}}"#
423                .to_vec(),
424        );
425
426        let config = SessionConfig::new("test-key")
427            .model(ModelId::from_static("models/gemini-2.0-flash-live-001"));
428        let handle = connect_with(config, no_reconnect_config(), transport, JsonCodec)
429            .await
430            .unwrap();
431
432        handle.wait_for_phase(SessionPhase::Active).await;
433        // Small delay to let the background task process
434        tokio::time::sleep(Duration::from_millis(50)).await;
435
436        // Disconnect gracefully
437        handle.disconnect().await.unwrap();
438
439        // Wait for disconnected phase
440        handle.wait_for_phase(SessionPhase::Disconnected).await;
441        assert_eq!(handle.phase(), SessionPhase::Disconnected);
442    }
443
444    #[test]
445    fn handle_server_msg_preserves_interruption() {
446        let (phase_tx, _phase_rx) = watch::channel(SessionPhase::Active);
447        let (event_tx, mut event_rx) = broadcast::channel(16);
448        let state = Arc::new(SessionState::with_events(phase_tx, event_tx.clone()));
449
450        let json = r#"{"serverContent":{"interrupted":true}}"#;
451        let msg = ServerMessage::parse(json).unwrap();
452        let action = handle_server_msg(msg, &state, &event_tx);
453
454        assert!(matches!(action, MessageAction::Continue));
455        // Should have emitted Interrupted event
456        let mut found_interrupted = false;
457        while let Ok(evt) = event_rx.try_recv() {
458            if matches!(evt, SessionEvent::Interrupted) {
459                found_interrupted = true;
460            }
461        }
462        assert!(found_interrupted, "should emit Interrupted event");
463    }
464
465    #[test]
466    fn handle_server_msg_routes_avatar_video_away_from_audio() {
467        let (phase_tx, _phase_rx) = watch::channel(SessionPhase::Active);
468        let (event_tx, mut event_rx) = broadcast::channel(16);
469        let state = Arc::new(SessionState::with_events(phase_tx, event_tx.clone()));
470
471        // "AAEC" = [0, 1, 2]; "AwQF" = [3, 4, 5].
472        let json = r#"{"serverContent":{"modelTurn":{"parts":[
473            {"inlineData":{"mimeType":"audio/pcm;rate=24000","data":"AAEC"}},
474            {"inlineData":{"mimeType":"video/mp4","data":"AwQF"}}
475        ]}}}"#;
476        let msg = ServerMessage::parse(json).unwrap();
477        handle_server_msg(msg, &state, &event_tx);
478
479        let mut audio = Vec::new();
480        let mut media = Vec::new();
481        while let Ok(evt) = event_rx.try_recv() {
482            match evt {
483                SessionEvent::AudioData(bytes) => audio.push(bytes.to_vec()),
484                SessionEvent::Media(m) => media.push(m),
485                _ => {}
486            }
487        }
488        assert_eq!(audio, [vec![0u8, 1, 2]], "only the audio part is audio");
489        assert_eq!(media.len(), 1);
490        assert!(media[0].is_video());
491        assert_eq!(media[0].mime_type, "video/mp4");
492        assert_eq!(media[0].data.as_ref(), [3u8, 4, 5]);
493    }
494
495    #[test]
496    fn a_text_session_on_a_speech_only_model_reads_the_transcript_as_text() {
497        let (phase_tx, _phase_rx) = watch::channel(SessionPhase::Active);
498        let (event_tx, mut event_rx) = broadcast::channel(32);
499        let state = Arc::new(SessionState::with_events(phase_tx, event_tx.clone()));
500        state.set_text_from_transcription(true);
501
502        for json in [
503            r#"{"serverContent":{"modelTurn":{"parts":[{"inlineData":{"mimeType":"audio/pcm;rate=24000","data":"AAEC"}}]}}}"#,
504            r#"{"serverContent":{"outputTranscription":{"text":"A table "}}}"#,
505            r#"{"serverContent":{"outputTranscription":{"text":"for two."}}}"#,
506            r#"{"serverContent":{"turnComplete":true}}"#,
507        ] {
508            handle_server_msg(ServerMessage::parse(json).unwrap(), &state, &event_tx);
509        }
510        let mut deltas = Vec::new();
511        let mut complete = None;
512        while let Ok(evt) = event_rx.try_recv() {
513            match evt {
514                SessionEvent::TextDelta(t) => deltas.push(t),
515                SessionEvent::TextComplete(t) => complete = Some(t),
516                SessionEvent::AudioData(_) => panic!("a text session plays no audio"),
517                SessionEvent::OutputTranscription(_) => panic!("the transcript is the text"),
518                _ => {}
519            }
520        }
521        assert_eq!(deltas, ["A table ", "for two."]);
522        assert_eq!(complete.as_deref(), Some("A table for two."));
523    }
524
525    #[test]
526    fn handle_server_msg_go_away() {
527        let (phase_tx, _phase_rx) = watch::channel(SessionPhase::Active);
528        let (event_tx, _event_rx) = broadcast::channel(16);
529        let state = Arc::new(SessionState::with_events(phase_tx, event_tx.clone()));
530
531        let json = r#"{"goAway":{"timeLeft":"30s"}}"#;
532        let msg = ServerMessage::parse(json).unwrap();
533        let action = handle_server_msg(msg, &state, &event_tx);
534
535        assert!(matches!(action, MessageAction::GoAway(Some(_))));
536    }
537
538    #[test]
539    fn handle_server_msg_unknown_is_continue() {
540        let (phase_tx, _phase_rx) = watch::channel(SessionPhase::Active);
541        let (event_tx, _event_rx) = broadcast::channel(16);
542        let state = Arc::new(SessionState::with_events(phase_tx, event_tx.clone()));
543
544        let json = r#"{"unknownField":{"data":"test"}}"#;
545        let msg = ServerMessage::parse(json).unwrap();
546        let action = handle_server_msg(msg, &state, &event_tx);
547
548        assert!(matches!(action, MessageAction::Continue));
549    }
550
551    #[tokio::test]
552    async fn session_handle_join_after_disconnect() {
553        let mut transport = MockTransport::new();
554        transport.script_recv(br#"{"setupComplete":{}}"#.to_vec());
555
556        let config = SessionConfig::new("test-key")
557            .model(ModelId::from_static("models/gemini-2.0-flash-live-001"));
558        let handle = connect_with(config, no_reconnect_config(), transport, JsonCodec)
559            .await
560            .unwrap();
561
562        handle.wait_for_phase(SessionPhase::Active).await;
563
564        // Disconnect to end the connection loop task
565        handle.disconnect().await.unwrap();
566        handle.wait_for_phase(SessionPhase::Disconnected).await;
567
568        // join() should return Ok after the task completes
569        let result = handle.join().await;
570        assert!(result.is_ok(), "join() should succeed after disconnect");
571    }
572
573    #[tokio::test]
574    async fn session_handle_join_after_command_channel_closed() {
575        let mut transport = MockTransport::new();
576        transport.script_recv(br#"{"setupComplete":{}}"#.to_vec());
577
578        let config = SessionConfig::new("test-key")
579            .model(ModelId::from_static("models/gemini-2.0-flash-live-001"));
580        let handle = connect_with(config, no_reconnect_config(), transport, JsonCodec)
581            .await
582            .unwrap();
583
584        handle.wait_for_phase(SessionPhase::Active).await;
585
586        // Drop all senders to close the command channel, which triggers disconnect
587        // We need to get the handle before dropping the original
588        let join_handle = handle.clone();
589
590        // Drop command_tx by dropping the handle — but we cloned it first.
591        // Instead, disconnect and then join.
592        handle.disconnect().await.unwrap();
593
594        let result = join_handle.join().await;
595        assert!(result.is_ok(), "join() should succeed after channel close");
596    }
597
598    #[test]
599    fn reconnect_delay_exponential_backoff() {
600        let config = TransportConfig::default();
601        let d1 = reconnect_delay(1, &config);
602        let d2 = reconnect_delay(2, &config);
603        let d3 = reconnect_delay(3, &config);
604        // Each step should roughly double (plus jitter)
605        assert!(d2 > d1);
606        assert!(d3 > d2);
607        // Should not exceed max
608        let d_large = reconnect_delay(100, &config);
609        let max_with_jitter = Duration::from_millis(
610            config.reconnect_max_delay_ms as u64 + config.reconnect_max_delay_ms as u64 / 4,
611        );
612        assert!(d_large <= max_with_jitter);
613    }
614}