flare_core/server/
config.rs1use crate::common::compression::CompressionAlgorithm;
4use crate::common::config_types::{HeartbeatConfig, TlsConfig, TransportProtocol};
5use crate::common::device::DeviceConflictStrategy;
6use crate::common::encryption::EncryptionAlgorithm;
7use crate::common::protocol::SerializationFormat;
8use std::time::Duration;
9
10#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
12pub struct ServerConfig {
13 pub bind_address: String,
15 pub transport: TransportProtocol,
17 pub transports: Option<Vec<TransportProtocol>>,
19 pub protocol_addresses: Option<std::collections::HashMap<TransportProtocol, String>>,
22 pub default_serialization_format: SerializationFormat,
24 pub default_compression: CompressionAlgorithm,
26 pub default_encryption: EncryptionAlgorithm,
28 pub max_connections: usize,
30 pub handshake_timeout: Duration,
32 pub max_handshake_concurrency: usize,
34 pub write_timeout: Duration,
36 pub fanout_concurrency: usize,
38 pub connection_timeout: Duration,
40 pub default_heartbeat: HeartbeatConfig,
42 pub tls: TlsConfig,
44 pub max_message_size: usize,
46 pub device_conflict_strategy: DeviceConflictStrategy,
48 pub auth_enabled: bool,
50 pub auth_timeout: Duration,
52}
53
54impl Default for ServerConfig {
55 fn default() -> Self {
56 Self {
57 bind_address: "0.0.0.0:8080".to_string(),
58 transport: TransportProtocol::WebSocket,
59 transports: None,
60 protocol_addresses: None,
61 default_serialization_format: SerializationFormat::Protobuf,
62 default_compression: CompressionAlgorithm::None,
63 default_encryption: EncryptionAlgorithm::None,
64 max_connections: 10000,
65 handshake_timeout: Duration::from_secs(10),
66 max_handshake_concurrency: 1024,
67 write_timeout: Duration::from_secs(10),
68 fanout_concurrency: 256,
69 connection_timeout: Duration::from_secs(300),
70 default_heartbeat: HeartbeatConfig::default(),
71 tls: TlsConfig::none(),
72 max_message_size: 10 * 1024 * 1024, device_conflict_strategy: DeviceConflictStrategy::default(),
74 auth_enabled: false, auth_timeout: Duration::from_secs(30), }
77 }
78}
79
80impl ServerConfig {
81 pub fn new(bind_address: String) -> Self {
83 Self {
84 bind_address,
85 ..Default::default()
86 }
87 }
88
89 pub fn websocket(mut self) -> Self {
91 self.transport = TransportProtocol::WebSocket;
92 self
93 }
94
95 pub fn quic(mut self) -> Self {
97 self.transport = TransportProtocol::QUIC;
98 self
99 }
100
101 pub fn tcp(mut self) -> Self {
103 self.transport = TransportProtocol::TCP;
104 self
105 }
106
107 pub fn with_format(mut self, format: SerializationFormat) -> Self {
109 self.default_serialization_format = format;
110 self
111 }
112
113 pub fn with_compression(mut self, compression: CompressionAlgorithm) -> Self {
115 self.default_compression = compression;
116 self
117 }
118
119 pub fn with_encryption(mut self, encryption: EncryptionAlgorithm) -> Self {
121 self.default_encryption = encryption;
122 self
123 }
124
125 pub fn with_max_connections(mut self, max: usize) -> Self {
127 self.max_connections = max;
128 self
129 }
130
131 pub fn with_handshake_timeout(mut self, timeout: Duration) -> Self {
133 self.handshake_timeout = timeout;
134 self
135 }
136
137 pub fn with_max_handshake_concurrency(mut self, max: usize) -> Self {
139 self.max_handshake_concurrency = max.max(1);
140 self
141 }
142
143 pub fn with_write_timeout(mut self, timeout: Duration) -> Self {
145 self.write_timeout = timeout;
146 self
147 }
148
149 pub fn with_fanout_concurrency(mut self, max: usize) -> Self {
151 self.fanout_concurrency = max.max(1);
152 self
153 }
154
155 pub fn with_protocols(mut self, protocols: Vec<TransportProtocol>) -> Self {
157 self.transports = Some(protocols);
158 self
159 }
160
161 pub fn with_protocol_address(mut self, protocol: TransportProtocol, address: String) -> Self {
163 if self.protocol_addresses.is_none() {
164 self.protocol_addresses = Some(std::collections::HashMap::new());
165 }
166 if let Some(ref mut addresses) = self.protocol_addresses {
167 addresses.insert(protocol, address);
168 }
169 self
170 }
171
172 pub fn with_protocol_addresses(
174 mut self,
175 addresses: std::collections::HashMap<TransportProtocol, String>,
176 ) -> Self {
177 self.protocol_addresses = Some(addresses);
178 self
179 }
180
181 pub fn get_protocol_address(&self, protocol: &TransportProtocol) -> String {
183 let endpoint = if let Some(ref addresses) = self.protocol_addresses
184 && let Some(addr) = addresses.get(protocol)
185 {
186 addr.as_str()
187 } else {
188 self.bind_address.as_str()
189 };
190
191 TransportProtocol::normalize_server_bind_address(endpoint)
192 }
193
194 pub fn with_heartbeat(mut self, heartbeat: HeartbeatConfig) -> Self {
196 self.default_heartbeat = heartbeat;
197 self
198 }
199
200 pub fn with_tls(mut self, tls: TlsConfig) -> Self {
202 self.tls = tls;
203 self
204 }
205
206 pub fn with_connection_timeout(mut self, timeout: Duration) -> Self {
208 self.connection_timeout = timeout;
209 self
210 }
211
212 pub fn get_protocols(&self) -> Vec<TransportProtocol> {
214 if let Some(ref protocols) = self.transports {
215 protocols.clone()
216 } else {
217 vec![self.transport]
218 }
219 }
220
221 pub fn with_device_conflict_strategy(mut self, strategy: DeviceConflictStrategy) -> Self {
223 self.device_conflict_strategy = strategy;
224 self
225 }
226
227 pub fn enable_auth(mut self) -> Self {
231 self.auth_enabled = true;
232 self
233 }
234
235 pub fn disable_auth(mut self) -> Self {
237 self.auth_enabled = false;
238 self
239 }
240
241 pub fn with_auth_timeout(mut self, timeout: Duration) -> Self {
246 self.auth_timeout = timeout;
247 self
248 }
249}
250
251#[cfg(test)]
252mod tests {
253 use super::*;
254
255 #[test]
256 fn default_connection_admission_limits_are_enabled() {
257 let config = ServerConfig::default();
258
259 assert_eq!(config.handshake_timeout, Duration::from_secs(10));
260 assert_eq!(config.max_handshake_concurrency, 1024);
261 assert_eq!(config.write_timeout, Duration::from_secs(10));
262 assert_eq!(config.fanout_concurrency, 256);
263 }
264
265 #[test]
266 fn connection_admission_limits_are_configurable() {
267 let config = ServerConfig::default()
268 .with_handshake_timeout(Duration::from_secs(3))
269 .with_max_handshake_concurrency(32)
270 .with_write_timeout(Duration::from_secs(2))
271 .with_fanout_concurrency(8);
272
273 assert_eq!(config.handshake_timeout, Duration::from_secs(3));
274 assert_eq!(config.max_handshake_concurrency, 32);
275 assert_eq!(config.write_timeout, Duration::from_secs(2));
276 assert_eq!(config.fanout_concurrency, 8);
277 }
278
279 #[test]
280 fn derives_bind_addresses_without_client_schemes() {
281 let config = ServerConfig::new("ws://0.0.0.0:8080".to_string());
282
283 assert_eq!(
284 config.get_protocol_address(&TransportProtocol::WebSocket),
285 "0.0.0.0:8080"
286 );
287 assert_eq!(
288 config.get_protocol_address(&TransportProtocol::QUIC),
289 "0.0.0.0:8080"
290 );
291 assert_eq!(
292 config.get_protocol_address(&TransportProtocol::TCP),
293 "0.0.0.0:8080"
294 );
295 }
296
297 #[test]
298 fn normalizes_explicit_protocol_bind_address() {
299 let config = ServerConfig::default()
300 .with_protocol_address(TransportProtocol::TCP, "tcp://127.0.0.1:19090".to_string());
301
302 assert_eq!(
303 config.get_protocol_address(&TransportProtocol::TCP),
304 "127.0.0.1:19090"
305 );
306 }
307}