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
19enum 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 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 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 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 }
142 }
143 Message::Pong(_) => {
144 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 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 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 }
209 }
210 Message::Pong(_) => {
211 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 #[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 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 fn build_pong_frame() -> Result<Vec<u8>> {
274 let pong_frame = frame_with_system_command(pong(), Reliability::BestEffort);
276
277 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 #[allow(clippy::type_complexity)]
288 async fn send_pong_frame_tls(
289 sink: &Arc<Mutex<SplitSink<WebSocketStream<MaybeTlsStream<TcpStream>>, Message>>>,
290 ) -> Result<()> {
291 let pong_data = Self::build_pong_frame()?;
293
294 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 async fn send_pong_frame_plain(
305 sink: &Arc<Mutex<SplitSink<WebSocketStream<TcpStream>, Message>>>,
306 ) -> Result<()> {
307 let pong_data = Self::build_pong_frame()?;
309
310 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 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 }
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}