flare_core/server/connection/
negotiation.rs1use 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#[derive(Debug, Clone)]
41pub struct NegotiationResult {
42 pub serialization_format: SerializationFormat,
44 pub serialization_format_specified: bool,
46 pub compression: CompressionAlgorithm,
48 pub encryption: EncryptionAlgorithm,
50 pub is_forced: bool,
52 pub device_info: Option<DeviceInfo>,
54 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
72pub 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
180pub fn create_connect_ack(
200 format: SerializationFormat,
201 compression: CompressionAlgorithm,
202 encryption: EncryptionAlgorithm,
203 additional_metadata: Option<HashMap<String, Vec<u8>>>,
204) -> SystemCommand {
205 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 compression_str = "none".to_string();
214 }
215
216 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 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 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}