Skip to main content

flare_core/transport/
websocket.rs

1use crate::common::error::{FlareError, Result};
2use crate::common::protocol::{Reliability, frame_with_system_command, pong};
3use crate::transport::connection::Connection;
4use crate::transport::events::{
5    ArcObserver, ConnectionEvent, notify_observers as notify_connection_observers,
6    notify_observers_and_clear as notify_connection_observers_and_clear,
7};
8use async_trait::async_trait;
9use bytes::Bytes;
10use futures_util::SinkExt;
11use futures_util::stream::{SplitSink, SplitStream, StreamExt};
12use prost::Message as ProstMessage;
13use std::sync::Arc;
14use tokio::net::TcpStream;
15use tokio::sync::Mutex;
16use tokio_tungstenite::tungstenite::{Error as WsError, Message, error::ProtocolError};
17use tokio_tungstenite::{MaybeTlsStream, WebSocketStream};
18
19// 使用枚举来支持两种类型的 WebSocketStream
20enum WebSocketSink {
21    Tls(Arc<Mutex<SplitSink<WebSocketStream<MaybeTlsStream<TcpStream>>, Message>>>),
22    Plain(Arc<Mutex<SplitSink<WebSocketStream<TcpStream>, Message>>>),
23}
24
25pub struct WebSocketTransport {
26    sink: WebSocketSink,
27    observers: Arc<std::sync::Mutex<Vec<ArcObserver>>>,
28    last_active: Arc<std::sync::Mutex<std::time::Instant>>,
29}
30
31impl WebSocketTransport {
32    pub fn new(stream: WebSocketStream<MaybeTlsStream<TcpStream>>) -> Self {
33        Self::from_stream(stream)
34    }
35
36    fn event_from_websocket_error(error: WsError) -> ConnectionEvent {
37        match error {
38            WsError::ConnectionClosed | WsError::AlreadyClosed => {
39                ConnectionEvent::Disconnected("WebSocket connection closed by peer".to_string())
40            }
41            WsError::Protocol(ProtocolError::ResetWithoutClosingHandshake) => {
42                ConnectionEvent::Disconnected(
43                    "WebSocket peer disconnected without close handshake".to_string(),
44                )
45            }
46            WsError::Io(err)
47                if matches!(
48                    err.kind(),
49                    std::io::ErrorKind::ConnectionReset
50                        | std::io::ErrorKind::ConnectionAborted
51                        | std::io::ErrorKind::BrokenPipe
52                        | std::io::ErrorKind::UnexpectedEof
53                ) =>
54            {
55                ConnectionEvent::Disconnected(format!("WebSocket peer disconnected: {}", err))
56            }
57            other => ConnectionEvent::Error(FlareError::connection_failed(other.to_string())),
58        }
59    }
60
61    /// 从 `WebSocketStream<TcpStream>` 创建(在没有 TLS 时使用)
62    ///
63    /// 使用单独的 Plain 类型,避免 unsafe transmute
64    pub fn from_tcp_stream(stream: WebSocketStream<TcpStream>) -> Self {
65        let (sink_plain, receiver_plain) = stream.split();
66
67        let observers = Arc::new(std::sync::Mutex::new(Vec::new()));
68        let sink_arc = Arc::new(Mutex::new(sink_plain));
69        let last_active = Arc::new(std::sync::Mutex::new(std::time::Instant::now()));
70
71        let task_observers = Arc::clone(&observers);
72        let task_sink = Arc::clone(&sink_arc);
73        let task_last_active = Arc::clone(&last_active);
74        tokio::spawn(async move {
75            Self::receiver_task_plain(receiver_plain, task_observers, task_sink, task_last_active)
76                .await;
77        });
78
79        Self {
80            sink: WebSocketSink::Plain(sink_arc),
81            observers,
82            last_active,
83        }
84    }
85
86    fn from_stream(stream: WebSocketStream<MaybeTlsStream<TcpStream>>) -> Self {
87        let (sink, receiver) = stream.split();
88        let observers = Arc::new(std::sync::Mutex::new(Vec::new()));
89        let sink_arc = Arc::new(Mutex::new(sink));
90        let last_active = Arc::new(std::sync::Mutex::new(std::time::Instant::now()));
91
92        let task_observers = Arc::clone(&observers);
93        let task_sink = Arc::clone(&sink_arc);
94        let task_last_active = Arc::clone(&last_active);
95        tokio::spawn(Self::receiver_task(
96            receiver,
97            task_observers,
98            task_sink,
99            task_last_active,
100        ));
101
102        Self {
103            sink: WebSocketSink::Tls(sink_arc),
104            observers,
105            last_active,
106        }
107    }
108
109    async fn receiver_task(
110        mut receiver: SplitStream<WebSocketStream<MaybeTlsStream<TcpStream>>>,
111        observers_arc: Arc<std::sync::Mutex<Vec<ArcObserver>>>,
112        sink_arc: Arc<Mutex<SplitSink<WebSocketStream<MaybeTlsStream<TcpStream>>, Message>>>,
113        last_active: Arc<std::sync::Mutex<std::time::Instant>>,
114    ) {
115        while let Some(message) = receiver.next().await {
116            // 收到消息时更新活跃时间
117            if let Ok(mut active) = last_active.lock() {
118                *active = std::time::Instant::now();
119            }
120
121            let event = match message {
122                Ok(msg) => match msg {
123                    Message::Text(text) => Some(ConnectionEvent::Message(text.as_bytes().to_vec())),
124                    Message::Binary(data) => Some(ConnectionEvent::Message(data.to_vec())),
125                    Message::Close(frame) => {
126                        let reason = frame
127                            .map(|f| f.reason.to_string())
128                            .unwrap_or_else(|| "Connection closed by peer".to_string());
129                        Some(ConnectionEvent::Disconnected(reason))
130                    }
131                    Message::Ping(data) => {
132                        // 收到 WebSocket 协议层的 PING
133                        // 1. 先回复 WebSocket 协议层的 PONG(保持连接)
134                        // 2. 然后使用 builder 构建应用层的 PONG Frame 并发送
135                        if let Err(e) = Self::send_pong_response_tls(&sink_arc, &data).await {
136                            Some(ConnectionEvent::Error(e))
137                        } else if let Err(e) = Self::send_pong_frame_tls(&sink_arc).await {
138                            Some(ConnectionEvent::Error(e))
139                        } else {
140                            None // PING/PONG 已处理,不需要触发事件
141                        }
142                    }
143                    Message::Pong(_) => {
144                        // 收到 WebSocket 协议层的 PONG,这是对我们之前发送的 PING 的响应
145                        // 使用 builder 构建应用层的 PONG Frame,通过事件通知上层处理
146                        match Self::build_pong_frame() {
147                            Ok(pong_data) => Some(ConnectionEvent::Message(pong_data)),
148                            Err(e) => Some(ConnectionEvent::Error(e)),
149                        }
150                    }
151                    _ => None,
152                },
153                Err(e) => Some(Self::event_from_websocket_error(e)),
154            };
155
156            if let Some(event) = event {
157                let is_terminal = matches!(
158                    event,
159                    ConnectionEvent::Disconnected(_) | ConnectionEvent::Error(_)
160                );
161
162                if is_terminal {
163                    notify_connection_observers_and_clear(
164                        &observers_arc,
165                        &event,
166                        "websocket observers",
167                    );
168                    break;
169                } else {
170                    notify_connection_observers(&observers_arc, &event, "websocket observers");
171                }
172            }
173        }
174    }
175
176    async fn receiver_task_plain(
177        mut receiver: SplitStream<WebSocketStream<TcpStream>>,
178        observers_arc: Arc<std::sync::Mutex<Vec<ArcObserver>>>,
179        sink_arc: Arc<Mutex<SplitSink<WebSocketStream<TcpStream>, Message>>>,
180        last_active: Arc<std::sync::Mutex<std::time::Instant>>,
181    ) {
182        while let Some(message) = receiver.next().await {
183            // 收到消息时更新活跃时间
184            if let Ok(mut active) = last_active.lock() {
185                *active = std::time::Instant::now();
186            }
187
188            let event = match message {
189                Ok(msg) => match msg {
190                    Message::Text(text) => Some(ConnectionEvent::Message(text.as_bytes().to_vec())),
191                    Message::Binary(data) => Some(ConnectionEvent::Message(data.to_vec())),
192                    Message::Close(frame) => {
193                        let reason = frame
194                            .map(|f| f.reason.to_string())
195                            .unwrap_or_else(|| "Connection closed by peer".to_string());
196                        Some(ConnectionEvent::Disconnected(reason))
197                    }
198                    Message::Ping(data) => {
199                        // 收到 WebSocket 协议层的 PING
200                        // 1. 先回复 WebSocket 协议层的 PONG(保持连接)
201                        // 2. 然后使用 builder 构建应用层的 PONG Frame 并发送
202                        if let Err(e) = Self::send_pong_response_plain(&sink_arc, &data).await {
203                            Some(ConnectionEvent::Error(e))
204                        } else if let Err(e) = Self::send_pong_frame_plain(&sink_arc).await {
205                            Some(ConnectionEvent::Error(e))
206                        } else {
207                            None // PING/PONG 已处理,不需要触发事件
208                        }
209                    }
210                    Message::Pong(_) => {
211                        // 收到 WebSocket 协议层的 PONG,这是对我们之前发送的 PING 的响应
212                        // 使用 builder 构建应用层的 PONG Frame,通过事件通知上层处理
213                        match Self::build_pong_frame() {
214                            Ok(pong_data) => Some(ConnectionEvent::Message(pong_data)),
215                            Err(e) => Some(ConnectionEvent::Error(e)),
216                        }
217                    }
218                    _ => None,
219                },
220                Err(e) => Some(Self::event_from_websocket_error(e)),
221            };
222
223            if let Some(event) = event {
224                let is_terminal = matches!(
225                    event,
226                    ConnectionEvent::Disconnected(_) | ConnectionEvent::Error(_)
227                );
228
229                if is_terminal {
230                    notify_connection_observers_and_clear(
231                        &observers_arc,
232                        &event,
233                        "websocket observers",
234                    );
235                    break;
236                } else {
237                    notify_connection_observers(&observers_arc, &event, "websocket observers");
238                }
239            }
240        }
241    }
242
243    fn notify_observers_and_clear(&self, event: &ConnectionEvent) {
244        notify_connection_observers_and_clear(&self.observers, event, "websocket observers");
245    }
246
247    /// 发送 WebSocket 协议层的 PONG 响应 (TLS)
248    #[allow(clippy::type_complexity)]
249    async fn send_pong_response_tls(
250        sink: &Arc<Mutex<SplitSink<WebSocketStream<MaybeTlsStream<TcpStream>>, Message>>>,
251        data: &[u8],
252    ) -> Result<()> {
253        let mut sink = sink.lock().await;
254        sink.send(Message::Pong(Bytes::from(data.to_vec())))
255            .await
256            .map_err(|e| FlareError::connection_failed(e.to_string()))?;
257        Ok(())
258    }
259
260    /// 发送 WebSocket 协议层的 PONG 响应 (Plain)
261    async fn send_pong_response_plain(
262        sink: &Arc<Mutex<SplitSink<WebSocketStream<TcpStream>, Message>>>,
263        data: &[u8],
264    ) -> Result<()> {
265        let mut sink = sink.lock().await;
266        sink.send(Message::Pong(Bytes::from(data.to_vec())))
267            .await
268            .map_err(|e| FlareError::connection_failed(e.to_string()))?;
269        Ok(())
270    }
271
272    /// 构建 PONG Frame 并返回序列化后的数据(用于事件通知)
273    fn build_pong_frame() -> Result<Vec<u8>> {
274        // 使用 builder 构建 PONG Frame
275        let pong_frame = frame_with_system_command(pong(), Reliability::BestEffort);
276
277        // 序列化为 protobuf
278        let mut buf = Vec::new();
279        pong_frame
280            .encode(&mut buf)
281            .map_err(|e| FlareError::encoding_error(e.to_string()))?;
282
283        Ok(buf)
284    }
285
286    /// 发送应用层的 PONG Frame 消息 (TLS)
287    #[allow(clippy::type_complexity)]
288    async fn send_pong_frame_tls(
289        sink: &Arc<Mutex<SplitSink<WebSocketStream<MaybeTlsStream<TcpStream>>, Message>>>,
290    ) -> Result<()> {
291        // 构建 PONG Frame
292        let pong_data = Self::build_pong_frame()?;
293
294        // 通过 WebSocket 发送
295        let mut sink = sink.lock().await;
296        sink.send(Message::Binary(Bytes::from(pong_data)))
297            .await
298            .map_err(|e| FlareError::connection_failed(e.to_string()))?;
299
300        Ok(())
301    }
302
303    /// 发送应用层的 PONG Frame 消息 (Plain)
304    async fn send_pong_frame_plain(
305        sink: &Arc<Mutex<SplitSink<WebSocketStream<TcpStream>, Message>>>,
306    ) -> Result<()> {
307        // 构建 PONG Frame
308        let pong_data = Self::build_pong_frame()?;
309
310        // 通过 WebSocket 发送
311        let mut sink = sink.lock().await;
312        sink.send(Message::Binary(Bytes::from(pong_data)))
313            .await
314            .map_err(|e| FlareError::connection_failed(e.to_string()))?;
315
316        Ok(())
317    }
318}
319
320#[async_trait]
321impl Connection for WebSocketTransport {
322    fn add_observer(&mut self, observer: ArcObserver) {
323        observer.on_event(&ConnectionEvent::Connected);
324        if let Ok(mut observers) = self.observers.lock() {
325            observers.push(observer);
326        }
327    }
328
329    fn remove_observer(&mut self, observer: ArcObserver) {
330        if let Ok(mut observers) = self.observers.lock() {
331            observers.retain(|o| !Arc::ptr_eq(o, &observer));
332        }
333    }
334
335    async fn send(&mut self, data: &[u8]) -> Result<()> {
336        if let Ok(mut active) = self.last_active.lock() {
337            *active = std::time::Instant::now();
338        }
339
340        let message = Message::Binary(Bytes::from(data.to_vec()));
341
342        match &mut self.sink {
343            WebSocketSink::Tls(sink) => {
344                let mut s = sink.lock().await;
345                s.send(message)
346                    .await
347                    .map_err(|e| FlareError::connection_failed(e.to_string()))?;
348            }
349            WebSocketSink::Plain(sink) => {
350                let mut s = sink.lock().await;
351                s.send(message)
352                    .await
353                    .map_err(|e| FlareError::connection_failed(e.to_string()))?;
354            }
355        }
356        Ok(())
357    }
358
359    async fn close(&mut self) -> Result<()> {
360        let close_result = match &mut self.sink {
361            WebSocketSink::Tls(sink) => {
362                let mut s = sink.lock().await;
363                s.close()
364                    .await
365                    .map_err(|e| FlareError::connection_failed(e.to_string()))
366            }
367            WebSocketSink::Plain(sink) => {
368                let mut s = sink.lock().await;
369                s.close()
370                    .await
371                    .map_err(|e| FlareError::connection_failed(e.to_string()))
372            }
373        };
374        self.notify_observers_and_clear(&ConnectionEvent::Disconnected(
375            "Closed by client".to_string(),
376        ));
377        close_result
378    }
379
380    fn last_active_time(&self) -> std::time::Instant {
381        self.last_active
382            .lock()
383            .map(|guard| *guard)
384            .unwrap_or_else(|_| {
385                // 如果锁被 poison,返回当前时间减去一个较大值,表示连接可能有问题
386                std::time::Instant::now() - std::time::Duration::from_secs(3600)
387            })
388    }
389
390    fn update_active_time(&mut self) {
391        if let Ok(mut active) = self.last_active.lock() {
392            *active = std::time::Instant::now();
393        }
394        // 如果锁被 poison,忽略更新(连接可能已经出问题)
395    }
396}
397
398#[cfg(test)]
399mod tests {
400    use super::*;
401
402    #[test]
403    fn reset_without_close_handshake_is_disconnected_event() {
404        let event = WebSocketTransport::event_from_websocket_error(WsError::Protocol(
405            ProtocolError::ResetWithoutClosingHandshake,
406        ));
407
408        assert!(matches!(event, ConnectionEvent::Disconnected(_)));
409        assert!(!event.is_error());
410    }
411
412    #[test]
413    fn transport_peer_disconnect_io_errors_are_disconnected_events() {
414        for kind in [
415            std::io::ErrorKind::ConnectionReset,
416            std::io::ErrorKind::ConnectionAborted,
417            std::io::ErrorKind::BrokenPipe,
418            std::io::ErrorKind::UnexpectedEof,
419        ] {
420            let event = WebSocketTransport::event_from_websocket_error(WsError::Io(
421                std::io::Error::new(kind, "peer closed"),
422            ));
423
424            assert!(
425                matches!(event, ConnectionEvent::Disconnected(_)),
426                "expected {kind:?} to be classified as disconnected, got {event:?}"
427            );
428        }
429    }
430
431    #[test]
432    fn malformed_websocket_protocol_errors_remain_error_events() {
433        let event = WebSocketTransport::event_from_websocket_error(WsError::Protocol(
434            ProtocolError::InvalidOpcode(0x0f),
435        ));
436
437        assert!(matches!(event, ConnectionEvent::Error(_)));
438    }
439}