use crate::common::compression::CompressionAlgorithm;
use crate::common::config_types::{HeartbeatConfig, TlsConfig, TransportProtocol};
use crate::common::device::DeviceConflictStrategy;
use crate::common::encryption::EncryptionAlgorithm;
use crate::common::protocol::SerializationFormat;
use std::time::Duration;
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
pub struct ServerConfig {
pub bind_address: String,
pub transport: TransportProtocol,
pub transports: Option<Vec<TransportProtocol>>,
pub protocol_addresses: Option<std::collections::HashMap<TransportProtocol, String>>,
pub default_serialization_format: SerializationFormat,
pub default_compression: CompressionAlgorithm,
pub default_encryption: EncryptionAlgorithm,
pub max_connections: usize,
pub handshake_timeout: Duration,
pub max_handshake_concurrency: usize,
pub write_timeout: Duration,
pub fanout_concurrency: usize,
pub connection_timeout: Duration,
pub default_heartbeat: HeartbeatConfig,
pub tls: TlsConfig,
pub max_message_size: usize,
pub device_conflict_strategy: DeviceConflictStrategy,
pub auth_enabled: bool,
pub auth_timeout: Duration,
}
impl Default for ServerConfig {
fn default() -> Self {
Self {
bind_address: "0.0.0.0:8080".to_string(),
transport: TransportProtocol::WebSocket,
transports: None,
protocol_addresses: None,
default_serialization_format: SerializationFormat::Protobuf,
default_compression: CompressionAlgorithm::None,
default_encryption: EncryptionAlgorithm::None,
max_connections: 10000,
handshake_timeout: Duration::from_secs(10),
max_handshake_concurrency: 1024,
write_timeout: Duration::from_secs(10),
fanout_concurrency: 256,
connection_timeout: Duration::from_secs(300),
default_heartbeat: HeartbeatConfig::default(),
tls: TlsConfig::none(),
max_message_size: 10 * 1024 * 1024, device_conflict_strategy: DeviceConflictStrategy::default(),
auth_enabled: false, auth_timeout: Duration::from_secs(30), }
}
}
impl ServerConfig {
pub fn new(bind_address: String) -> Self {
Self {
bind_address,
..Default::default()
}
}
pub fn websocket(mut self) -> Self {
self.transport = TransportProtocol::WebSocket;
self
}
pub fn quic(mut self) -> Self {
self.transport = TransportProtocol::QUIC;
self
}
pub fn tcp(mut self) -> Self {
self.transport = TransportProtocol::TCP;
self
}
pub fn with_format(mut self, format: SerializationFormat) -> Self {
self.default_serialization_format = format;
self
}
pub fn with_compression(mut self, compression: CompressionAlgorithm) -> Self {
self.default_compression = compression;
self
}
pub fn with_encryption(mut self, encryption: EncryptionAlgorithm) -> Self {
self.default_encryption = encryption;
self
}
pub fn with_max_connections(mut self, max: usize) -> Self {
self.max_connections = max;
self
}
pub fn with_handshake_timeout(mut self, timeout: Duration) -> Self {
self.handshake_timeout = timeout;
self
}
pub fn with_max_handshake_concurrency(mut self, max: usize) -> Self {
self.max_handshake_concurrency = max.max(1);
self
}
pub fn with_write_timeout(mut self, timeout: Duration) -> Self {
self.write_timeout = timeout;
self
}
pub fn with_fanout_concurrency(mut self, max: usize) -> Self {
self.fanout_concurrency = max.max(1);
self
}
pub fn with_protocols(mut self, protocols: Vec<TransportProtocol>) -> Self {
self.transports = Some(protocols);
self
}
pub fn with_protocol_address(mut self, protocol: TransportProtocol, address: String) -> Self {
if self.protocol_addresses.is_none() {
self.protocol_addresses = Some(std::collections::HashMap::new());
}
if let Some(ref mut addresses) = self.protocol_addresses {
addresses.insert(protocol, address);
}
self
}
pub fn with_protocol_addresses(
mut self,
addresses: std::collections::HashMap<TransportProtocol, String>,
) -> Self {
self.protocol_addresses = Some(addresses);
self
}
pub fn get_protocol_address(&self, protocol: &TransportProtocol) -> String {
let endpoint = if let Some(ref addresses) = self.protocol_addresses
&& let Some(addr) = addresses.get(protocol)
{
addr.as_str()
} else {
self.bind_address.as_str()
};
TransportProtocol::normalize_server_bind_address(endpoint)
}
pub fn with_heartbeat(mut self, heartbeat: HeartbeatConfig) -> Self {
self.default_heartbeat = heartbeat;
self
}
pub fn with_tls(mut self, tls: TlsConfig) -> Self {
self.tls = tls;
self
}
pub fn with_connection_timeout(mut self, timeout: Duration) -> Self {
self.connection_timeout = timeout;
self
}
pub fn get_protocols(&self) -> Vec<TransportProtocol> {
if let Some(ref protocols) = self.transports {
protocols.clone()
} else {
vec![self.transport]
}
}
pub fn with_device_conflict_strategy(mut self, strategy: DeviceConflictStrategy) -> Self {
self.device_conflict_strategy = strategy;
self
}
pub fn enable_auth(mut self) -> Self {
self.auth_enabled = true;
self
}
pub fn disable_auth(mut self) -> Self {
self.auth_enabled = false;
self
}
pub fn with_auth_timeout(mut self, timeout: Duration) -> Self {
self.auth_timeout = timeout;
self
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn default_connection_admission_limits_are_enabled() {
let config = ServerConfig::default();
assert_eq!(config.handshake_timeout, Duration::from_secs(10));
assert_eq!(config.max_handshake_concurrency, 1024);
assert_eq!(config.write_timeout, Duration::from_secs(10));
assert_eq!(config.fanout_concurrency, 256);
}
#[test]
fn connection_admission_limits_are_configurable() {
let config = ServerConfig::default()
.with_handshake_timeout(Duration::from_secs(3))
.with_max_handshake_concurrency(32)
.with_write_timeout(Duration::from_secs(2))
.with_fanout_concurrency(8);
assert_eq!(config.handshake_timeout, Duration::from_secs(3));
assert_eq!(config.max_handshake_concurrency, 32);
assert_eq!(config.write_timeout, Duration::from_secs(2));
assert_eq!(config.fanout_concurrency, 8);
}
#[test]
fn derives_bind_addresses_without_client_schemes() {
let config = ServerConfig::new("ws://0.0.0.0:8080".to_string());
assert_eq!(
config.get_protocol_address(&TransportProtocol::WebSocket),
"0.0.0.0:8080"
);
assert_eq!(
config.get_protocol_address(&TransportProtocol::QUIC),
"0.0.0.0:8080"
);
assert_eq!(
config.get_protocol_address(&TransportProtocol::TCP),
"0.0.0.0:8080"
);
}
#[test]
fn normalizes_explicit_protocol_bind_address() {
let config = ServerConfig::default()
.with_protocol_address(TransportProtocol::TCP, "tcp://127.0.0.1:19090".to_string());
assert_eq!(
config.get_protocol_address(&TransportProtocol::TCP),
"127.0.0.1:19090"
);
}
}