Skip to main content

flare_core/client/
config.rs

1//! 客户端配置模块
2
3use 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/// 客户端配置
10#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
11pub struct ClientConfig {
12    /// 服务器地址(单个协议时使用,多协议时用作默认地址)
13    pub server_url: String,
14    /// 传输协议(单个)
15    pub transport: TransportProtocol,
16    /// 传输协议列表(用于竞速,如果设置了此列表,transport 将被忽略)
17    /// 列表顺序就是优先级顺序,前面的优先级更高
18    pub transports: Option<Vec<TransportProtocol>>,
19    /// 每个协议的独立地址配置(协议 -> 地址映射)
20    /// 如果设置了此映射,每个协议将使用对应的地址
21    pub protocol_urls: Option<std::collections::HashMap<TransportProtocol, String>>,
22    /// 协议竞速超时时间(如果启用多协议竞速)
23    pub race_timeout: Option<Duration>,
24    /// 序列化格式(首选格式,用于协商)
25    pub serialization_format: SerializationFormat,
26    /// 压缩算法(首选算法,用于协商)
27    pub compression: CompressionAlgorithm,
28    /// 强制指定的序列化格式(如果设置,则强制使用此格式,不进行协商)
29    /// 适用于某些端不支持 protobuf 等格式的场景
30    pub force_serialization_format: Option<SerializationFormat>,
31    /// 强制指定的压缩算法(如果设置,则强制使用此算法,不进行协商)
32    pub force_compression: Option<CompressionAlgorithm>,
33    /// 连接超时时间
34    pub connect_timeout: Duration,
35    /// 重连间隔
36    pub reconnect_interval: Duration,
37    /// 最大重连次数(None 表示无限重连)
38    pub max_reconnect_attempts: Option<u32>,
39    /// 心跳配置
40    pub heartbeat: HeartbeatConfig,
41    /// TLS 配置
42    pub tls: TlsConfig,
43    /// 自定义连接 ID(如果为 None,则自动生成)
44    pub connection_id: Option<String>,
45    /// 用户 ID(用于认证)
46    pub user_id: Option<String>,
47    /// 额外的元数据
48    pub metadata: std::collections::HashMap<String, String>,
49    /// 是否启用消息路由(默认 false)
50    pub enable_router: bool,
51    /// 设备信息(用于协商和设备管理)
52    pub device_info: Option<DeviceInfo>,
53    /// Token(用于认证,如果服务端启用认证,必须提供)
54    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            // 默认使用 JSON 序列化,不压缩(客户端可以在协商时指定首选格式)
66            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    /// 创建新的客户端配置
87    pub fn new(server_url: String) -> Self {
88        Self {
89            server_url,
90            ..Default::default()
91        }
92    }
93
94    /// 使用 WebSocket 协议
95    pub fn websocket(mut self) -> Self {
96        self.transport = TransportProtocol::WebSocket;
97        self
98    }
99
100    /// 使用 QUIC 协议
101    pub fn quic(mut self) -> Self {
102        self.transport = TransportProtocol::QUIC;
103        self
104    }
105
106    /// 使用 TCP 协议
107    pub fn tcp(mut self) -> Self {
108        self.transport = TransportProtocol::TCP;
109        self
110    }
111
112    /// 设置序列化格式
113    pub fn with_format(mut self, format: SerializationFormat) -> Self {
114        self.serialization_format = format;
115        self
116    }
117
118    /// 设置压缩算法
119    pub fn with_compression(mut self, compression: CompressionAlgorithm) -> Self {
120        self.compression = compression;
121        self
122    }
123
124    /// 设置用户 ID
125    pub fn with_user_id(mut self, user_id: String) -> Self {
126        self.user_id = Some(user_id);
127        self
128    }
129
130    /// 设置 Token(用于认证,如果服务端启用认证,必须提供)
131    pub fn with_token(mut self, token: String) -> Self {
132        self.token = Some(token);
133        self
134    }
135
136    /// 启用多协议竞速
137    ///
138    /// 协议列表的顺序就是优先级顺序,前面的协议优先级更高
139    /// 例如:with_protocol_race(vec![QUIC, WebSocket]) 表示 QUIC 优先级高于 WebSocket
140    pub fn with_protocol_race(mut self, protocols: Vec<TransportProtocol>) -> Self {
141        self.transports = Some(protocols);
142        self
143    }
144
145    /// 为特定协议设置服务器地址
146    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    /// 批量设置协议地址映射
157    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    /// 获取指定协议的地址
166    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    /// 设置竞速超时时间
179    pub fn with_race_timeout(mut self, timeout: Duration) -> Self {
180        self.race_timeout = Some(timeout);
181        self
182    }
183
184    /// 设置心跳配置
185    pub fn with_heartbeat(mut self, heartbeat: HeartbeatConfig) -> Self {
186        self.heartbeat = heartbeat;
187        self
188    }
189
190    /// 设置 TLS 配置
191    pub fn with_tls(mut self, tls: TlsConfig) -> Self {
192        self.tls = tls;
193        self
194    }
195
196    /// 设置连接超时
197    pub fn with_connect_timeout(mut self, timeout: Duration) -> Self {
198        self.connect_timeout = timeout;
199        self
200    }
201
202    /// 设置重连间隔
203    pub fn with_reconnect_interval(mut self, interval: Duration) -> Self {
204        self.reconnect_interval = interval;
205        self
206    }
207
208    /// 设置最大重连次数(None 表示无限重连)
209    pub fn with_max_reconnect_attempts(mut self, max: Option<u32>) -> Self {
210        self.max_reconnect_attempts = max;
211        self
212    }
213
214    /// 启用消息路由
215    pub fn enable_router(mut self) -> Self {
216        self.enable_router = true;
217        self
218    }
219
220    /// 设置设备信息(用于协商和设备管理)
221    pub fn with_device_info(mut self, device_info: DeviceInfo) -> Self {
222        self.device_info = Some(device_info);
223        self
224    }
225
226    /// 强制指定序列化格式(不进行协商,直接使用此格式)
227    ///
228    /// 适用于某些端不支持 protobuf 等格式的场景
229    /// 如果设置了强制格式,客户端将直接使用此格式,服务端必须接受
230    pub fn force_format(mut self, format: SerializationFormat) -> Self {
231        self.force_serialization_format = Some(format);
232        self
233    }
234
235    /// 强制指定压缩算法(不进行协商,直接使用此算法)
236    ///
237    /// 如果设置了强制压缩,客户端将直接使用此算法,服务端必须接受
238    pub fn force_compression(mut self, compression: CompressionAlgorithm) -> Self {
239        self.force_compression = Some(compression);
240        self
241    }
242
243    /// 检查是否强制指定了格式(不进行协商)
244    pub fn is_force_format(&self) -> bool {
245        self.force_serialization_format.is_some() || self.force_compression.is_some()
246    }
247
248    /// 获取实际使用的序列化格式(强制格式优先,否则使用首选格式)
249    pub fn get_serialization_format(&self) -> SerializationFormat {
250        self.force_serialization_format
251            .unwrap_or(self.serialization_format)
252    }
253
254    /// 获取实际使用的压缩算法(强制算法优先,否则使用首选算法)
255    pub fn get_compression(&self) -> CompressionAlgorithm {
256        self.force_compression
257            .clone()
258            .unwrap_or_else(|| self.compression.clone())
259    }
260
261    /// 获取要使用的协议列表
262    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    /// 是否启用协议竞速
271    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}