gemini_genai_rs/transport/connection/
mod.rs1mod 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 #[tokio::test]
119 async fn connect_with_mock_transport() {
120 let mut transport = MockTransport::new();
121 transport.script_recv(br#"{"setupComplete":{}}"#.to_vec());
123 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 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 handle.wait_for_phase(SessionPhase::Active).await;
160
161 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 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 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 tokio::time::sleep(Duration::from_millis(50)).await;
247
248 handle.disconnect().await.unwrap();
250
251 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 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 handle.disconnect().await.unwrap();
318 handle.wait_for_phase(SessionPhase::Disconnected).await;
319
320 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 let join_handle = handle.clone();
341
342 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 assert!(d2 > d1);
358 assert!(d3 > d2);
359 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}