Skip to main content

flare_core/server/connection/manager/
trait_impl.rs

1use super::*;
2
3#[async_trait]
4impl ConnectionManagerTrait for ConnectionManager {
5    fn as_any(&self) -> &dyn std::any::Any {
6        self
7    }
8
9    async fn add_connection(
10        &self,
11        connection_id: String,
12        connection: Arc<Mutex<Box<dyn Connection>>>,
13        user_id: Option<String>,
14    ) -> Result<()> {
15        // 注意:trait 方法不能直接传递 requires_auth,我们需要从 ServerCore 获取
16        // 但这里我们暂时使用 true(需要认证),实际值应该在调用时通过 ServerCore 的 auth_enabled() 获取
17        // 由于 ConnectionManager 不知道 ServerCore,我们暂时使用 true
18        // 实际应用中,连接会在 CONNECT 消息处理时被标记为已验证
19        let requires_auth = true; // 默认需要认证,如果不需要认证,连接会在 CONNECT 消息处理时被标记为已验证
20
21        // 将 Arc<Mutex<Box<dyn Connection>>> 转换为 Box<dyn Connection>
22        // 注意:这需要从 Arc 中取出,但 Arc 可能被多个地方引用
23        // 对于默认实现,我们需要一个不同的方式
24        // 由于 ConnectionManager 内部使用 Arc<Mutex<Box<dyn Connection>>>,
25        // 我们需要保持一致性
26        self.reserve_connection_slot(usize::MAX)?;
27
28        let mut shard = match self.connection_shard(&connection_id).write() {
29            Ok(shard) => shard,
30            Err(_) => {
31                self.release_connection_slot();
32                return Err(FlareError::general_error("Failed to lock connection shard"));
33            }
34        };
35
36        if shard.contains_key(&connection_id) {
37            self.release_connection_slot();
38            return Err(FlareError::protocol_error(format!(
39                "Connection {} already exists",
40                connection_id
41            )));
42        }
43
44        let mut info = ConnectionInfo::new(connection_id.clone(), requires_auth);
45        info.user_id = user_id.clone();
46
47        let entry = self.new_connection_entry(&connection_id, Arc::clone(&connection), info);
48        shard.insert(connection_id.clone(), entry);
49
50        // 如果提供了用户 ID,添加到用户连接映射
51        if let Some(user_id) = user_id
52            && let Err(err) = self.insert_user_connection(user_id, &connection_id)
53        {
54            shard.remove(&connection_id);
55            self.release_connection_slot();
56            return Err(err);
57        }
58
59        Ok(())
60    }
61
62    async fn remove_connection(&self, connection_id: &str) -> Result<()> {
63        ConnectionManager::remove_connection(self, connection_id)
64    }
65
66    async fn get_connection(
67        &self,
68        connection_id: &str,
69    ) -> Option<(
70        Arc<Mutex<Box<dyn Connection>>>,
71        crate::server::connection::r#trait::ConnectionInfo,
72    )> {
73        ConnectionManager::get_connection(self, connection_id).map(|(conn, info)| {
74            // 转换 ConnectionInfo 格式(从 Instant 转换为 Unix 时间戳)
75            let now = std::time::SystemTime::now()
76                .duration_since(std::time::UNIX_EPOCH)
77                .unwrap_or_default()
78                .as_secs();
79            let created_at_secs = now.saturating_sub(info.created_at.elapsed().as_secs());
80            let last_active_secs = now.saturating_sub(info.last_active.elapsed().as_secs());
81
82            let trait_info = crate::server::connection::r#trait::ConnectionInfo {
83                connection_id: info.connection_id,
84                user_id: info.user_id,
85                created_at: created_at_secs,
86                last_active: last_active_secs,
87                metadata: info.metadata,
88                device_info: info.device_info.clone(),
89                serialization_format: info.serialization_format,
90                compression: info.compression,
91                encryption: info.encryption,
92                authenticated: info.authenticated,
93                authenticated_at: info.authenticated_at,
94                negotiation_completed: info.negotiation_completed,
95                negotiation_confirmed: info.negotiation_confirmed,
96                cached_parser: info.cached_parser.clone(),
97                cached_pipeline: info.cached_pipeline.clone(),
98            };
99            (conn, trait_info)
100        })
101    }
102
103    async fn get_user_connections(&self, user_id: &str) -> Vec<String> {
104        ConnectionManager::get_user_connections(self, user_id)
105    }
106
107    async fn bind_user(&self, connection_id: &str, user_id: String) -> Result<()> {
108        ConnectionManager::bind_user(self, connection_id, user_id)
109    }
110
111    async fn update_connection_active(&self, connection_id: &str) -> Result<()> {
112        ConnectionManager::update_connection_active(self, connection_id)
113    }
114
115    async fn set_connection_authenticated(
116        &self,
117        connection_id: &str,
118        user_id: Option<String>,
119    ) -> Result<()> {
120        // ConnectionManager::set_connection_authenticated 是同步方法,直接调用
121        ConnectionManager::set_connection_authenticated(self, connection_id, user_id)
122    }
123
124    async fn list_connections(&self) -> Vec<String> {
125        ConnectionManager::list_connections(self)
126    }
127
128    async fn connection_count(&self) -> usize {
129        ConnectionManager::connection_count(self)
130    }
131
132    fn connection_count_snapshot(&self) -> usize {
133        ConnectionManager::connection_count(self)
134    }
135
136    fn user_count_snapshot(&self) -> usize {
137        ConnectionManager::user_count(self)
138    }
139
140    async fn cleanup_timeout_connections(&self, timeout: Duration) -> Vec<String> {
141        let timeout_connections = self.timeout_connection_snapshots(timeout);
142
143        for (_, connection, _) in &timeout_connections {
144            let mut conn = connection.lock().await;
145            let _ = conn.close().await;
146        }
147
148        self.remove_connection_snapshots(
149            timeout_connections
150                .iter()
151                .map(|(connection_id, _, _)| connection_id.clone()),
152        )
153    }
154
155    async fn send_to_connection(&self, connection_id: &str, data: &[u8]) -> Result<()> {
156        let (_, connection, _) = self.get_connection_snapshot(connection_id).ok_or_else(|| {
157            FlareError::protocol_error(format!("Connection {} not found", connection_id))
158        })?;
159
160        self.send_to_connection_handle(connection_id, connection, data)
161            .await
162    }
163
164    async fn send_to_user(&self, user_id: &str, data: &[u8]) -> Result<()> {
165        let connections =
166            self.connection_handles_for_ids(ConnectionManager::get_user_connections(self, user_id));
167
168        stream::iter(connections)
169            .for_each_concurrent(
170                self.fanout_concurrency,
171                |(connection_id, connection)| async move {
172                    if let Err(e) = self
173                        .send_to_connection_handle(&connection_id, connection, data)
174                        .await
175                    {
176                        tracing::warn!("Failed to send to connection {}: {:?}", connection_id, e);
177                    }
178                },
179            )
180            .await;
181
182        Ok(())
183    }
184
185    async fn broadcast(&self, data: &[u8]) -> Result<()> {
186        let connections = self.connection_handles();
187
188        stream::iter(connections)
189            .for_each_concurrent(
190                self.fanout_concurrency,
191                |(connection_id, connection)| async move {
192                    if let Err(e) = self
193                        .send_to_connection_handle(&connection_id, connection, data)
194                        .await
195                    {
196                        tracing::warn!(
197                            "Failed to broadcast to connection {}: {:?}",
198                            connection_id,
199                            e
200                        );
201                    }
202                },
203            )
204            .await;
205
206        Ok(())
207    }
208
209    async fn broadcast_except(&self, data: &[u8], exclude_connection_id: &str) -> Result<()> {
210        let connections = self.connection_handles_except(exclude_connection_id);
211
212        stream::iter(connections)
213            .for_each_concurrent(
214                self.fanout_concurrency,
215                |(connection_id, connection)| async move {
216                    if let Err(e) = self
217                        .send_to_connection_handle(&connection_id, connection, data)
218                        .await
219                    {
220                        tracing::warn!(
221                            "Failed to broadcast to connection {}: {:?}",
222                            connection_id,
223                            e
224                        );
225                    }
226                },
227            )
228            .await;
229
230        Ok(())
231    }
232
233    async fn send_frame_to(
234        &self,
235        connection_id: &str,
236        frame: &crate::common::protocol::Frame,
237        parser: Option<&crate::common::MessageParser>,
238    ) -> Result<()> {
239        let snapshot = self.get_connection_snapshot(connection_id).ok_or_else(|| {
240            FlareError::connection_failed(format!("连接 {} 不存在", connection_id))
241        })?;
242
243        self.send_frame_to_snapshot(snapshot, frame, parser).await
244    }
245
246    async fn send_frame_to_user(
247        &self,
248        user_id: &str,
249        frame: &crate::common::protocol::Frame,
250        parser: Option<&crate::common::MessageParser>,
251    ) -> Result<()> {
252        let connection_ids = ConnectionManager::get_user_connections(self, user_id);
253
254        if let Some(parser) = parser {
255            let connections = self.connection_auth_snapshots_for_ids(connection_ids);
256            let data = match parser.serialize(frame) {
257                Ok(data) => data,
258                Err(e) => {
259                    tracing::warn!("Failed to serialize frame for user {}: {:?}", user_id, e);
260                    return Ok(());
261                }
262            };
263
264            let successful_ids = Arc::new(std::sync::Mutex::new(Vec::new()));
265            stream::iter(connections)
266                .for_each_concurrent(self.fanout_concurrency, |snapshot| {
267                    let successful_ids = Arc::clone(&successful_ids);
268                    let data = data.as_slice();
269                    async move {
270                        let connection_id = snapshot.0.clone();
271                        let result = self
272                            .send_serialized_frame_to_auth_snapshot_without_active(
273                                snapshot, frame, data,
274                            )
275                            .await;
276                        match result {
277                            Ok(connection_id) => {
278                                Self::record_successful_connection_id(
279                                    &successful_ids,
280                                    connection_id,
281                                );
282                            }
283                            Err(e) => {
284                                tracing::warn!(
285                                    "Failed to send frame to connection {}: {:?}",
286                                    connection_id,
287                                    e
288                                );
289                            }
290                        }
291                    }
292                })
293                .await;
294            self.update_connections_active(Self::take_successful_connection_ids(successful_ids));
295
296            return Ok(());
297        }
298
299        let connections = self.connection_snapshots_for_ids(connection_ids);
300        let successful_ids = Arc::new(std::sync::Mutex::new(Vec::new()));
301        stream::iter(connections)
302            .for_each_concurrent(self.fanout_concurrency, |snapshot| {
303                let successful_ids = Arc::clone(&successful_ids);
304                async move {
305                    let connection_id = snapshot.0.clone();
306                    let result = self
307                        .send_frame_to_snapshot_without_active(snapshot, frame, parser)
308                        .await;
309                    match result {
310                        Ok(connection_id) => {
311                            Self::record_successful_connection_id(&successful_ids, connection_id);
312                        }
313                        Err(e) => {
314                            tracing::warn!(
315                                "Failed to send frame to connection {}: {:?}",
316                                connection_id,
317                                e
318                            );
319                        }
320                    }
321                }
322            })
323            .await;
324        self.update_connections_active(Self::take_successful_connection_ids(successful_ids));
325
326        Ok(())
327    }
328
329    async fn broadcast_frame(
330        &self,
331        frame: &crate::common::protocol::Frame,
332        parser: Option<&crate::common::MessageParser>,
333    ) -> Result<()> {
334        if let Some(parser) = parser {
335            let connections = self.connection_auth_snapshots();
336            let data = match parser.serialize(frame) {
337                Ok(data) => data,
338                Err(e) => {
339                    tracing::warn!("Failed to serialize broadcast frame: {:?}", e);
340                    return Ok(());
341                }
342            };
343
344            let successful_ids = Arc::new(std::sync::Mutex::new(Vec::new()));
345            stream::iter(connections)
346                .for_each_concurrent(self.fanout_concurrency, |snapshot| {
347                    let successful_ids = Arc::clone(&successful_ids);
348                    let data = data.as_slice();
349                    async move {
350                        let connection_id = snapshot.0.clone();
351                        let result = self
352                            .send_serialized_frame_to_auth_snapshot_without_active(
353                                snapshot, frame, data,
354                            )
355                            .await;
356                        match result {
357                            Ok(connection_id) => {
358                                Self::record_successful_connection_id(
359                                    &successful_ids,
360                                    connection_id,
361                                );
362                            }
363                            Err(e) => {
364                                tracing::warn!(
365                                    "Failed to broadcast frame to connection {}: {:?}",
366                                    connection_id,
367                                    e
368                                );
369                            }
370                        }
371                    }
372                })
373                .await;
374            self.update_connections_active(Self::take_successful_connection_ids(successful_ids));
375
376            return Ok(());
377        }
378
379        let connections = self.connection_snapshots();
380        let successful_ids = Arc::new(std::sync::Mutex::new(Vec::new()));
381        stream::iter(connections)
382            .for_each_concurrent(self.fanout_concurrency, |snapshot| {
383                let successful_ids = Arc::clone(&successful_ids);
384                async move {
385                    let connection_id = snapshot.0.clone();
386                    let result = self
387                        .send_frame_to_snapshot_without_active(snapshot, frame, parser)
388                        .await;
389                    match result {
390                        Ok(connection_id) => {
391                            Self::record_successful_connection_id(&successful_ids, connection_id);
392                        }
393                        Err(e) => {
394                            tracing::warn!(
395                                "Failed to broadcast frame to connection {}: {:?}",
396                                connection_id,
397                                e
398                            );
399                        }
400                    }
401                }
402            })
403            .await;
404        self.update_connections_active(Self::take_successful_connection_ids(successful_ids));
405
406        Ok(())
407    }
408
409    async fn broadcast_frame_except(
410        &self,
411        frame: &crate::common::protocol::Frame,
412        exclude_connection_id: &str,
413        parser: Option<&crate::common::MessageParser>,
414    ) -> Result<()> {
415        if let Some(parser) = parser {
416            let connections = self.connection_auth_snapshots_except(exclude_connection_id);
417            let data = match parser.serialize(frame) {
418                Ok(data) => data,
419                Err(e) => {
420                    tracing::warn!("Failed to serialize broadcast frame: {:?}", e);
421                    return Ok(());
422                }
423            };
424
425            let successful_ids = Arc::new(std::sync::Mutex::new(Vec::new()));
426            stream::iter(connections)
427                .for_each_concurrent(self.fanout_concurrency, |snapshot| {
428                    let successful_ids = Arc::clone(&successful_ids);
429                    let data = data.as_slice();
430                    async move {
431                        let connection_id = snapshot.0.clone();
432                        let result = self
433                            .send_serialized_frame_to_auth_snapshot_without_active(
434                                snapshot, frame, data,
435                            )
436                            .await;
437                        match result {
438                            Ok(connection_id) => {
439                                Self::record_successful_connection_id(
440                                    &successful_ids,
441                                    connection_id,
442                                );
443                            }
444                            Err(e) => {
445                                tracing::warn!(
446                                    "Failed to broadcast frame to connection {}: {:?}",
447                                    connection_id,
448                                    e
449                                );
450                            }
451                        }
452                    }
453                })
454                .await;
455            self.update_connections_active(Self::take_successful_connection_ids(successful_ids));
456
457            return Ok(());
458        }
459
460        let connections = self.connection_snapshots_except(exclude_connection_id);
461        let successful_ids = Arc::new(std::sync::Mutex::new(Vec::new()));
462        stream::iter(connections)
463            .for_each_concurrent(self.fanout_concurrency, |snapshot| {
464                let successful_ids = Arc::clone(&successful_ids);
465                async move {
466                    let connection_id = snapshot.0.clone();
467                    let result = self
468                        .send_frame_to_snapshot_without_active(snapshot, frame, parser)
469                        .await;
470                    match result {
471                        Ok(connection_id) => {
472                            Self::record_successful_connection_id(&successful_ids, connection_id);
473                        }
474                        Err(e) => {
475                            tracing::warn!(
476                                "Failed to broadcast frame to connection {}: {:?}",
477                                connection_id,
478                                e
479                            );
480                        }
481                    }
482                }
483            })
484            .await;
485        self.update_connections_active(Self::take_successful_connection_ids(successful_ids));
486
487        Ok(())
488    }
489}