gemini_genai_rs/transport/
builder.rs1use crate::protocol::types::SessionConfig;
4use crate::session::{SessionError, SessionHandle};
5use crate::transport::TransportConfig;
6use crate::transport::codec::{Codec, JsonCodec};
7use crate::transport::connection::connect_with;
8use crate::transport::ws::{Transport, TungsteniteTransport};
9
10pub struct ConnectBuilder<T = TungsteniteTransport, C = JsonCodec> {
28 config: SessionConfig,
29 transport_config: TransportConfig,
30 transport: T,
31 codec: C,
32}
33
34impl ConnectBuilder {
35 pub fn new(config: SessionConfig) -> Self {
37 Self {
38 config,
39 transport_config: TransportConfig::default(),
40 transport: TungsteniteTransport::new(),
41 codec: JsonCodec,
42 }
43 }
44}
45
46impl<T: Transport, C: Codec> ConnectBuilder<T, C> {
47 pub fn transport_config(mut self, tc: TransportConfig) -> Self {
49 self.transport_config = tc;
50 self
51 }
52
53 pub fn configure(mut self, f: impl FnOnce(SessionConfig) -> SessionConfig) -> Self {
56 self.config = f(self.config);
57 self
58 }
59
60 pub fn transport<T2: Transport>(self, transport: T2) -> ConnectBuilder<T2, C> {
62 ConnectBuilder {
63 config: self.config,
64 transport_config: self.transport_config,
65 transport,
66 codec: self.codec,
67 }
68 }
69
70 pub fn record_wire(
75 mut self,
76 recorder: std::sync::Arc<dyn crate::transport::recording::WireRecorder>,
77 ) -> Self {
78 self.config = self.config.record_wire(recorder);
79 self
80 }
81
82 pub fn codec<C2: Codec>(self, codec: C2) -> ConnectBuilder<T, C2> {
84 ConnectBuilder {
85 config: self.config,
86 transport_config: self.transport_config,
87 transport: self.transport,
88 codec,
89 }
90 }
91
92 pub async fn connect(self) -> Result<SessionHandle, SessionError> {
94 connect_with(
95 self.config,
96 self.transport_config,
97 self.transport,
98 self.codec,
99 )
100 .await
101 }
102}
103
104#[cfg(test)]
105mod tests {
106 use super::*;
107 use crate::protocol::types::*;
108 use crate::transport::ws::MockTransport;
109
110 #[test]
111 fn builder_compiles_with_defaults() {
112 let config = SessionConfig::new("key")
113 .model(ModelId::from_static("models/gemini-2.0-flash-live-001"));
114 let _builder = ConnectBuilder::new(config);
115 }
116
117 #[test]
118 fn builder_with_custom_transport_config() {
119 let config = SessionConfig::new("key");
120 let _builder = ConnectBuilder::new(config).transport_config(TransportConfig {
121 connect_timeout_secs: 30,
122 ..Default::default()
123 });
124 }
125
126 #[test]
127 fn builder_with_mock_transport() {
128 let config = SessionConfig::new("key");
129 let mock = MockTransport::new();
130 let _builder = ConnectBuilder::new(config).transport(mock);
131 }
132
133 #[test]
134 fn builder_with_custom_codec() {
135 let config = SessionConfig::new("key");
136 let _builder = ConnectBuilder::new(config).codec(JsonCodec);
137 }
138
139 #[tokio::test]
140 async fn builder_with_mock_builds() {
141 let mut mock = MockTransport::new();
142 mock.script_recv(br#"{"setupComplete":{}}"#.to_vec());
143
144 let config = SessionConfig::new("key")
145 .model(ModelId::from_static("models/gemini-2.0-flash-live-001"));
146 let handle = ConnectBuilder::new(config)
147 .transport(mock)
148 .connect()
149 .await
150 .unwrap();
151
152 handle
153 .wait_for_phase(crate::session::SessionPhase::Active)
154 .await;
155 assert_eq!(handle.phase(), crate::session::SessionPhase::Active);
156 }
157}