Skip to main content

flare_core/server/connection/
negotiation.rs

1//! 连接协商模块
2//!
3//! 处理连接建立时的序列化格式、压缩算法和设备信息协商
4
5use crate::common::compression::CompressionAlgorithm;
6use crate::common::device::{DeviceInfo, DevicePlatform};
7use crate::common::encryption::EncryptionAlgorithm;
8use crate::common::error::Result;
9use crate::common::protocol::flare::core::commands::command::Type as CommandType;
10use crate::common::protocol::flare::core::commands::system_command::SerializationFormat;
11use crate::common::protocol::flare::core::commands::system_command::Type as SystemType;
12use crate::common::protocol::{Frame, SystemCommand};
13use std::collections::HashMap;
14
15const METADATA_COMPRESSION: &str = "compression";
16const METADATA_FORMAT: &str = "format";
17const METADATA_ENCRYPTION: &str = "encryption";
18const METADATA_DEVICE_ID: &str = "device_id";
19const METADATA_PLATFORM: &str = "platform";
20const METADATA_MODEL: &str = "model";
21const METADATA_APP_VERSION: &str = "app_version";
22const METADATA_SYSTEM_VERSION: &str = "system_version";
23const METADATA_USER_ID: &str = "user_id";
24const METADATA_FORCE_FORMAT: &str = "force_format";
25
26const DEVICE_RESERVED_METADATA_KEYS: &[&str] = &[
27    METADATA_COMPRESSION,
28    METADATA_FORMAT,
29    METADATA_ENCRYPTION,
30    METADATA_DEVICE_ID,
31    METADATA_PLATFORM,
32    METADATA_MODEL,
33    METADATA_APP_VERSION,
34    METADATA_SYSTEM_VERSION,
35    METADATA_USER_ID,
36    METADATA_FORCE_FORMAT,
37];
38
39/// 连接协商结果
40#[derive(Debug, Clone)]
41pub struct NegotiationResult {
42    /// 序列化格式(客户端请求的格式)
43    pub serialization_format: SerializationFormat,
44    /// 客户端是否显式请求了序列化格式
45    pub serialization_format_specified: bool,
46    /// 压缩算法(客户端请求的压缩方式)
47    pub compression: CompressionAlgorithm,
48    /// 加密方式
49    pub encryption: EncryptionAlgorithm,
50    /// 是否强制指定格式(客户端强制模式,服务端必须使用客户端指定的格式)
51    pub is_forced: bool,
52    /// 设备信息(如果客户端提供)
53    pub device_info: Option<DeviceInfo>,
54    /// 用户 ID(如果客户端在 CONNECT 中提供)
55    pub user_id: Option<String>,
56}
57
58impl Default for NegotiationResult {
59    fn default() -> Self {
60        Self {
61            serialization_format: SerializationFormat::Json,
62            serialization_format_specified: false,
63            compression: CompressionAlgorithm::None,
64            encryption: EncryptionAlgorithm::None,
65            is_forced: false,
66            device_info: None,
67            user_id: None,
68        }
69    }
70}
71
72/// 解析 CONNECT 消息,提取客户端协商信息
73///
74/// # 参数
75/// - `frame`: CONNECT 消息的 Frame
76///
77/// # 返回
78/// 协商结果,包含序列化格式、压缩算法、设备信息等
79pub fn parse_connect_message(frame: &Frame) -> Result<NegotiationResult> {
80    let Some(sys_cmd) = connect_system_command(frame) else {
81        return Ok(NegotiationResult::default());
82    };
83
84    if sys_cmd.r#type != SystemType::Connect as i32 {
85        return Ok(NegotiationResult::default());
86    }
87
88    let metadata = &sys_cmd.metadata;
89    let (serialization_format, serialization_format_specified) =
90        parse_serialization_format(metadata);
91    let result = NegotiationResult {
92        serialization_format,
93        serialization_format_specified,
94        compression: parse_compression(metadata),
95        encryption: parse_encryption(metadata),
96        is_forced: metadata_string(metadata, METADATA_FORCE_FORMAT).is_some_and(|v| v == "true"),
97        device_info: parse_device_info(metadata),
98        user_id: metadata_string(metadata, METADATA_USER_ID),
99    };
100
101    Ok(result)
102}
103
104fn connect_system_command(frame: &Frame) -> Option<&SystemCommand> {
105    match frame.command.as_ref()?.r#type.as_ref()? {
106        CommandType::System(sys_cmd) => Some(sys_cmd),
107        _ => None,
108    }
109}
110
111fn parse_serialization_format(metadata: &HashMap<String, Vec<u8>>) -> (SerializationFormat, bool) {
112    let Some(format) = metadata_string(metadata, METADATA_FORMAT) else {
113        return (SerializationFormat::Json, false);
114    };
115
116    let format = match format.as_str() {
117        value if value.eq_ignore_ascii_case("protobuf") || value.eq_ignore_ascii_case("proto") => {
118            SerializationFormat::Protobuf
119        }
120        value if value.eq_ignore_ascii_case("json") => SerializationFormat::Json,
121        _ => SerializationFormat::Json,
122    };
123
124    (format, true)
125}
126
127fn parse_compression(metadata: &HashMap<String, Vec<u8>>) -> CompressionAlgorithm {
128    metadata_string(metadata, METADATA_COMPRESSION)
129        .and_then(|value| CompressionAlgorithm::from_str(&value))
130        .unwrap_or(CompressionAlgorithm::None)
131}
132
133fn parse_encryption(metadata: &HashMap<String, Vec<u8>>) -> EncryptionAlgorithm {
134    metadata_string(metadata, METADATA_ENCRYPTION)
135        .and_then(|value| EncryptionAlgorithm::from_str(&value))
136        .unwrap_or(EncryptionAlgorithm::None)
137}
138
139fn parse_device_info(metadata: &HashMap<String, Vec<u8>>) -> Option<DeviceInfo> {
140    let device_id = metadata_string(metadata, METADATA_DEVICE_ID)?;
141    let platform = metadata_string(metadata, METADATA_PLATFORM)
142        .map(|value| DevicePlatform::from_str(&value))
143        .unwrap_or_else(|| DevicePlatform::Other("unknown".to_string()));
144
145    let mut device_info = DeviceInfo::new(device_id, platform);
146    if let Some(model) = metadata_string(metadata, METADATA_MODEL) {
147        device_info = device_info.with_model(model);
148    }
149    if let Some(app_version) = metadata_string(metadata, METADATA_APP_VERSION) {
150        device_info = device_info.with_app_version(app_version);
151    }
152    if let Some(system_version) = metadata_string(metadata, METADATA_SYSTEM_VERSION) {
153        device_info = device_info.with_system_version(system_version);
154    }
155
156    for (key, value) in metadata
157        .iter()
158        .filter(|(key, _)| !is_reserved_device_metadata_key(key))
159    {
160        if let Some(value) = bytes_to_string(value) {
161            device_info = device_info.with_metadata(key.clone(), value);
162        }
163    }
164
165    Some(device_info)
166}
167
168fn metadata_string(metadata: &HashMap<String, Vec<u8>>, key: &str) -> Option<String> {
169    metadata.get(key).and_then(|value| bytes_to_string(value))
170}
171
172fn bytes_to_string(value: &[u8]) -> Option<String> {
173    String::from_utf8(value.to_vec()).ok()
174}
175
176fn is_reserved_device_metadata_key(key: &str) -> bool {
177    DEVICE_RESERVED_METADATA_KEYS.contains(&key)
178}
179
180/// 创建 CONNECT_ACK 响应
181///
182/// # 参数
183/// - `format`: 确认使用的序列化格式
184/// - `compression`: 确认使用的压缩算法
185/// - `additional_metadata`: 额外的元数据(如设备冲突信息等)
186///
187/// # 返回
188/// CONNECT_ACK 命令
189/// 创建 CONNECT_ACK 消息
190///
191/// # 参数
192/// - `format`: 确认使用的序列化格式
193/// - `compression`: 确认使用的压缩算法
194/// - `encryption`: 确认使用的加密方式(目前为 "none",为未来扩展预留)
195/// - `additional_metadata`: 额外的元数据(如设备冲突信息等)
196///
197/// # 返回
198/// CONNECT_ACK 命令
199pub fn create_connect_ack(
200    format: SerializationFormat,
201    compression: CompressionAlgorithm,
202    encryption: EncryptionAlgorithm,
203    additional_metadata: Option<HashMap<String, Vec<u8>>>,
204) -> SystemCommand {
205    // 验证压缩算法是否已注册
206    let mut compression_str = compression.as_str();
207    if !crate::common::compression::CompressionUtil::is_registered(&compression_str) {
208        tracing::warn!(
209            "[Negotiation] 压缩算法 '{}' 未注册,将使用 'none'",
210            compression_str
211        );
212        // 如果未注册,回退到 none
213        compression_str = "none".to_string();
214    }
215
216    // 验证加密算法是否已注册
217    let mut encryption_str = encryption.as_str();
218    if encryption != EncryptionAlgorithm::None {
219        if !crate::common::encryption::EncryptionUtil::is_registered(&encryption_str) {
220            let registered = crate::common::encryption::EncryptionUtil::list_registered();
221            tracing::warn!(
222                "[Negotiation] 加密算法 '{}' 未注册,将使用 'none'。已注册的加密器: {:?}",
223                encryption_str,
224                registered
225            );
226            // 如果未注册,回退到 none
227            encryption_str = "none".to_string();
228        } else {
229            tracing::debug!(
230                "[Negotiation] 加密算法 '{}' 已注册,可以使用",
231                encryption_str
232            );
233        }
234    }
235
236    let mut metadata = HashMap::new();
237
238    // 添加额外的元数据
239    if let Some(extra) = additional_metadata {
240        for (key, value) in extra {
241            metadata.insert(key, value);
242        }
243    }
244
245    crate::common::protocol::connect_ack(
246        format,
247        Some(compression_str.as_str()),
248        Some(encryption_str.as_str()),
249        metadata,
250    )
251}
252
253#[cfg(test)]
254mod tests {
255    use super::*;
256    use crate::common::protocol::{Reliability, connect, frame_with_system_command, ping};
257
258    #[test]
259    fn parses_connect_metadata_into_negotiation_result() {
260        let metadata = HashMap::from([
261            (METADATA_COMPRESSION.to_string(), b"gzip".to_vec()),
262            (METADATA_FORMAT.to_string(), b"protobuf".to_vec()),
263            (METADATA_ENCRYPTION.to_string(), b"aes256gcm".to_vec()),
264            (METADATA_DEVICE_ID.to_string(), b"device-1".to_vec()),
265            (METADATA_PLATFORM.to_string(), b"ios".to_vec()),
266            (METADATA_MODEL.to_string(), b"iPhone".to_vec()),
267            (METADATA_APP_VERSION.to_string(), b"1.2.3".to_vec()),
268            (METADATA_SYSTEM_VERSION.to_string(), b"18.0".to_vec()),
269            (METADATA_USER_ID.to_string(), b"user-1".to_vec()),
270            (METADATA_FORCE_FORMAT.to_string(), b"true".to_vec()),
271            ("trace_id".to_string(), b"trace-1".to_vec()),
272        ]);
273        let frame = frame_with_system_command(
274            connect(SerializationFormat::Json, metadata),
275            Reliability::AtLeastOnce,
276        );
277
278        let result = parse_connect_message(&frame).expect("parse connect");
279
280        assert_eq!(result.serialization_format, SerializationFormat::Protobuf);
281        assert!(result.serialization_format_specified);
282        assert_eq!(result.compression, CompressionAlgorithm::Gzip);
283        assert_eq!(result.encryption, EncryptionAlgorithm::Aes256Gcm);
284        assert!(result.is_forced);
285        assert_eq!(result.user_id.as_deref(), Some("user-1"));
286
287        let device = result.device_info.expect("device info");
288        assert_eq!(device.device_id, "device-1");
289        assert_eq!(device.platform, DevicePlatform::IOS);
290        assert_eq!(device.model.as_deref(), Some("iPhone"));
291        assert_eq!(device.app_version.as_deref(), Some("1.2.3"));
292        assert_eq!(device.system_version.as_deref(), Some("18.0"));
293        assert_eq!(
294            device.metadata.get("trace_id").map(String::as_str),
295            Some("trace-1")
296        );
297    }
298
299    #[test]
300    fn returns_default_negotiation_for_non_connect_system_command() {
301        let frame = frame_with_system_command(ping(), Reliability::BestEffort);
302
303        let result = parse_connect_message(&frame).expect("parse non-connect");
304
305        assert_eq!(result.serialization_format, SerializationFormat::Json);
306        assert!(!result.serialization_format_specified);
307        assert_eq!(result.compression, CompressionAlgorithm::None);
308        assert_eq!(result.encryption, EncryptionAlgorithm::None);
309        assert!(!result.is_forced);
310        assert!(result.device_info.is_none());
311        assert!(result.user_id.is_none());
312    }
313
314    #[test]
315    fn parses_explicit_json_format_from_metadata() {
316        let metadata = HashMap::from([(METADATA_FORMAT.to_string(), b"json".to_vec())]);
317        let frame = frame_with_system_command(
318            connect(SerializationFormat::Json, metadata),
319            Reliability::AtLeastOnce,
320        );
321
322        let result = parse_connect_message(&frame).expect("parse connect");
323
324        assert_eq!(result.serialization_format, SerializationFormat::Json);
325        assert!(result.serialization_format_specified);
326    }
327
328    #[test]
329    fn treats_missing_format_metadata_as_unspecified() {
330        let frame = frame_with_system_command(
331            connect(SerializationFormat::Json, HashMap::new()),
332            Reliability::AtLeastOnce,
333        );
334
335        let result = parse_connect_message(&frame).expect("parse connect");
336
337        assert_eq!(result.serialization_format, SerializationFormat::Json);
338        assert!(!result.serialization_format_specified);
339    }
340}