flare_core/client/
config.rs1use crate::common::compression::CompressionAlgorithm;
4use crate::common::config_types::{HeartbeatConfig, TlsConfig, TransportProtocol};
5use crate::common::device::DeviceInfo;
6use crate::common::protocol::SerializationFormat;
7use std::time::Duration;
8
9#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
11pub struct ClientConfig {
12 pub server_url: String,
14 pub transport: TransportProtocol,
16 pub transports: Option<Vec<TransportProtocol>>,
19 pub protocol_urls: Option<std::collections::HashMap<TransportProtocol, String>>,
22 pub race_timeout: Option<Duration>,
24 pub serialization_format: SerializationFormat,
26 pub compression: CompressionAlgorithm,
28 pub force_serialization_format: Option<SerializationFormat>,
31 pub force_compression: Option<CompressionAlgorithm>,
33 pub connect_timeout: Duration,
35 pub reconnect_interval: Duration,
37 pub max_reconnect_attempts: Option<u32>,
39 pub heartbeat: HeartbeatConfig,
41 pub tls: TlsConfig,
43 pub connection_id: Option<String>,
45 pub user_id: Option<String>,
47 pub metadata: std::collections::HashMap<String, String>,
49 pub enable_router: bool,
51 pub device_info: Option<DeviceInfo>,
53 pub token: Option<String>,
55}
56
57impl Default for ClientConfig {
58 fn default() -> Self {
59 Self {
60 server_url: "ws://localhost:8080".to_string(),
61 transport: TransportProtocol::WebSocket,
62 transports: None,
63 protocol_urls: None,
64 race_timeout: Some(Duration::from_secs(5)),
65 serialization_format: SerializationFormat::Json,
67 compression: CompressionAlgorithm::None,
68 force_serialization_format: None,
69 force_compression: None,
70 connect_timeout: Duration::from_secs(30),
71 reconnect_interval: Duration::from_secs(5),
72 max_reconnect_attempts: Some(5),
73 heartbeat: HeartbeatConfig::default(),
74 tls: TlsConfig::none(),
75 connection_id: None,
76 user_id: None,
77 metadata: std::collections::HashMap::new(),
78 enable_router: false,
79 device_info: None,
80 token: None,
81 }
82 }
83}
84
85impl ClientConfig {
86 pub fn new(server_url: String) -> Self {
88 Self {
89 server_url,
90 ..Default::default()
91 }
92 }
93
94 pub fn websocket(mut self) -> Self {
96 self.transport = TransportProtocol::WebSocket;
97 self
98 }
99
100 pub fn quic(mut self) -> Self {
102 self.transport = TransportProtocol::QUIC;
103 self
104 }
105
106 pub fn tcp(mut self) -> Self {
108 self.transport = TransportProtocol::TCP;
109 self
110 }
111
112 pub fn with_format(mut self, format: SerializationFormat) -> Self {
114 self.serialization_format = format;
115 self
116 }
117
118 pub fn with_compression(mut self, compression: CompressionAlgorithm) -> Self {
120 self.compression = compression;
121 self
122 }
123
124 pub fn with_user_id(mut self, user_id: String) -> Self {
126 self.user_id = Some(user_id);
127 self
128 }
129
130 pub fn with_token(mut self, token: String) -> Self {
132 self.token = Some(token);
133 self
134 }
135
136 pub fn with_protocol_race(mut self, protocols: Vec<TransportProtocol>) -> Self {
141 self.transports = Some(protocols);
142 self
143 }
144
145 pub fn with_protocol_url(mut self, protocol: TransportProtocol, url: String) -> Self {
147 if self.protocol_urls.is_none() {
148 self.protocol_urls = Some(std::collections::HashMap::new());
149 }
150 if let Some(ref mut urls) = self.protocol_urls {
151 urls.insert(protocol, url);
152 }
153 self
154 }
155
156 pub fn with_protocol_urls(
158 mut self,
159 urls: std::collections::HashMap<TransportProtocol, String>,
160 ) -> Self {
161 self.protocol_urls = Some(urls);
162 self
163 }
164
165 pub fn get_protocol_url(&self, protocol: &TransportProtocol) -> String {
167 let endpoint = if let Some(ref urls) = self.protocol_urls
168 && let Some(url) = urls.get(protocol)
169 {
170 url.as_str()
171 } else {
172 self.server_url.as_str()
173 };
174
175 protocol.normalize_client_url(endpoint)
176 }
177
178 pub fn with_race_timeout(mut self, timeout: Duration) -> Self {
180 self.race_timeout = Some(timeout);
181 self
182 }
183
184 pub fn with_heartbeat(mut self, heartbeat: HeartbeatConfig) -> Self {
186 self.heartbeat = heartbeat;
187 self
188 }
189
190 pub fn with_tls(mut self, tls: TlsConfig) -> Self {
192 self.tls = tls;
193 self
194 }
195
196 pub fn with_connect_timeout(mut self, timeout: Duration) -> Self {
198 self.connect_timeout = timeout;
199 self
200 }
201
202 pub fn with_reconnect_interval(mut self, interval: Duration) -> Self {
204 self.reconnect_interval = interval;
205 self
206 }
207
208 pub fn with_max_reconnect_attempts(mut self, max: Option<u32>) -> Self {
210 self.max_reconnect_attempts = max;
211 self
212 }
213
214 pub fn enable_router(mut self) -> Self {
216 self.enable_router = true;
217 self
218 }
219
220 pub fn with_device_info(mut self, device_info: DeviceInfo) -> Self {
222 self.device_info = Some(device_info);
223 self
224 }
225
226 pub fn force_format(mut self, format: SerializationFormat) -> Self {
231 self.force_serialization_format = Some(format);
232 self
233 }
234
235 pub fn force_compression(mut self, compression: CompressionAlgorithm) -> Self {
239 self.force_compression = Some(compression);
240 self
241 }
242
243 pub fn is_force_format(&self) -> bool {
245 self.force_serialization_format.is_some() || self.force_compression.is_some()
246 }
247
248 pub fn get_serialization_format(&self) -> SerializationFormat {
250 self.force_serialization_format
251 .unwrap_or(self.serialization_format)
252 }
253
254 pub fn get_compression(&self) -> CompressionAlgorithm {
256 self.force_compression
257 .clone()
258 .unwrap_or_else(|| self.compression.clone())
259 }
260
261 pub fn get_protocols(&self) -> Vec<TransportProtocol> {
263 if let Some(ref protocols) = self.transports {
264 protocols.clone()
265 } else {
266 vec![self.transport]
267 }
268 }
269
270 pub fn is_race_mode(&self) -> bool {
272 self.transports.is_some() && self.transports.as_ref().unwrap().len() > 1
273 }
274}
275
276#[cfg(test)]
277mod tests {
278 use super::*;
279
280 #[test]
281 fn derives_protocol_specific_urls_from_default_websocket_url() {
282 let config = ClientConfig::default();
283
284 assert_eq!(
285 config.get_protocol_url(&TransportProtocol::WebSocket),
286 "ws://localhost:8080"
287 );
288 assert_eq!(
289 config.get_protocol_url(&TransportProtocol::QUIC),
290 "quic://localhost:8080"
291 );
292 assert_eq!(
293 config.get_protocol_url(&TransportProtocol::TCP),
294 "tcp://localhost:8080"
295 );
296 }
297
298 #[test]
299 fn normalizes_explicit_bare_protocol_url() {
300 let config = ClientConfig::default()
301 .with_protocol_url(TransportProtocol::TCP, "127.0.0.1:19090".to_string());
302
303 assert_eq!(
304 config.get_protocol_url(&TransportProtocol::TCP),
305 "tcp://127.0.0.1:19090"
306 );
307 }
308}