1mod 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
17pub 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
33pub(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 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 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 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 None => std::future::pending().await,
153 }
154 }
155
156 async fn close(&mut self) -> Result<(), Self::Error> {
157 Ok(())
158 }
159 }
160
161 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 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 transport.script_recv(br#"{"setupComplete":{}}"#.to_vec());
311 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 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 handle.wait_for_phase(SessionPhase::Active).await;
348
349 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 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 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 tokio::time::sleep(Duration::from_millis(50)).await;
435
436 handle.disconnect().await.unwrap();
438
439 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 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 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 handle.disconnect().await.unwrap();
566 handle.wait_for_phase(SessionPhase::Disconnected).await;
567
568 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 let join_handle = handle.clone();
589
590 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 assert!(d2 > d1);
606 assert!(d3 > d2);
607 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}