Skip to main content

eggress_protocol_websocket/
lib.rs

1pub mod error;
2
3use std::pin::Pin;
4use std::task::{Context, Poll};
5
6use bytes::{Buf, BytesMut};
7use eggress_core::BoxStream;
8use futures_util::stream::{SplitSink, SplitStream};
9use futures_util::{Sink, Stream, StreamExt};
10use tokio::io::{AsyncRead, AsyncWrite, ReadBuf};
11use tokio_tungstenite::tungstenite::handshake::server::{ErrorResponse, Request, Response};
12use tokio_tungstenite::tungstenite::Message;
13use tokio_tungstenite::WebSocketStream;
14
15use crate::error::WebSocketError;
16
17const DEFAULT_MAX_MESSAGE_SIZE: usize = 16 * 1024 * 1024;
18
19pub struct WebSocketStreamAdapter<S> {
20    read_half: SplitStream<WebSocketStream<S>>,
21    write_half: SplitSink<WebSocketStream<S>, Message>,
22    read_buf: BytesMut,
23    max_message_size: usize,
24}
25
26impl<S: tokio::io::AsyncRead + tokio::io::AsyncWrite + Unpin + Send + 'static>
27    WebSocketStreamAdapter<S>
28{
29    pub fn new(ws: WebSocketStream<S>, max_message_size: usize) -> Self {
30        let (write_half, read_half) = ws.split();
31        Self {
32            read_half,
33            write_half,
34            read_buf: BytesMut::new(),
35            max_message_size,
36        }
37    }
38
39    pub fn into_boxed(self) -> BoxStream {
40        Box::new(self)
41    }
42
43    fn poll_next_message(
44        mut self: Pin<&mut Self>,
45        cx: &mut Context<'_>,
46    ) -> Poll<Option<Result<Message, WebSocketError>>> {
47        match Pin::new(&mut self.read_half).poll_next(cx) {
48            Poll::Ready(Some(Ok(msg))) => Poll::Ready(Some(Ok(msg))),
49            Poll::Ready(Some(Err(e))) => {
50                Poll::Ready(Some(Err(WebSocketError::Protocol(e.to_string()))))
51            }
52            Poll::Ready(None) => Poll::Ready(None),
53            Poll::Pending => Poll::Pending,
54        }
55    }
56}
57
58impl<S: tokio::io::AsyncRead + tokio::io::AsyncWrite + Unpin + Send + 'static> AsyncRead
59    for WebSocketStreamAdapter<S>
60{
61    fn poll_read(
62        mut self: Pin<&mut Self>,
63        cx: &mut Context<'_>,
64        buf: &mut ReadBuf<'_>,
65    ) -> Poll<std::io::Result<()>> {
66        if !self.read_buf.is_empty() {
67            let to_copy = std::cmp::min(self.read_buf.len(), buf.remaining());
68            buf.put_slice(&self.read_buf[..to_copy]);
69            self.read_buf.advance(to_copy);
70            return Poll::Ready(Ok(()));
71        }
72
73        loop {
74            match self.as_mut().poll_next_message(cx) {
75                Poll::Ready(Some(Ok(Message::Binary(data)))) => {
76                    if data.len() > self.max_message_size {
77                        return Poll::Ready(Err(std::io::Error::new(
78                            std::io::ErrorKind::InvalidData,
79                            WebSocketError::MessageTooLarge {
80                                size: data.len(),
81                                max: self.max_message_size,
82                            },
83                        )));
84                    }
85                    if data.len() <= buf.remaining() {
86                        buf.put_slice(&data);
87                    } else {
88                        let to_copy = buf.remaining();
89                        buf.put_slice(&data[..to_copy]);
90                        self.read_buf.extend_from_slice(&data[to_copy..]);
91                    }
92                    return Poll::Ready(Ok(()));
93                }
94                Poll::Ready(Some(Ok(Message::Close(_)))) => {
95                    return Poll::Ready(Ok(()));
96                }
97                Poll::Ready(Some(Ok(Message::Ping(_)))) => {
98                    continue;
99                }
100                Poll::Ready(Some(Ok(Message::Pong(_)))) => {
101                    continue;
102                }
103                Poll::Ready(Some(Ok(Message::Text(_)))) => {
104                    tracing::warn!("received text frame on WebSocket tunnel, skipping");
105                    continue;
106                }
107                Poll::Ready(Some(Ok(Message::Frame(_)))) => {
108                    continue;
109                }
110                Poll::Ready(Some(Err(e))) => {
111                    return Poll::Ready(Err(std::io::Error::other(e)));
112                }
113                Poll::Ready(None) => {
114                    return Poll::Ready(Ok(()));
115                }
116                Poll::Pending => {
117                    return Poll::Pending;
118                }
119            }
120        }
121    }
122}
123
124impl<S: tokio::io::AsyncRead + tokio::io::AsyncWrite + Unpin + Send + 'static> AsyncWrite
125    for WebSocketStreamAdapter<S>
126{
127    fn poll_write(
128        mut self: Pin<&mut Self>,
129        cx: &mut Context<'_>,
130        buf: &[u8],
131    ) -> Poll<std::io::Result<usize>> {
132        match Pin::new(&mut self.write_half).poll_flush(cx) {
133            Poll::Ready(Ok(())) => {}
134            Poll::Ready(Err(e)) => {
135                return Poll::Ready(Err(std::io::Error::other(WebSocketError::Protocol(
136                    e.to_string(),
137                ))));
138            }
139            Poll::Pending => return Poll::Pending,
140        }
141
142        match Pin::new(&mut self.write_half).start_send(Message::Binary(buf.to_vec().into())) {
143            Ok(()) => {}
144            Err(e) => {
145                return Poll::Ready(Err(std::io::Error::other(WebSocketError::Protocol(
146                    e.to_string(),
147                ))));
148            }
149        }
150
151        match Pin::new(&mut self.write_half).poll_flush(cx) {
152            Poll::Ready(Ok(())) => Poll::Ready(Ok(buf.len())),
153            Poll::Ready(Err(e)) => Poll::Ready(Err(std::io::Error::other(
154                WebSocketError::Protocol(e.to_string()),
155            ))),
156            Poll::Pending => Poll::Ready(Ok(buf.len())),
157        }
158    }
159
160    fn poll_flush(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<std::io::Result<()>> {
161        match Pin::new(&mut self.write_half).poll_flush(cx) {
162            Poll::Ready(Ok(())) => Poll::Ready(Ok(())),
163            Poll::Ready(Err(e)) => Poll::Ready(Err(std::io::Error::other(
164                WebSocketError::Protocol(e.to_string()),
165            ))),
166            Poll::Pending => Poll::Pending,
167        }
168    }
169
170    fn poll_shutdown(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<std::io::Result<()>> {
171        match Pin::new(&mut self.write_half).poll_close(cx) {
172            Poll::Ready(Ok(())) => Poll::Ready(Ok(())),
173            Poll::Ready(Err(e)) => Poll::Ready(Err(std::io::Error::other(
174                WebSocketError::Protocol(e.to_string()),
175            ))),
176            Poll::Pending => Poll::Pending,
177        }
178    }
179}
180
181pub struct WebSocketTunnelServer {
182    max_message_size: usize,
183}
184
185impl WebSocketTunnelServer {
186    pub fn new(max_message_size: usize) -> Self {
187        Self { max_message_size }
188    }
189
190    pub fn with_default_config() -> Self {
191        Self {
192            max_message_size: DEFAULT_MAX_MESSAGE_SIZE,
193        }
194    }
195
196    pub async fn accept_upgrade(
197        &self,
198        stream: tokio::net::TcpStream,
199    ) -> Result<BoxStream, WebSocketError> {
200        self.accept_upgrade_over_stream(stream).await
201    }
202
203    pub async fn accept_upgrade_over_stream<S>(
204        &self,
205        stream: S,
206    ) -> Result<BoxStream, WebSocketError>
207    where
208        S: AsyncRead + AsyncWrite + Unpin + Send + 'static,
209    {
210        let ws_stream = tokio_tungstenite::accept_async(stream)
211            .await
212            .map_err(|e| WebSocketError::Handshake(e.to_string()))?;
213
214        Ok(WebSocketStreamAdapter::new(ws_stream, self.max_message_size).into_boxed())
215    }
216
217    pub async fn accept_upgrade_with_config(
218        &self,
219        stream: tokio::net::TcpStream,
220        config: tokio_tungstenite::tungstenite::protocol::WebSocketConfig,
221    ) -> Result<BoxStream, WebSocketError> {
222        self.accept_upgrade_with_config_over_stream(stream, config)
223            .await
224    }
225
226    pub async fn accept_upgrade_with_config_over_stream<S>(
227        &self,
228        stream: S,
229        config: tokio_tungstenite::tungstenite::protocol::WebSocketConfig,
230    ) -> Result<BoxStream, WebSocketError>
231    where
232        S: AsyncRead + AsyncWrite + Unpin + Send + 'static,
233    {
234        let ws_stream = tokio_tungstenite::accept_async_with_config(stream, Some(config))
235            .await
236            .map_err(|e| WebSocketError::Handshake(e.to_string()))?;
237
238        Ok(WebSocketStreamAdapter::new(ws_stream, self.max_message_size).into_boxed())
239    }
240}
241
242/// Complete a server-side WebSocket upgrade and validate an optional proxy
243/// Basic-Auth header. The returned username is present only when credentials
244/// were supplied and validated on this connection.
245#[allow(clippy::result_large_err)]
246pub async fn accept_upgrade_with_auth<S>(
247    stream: S,
248    credentials: Option<(&str, &str)>,
249) -> Result<(BoxStream, Option<String>), WebSocketError>
250where
251    S: AsyncRead + AsyncWrite + Unpin + Send + 'static,
252{
253    let expected = credentials.map(|(user, pass)| (user.to_string(), pass.to_string()));
254    let accepted_user = std::sync::Arc::new(std::sync::Mutex::new(None::<String>));
255    let accepted_user_for_callback = accepted_user.clone();
256    let ws_stream = tokio_tungstenite::accept_hdr_async(
257        stream,
258        move |request: &Request, response: Response| {
259            let Some((expected_user, expected_password)) = expected.as_ref() else {
260                return Ok(response);
261            };
262            let Some(header) = request
263                .headers()
264                .get("proxy-authorization")
265                .and_then(|value| value.to_str().ok())
266            else {
267                return Err(ErrorResponse::new(None));
268            };
269            let Some((user, password)) = parse_basic_auth(header) else {
270                return Err(ErrorResponse::new(None));
271            };
272            if user != *expected_user || password != *expected_password {
273                return Err(ErrorResponse::new(None));
274            }
275            *accepted_user_for_callback
276                .lock()
277                .unwrap_or_else(|e| e.into_inner()) = Some(user);
278            Ok(response)
279        },
280    )
281    .await
282    .map_err(|e| WebSocketError::Handshake(e.to_string()))?;
283
284    let user = accepted_user
285        .lock()
286        .unwrap_or_else(|e| e.into_inner())
287        .clone();
288    Ok((
289        WebSocketStreamAdapter::new(ws_stream, DEFAULT_MAX_MESSAGE_SIZE).into_boxed(),
290        user,
291    ))
292}
293
294fn parse_basic_auth(value: &str) -> Option<(String, String)> {
295    let encoded = value.strip_prefix("Basic ")?;
296    let bytes = encoded.as_bytes();
297    let mut decoded = Vec::with_capacity(bytes.len() * 3 / 4);
298    let table = b"ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz0123456789+/";
299    let mut buffer = 0u32;
300    let mut bits = 0u8;
301    for byte in bytes.iter().copied().filter(|byte| *byte != b'=') {
302        let value = table.iter().position(|candidate| *candidate == byte)? as u32;
303        buffer = (buffer << 6) | value;
304        bits += 6;
305        if bits >= 8 {
306            bits -= 8;
307            decoded.push((buffer >> bits) as u8);
308        }
309    }
310    let decoded = String::from_utf8(decoded).ok()?;
311    let (user, password) = decoded.split_once(':')?;
312    Some((user.to_string(), password.to_string()))
313}
314
315pub struct WebSocketTunnelClient {
316    max_message_size: usize,
317}
318
319impl WebSocketTunnelClient {
320    pub fn new(max_message_size: usize) -> Self {
321        Self { max_message_size }
322    }
323
324    pub fn with_default_config() -> Self {
325        Self {
326            max_message_size: DEFAULT_MAX_MESSAGE_SIZE,
327        }
328    }
329
330    pub async fn connect(&self, url: &str) -> Result<BoxStream, WebSocketError> {
331        let (ws_stream, _) = tokio_tungstenite::connect_async(url)
332            .await
333            .map_err(|e| WebSocketError::Connect(e.to_string()))?;
334
335        Ok(WebSocketStreamAdapter::new(ws_stream, self.max_message_size).into_boxed())
336    }
337
338    pub async fn connect_with_config(
339        &self,
340        url: &str,
341        config: tokio_tungstenite::tungstenite::protocol::WebSocketConfig,
342    ) -> Result<BoxStream, WebSocketError> {
343        let (ws_stream, _) = tokio_tungstenite::connect_async_with_config(url, Some(config), false)
344            .await
345            .map_err(|e| WebSocketError::Connect(e.to_string()))?;
346
347        Ok(WebSocketStreamAdapter::new(ws_stream, self.max_message_size).into_boxed())
348    }
349
350    pub async fn connect_over_stream<S>(
351        &self,
352        url: &str,
353        stream: S,
354    ) -> Result<BoxStream, WebSocketError>
355    where
356        S: AsyncRead + AsyncWrite + Unpin + Send + 'static,
357    {
358        let (ws_stream, _) = tokio_tungstenite::client_async(url, stream)
359            .await
360            .map_err(|e| WebSocketError::Connect(e.to_string()))?;
361
362        Ok(WebSocketStreamAdapter::new(ws_stream, self.max_message_size).into_boxed())
363    }
364
365    pub async fn connect_over_stream_with_config<S>(
366        &self,
367        url: &str,
368        stream: S,
369        config: tokio_tungstenite::tungstenite::protocol::WebSocketConfig,
370    ) -> Result<BoxStream, WebSocketError>
371    where
372        S: AsyncRead + AsyncWrite + Unpin + Send + 'static,
373    {
374        let (ws_stream, _) = tokio_tungstenite::client_async_with_config(url, stream, Some(config))
375            .await
376            .map_err(|e| WebSocketError::Connect(e.to_string()))?;
377
378        Ok(WebSocketStreamAdapter::new(ws_stream, self.max_message_size).into_boxed())
379    }
380}
381
382#[cfg(test)]
383mod tests {
384    use super::*;
385    use futures_util::SinkExt;
386    use tokio::io::{AsyncReadExt, AsyncWriteExt};
387
388    #[tokio::test]
389    async fn test_websocket_echo() {
390        let server_listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
391        let server_addr = server_listener.local_addr().unwrap();
392
393        let server_handle = tokio::spawn(async move {
394            let (stream, _) = server_listener.accept().await.unwrap();
395            let server = WebSocketTunnelServer::with_default_config();
396            let mut bs = server.accept_upgrade(stream).await.unwrap();
397            let mut buf = [0u8; 15];
398            bs.read_exact(&mut buf).await.unwrap();
399            bs.write_all(&buf).await.unwrap();
400            bs.shutdown().await.unwrap();
401        });
402
403        let (ws_stream, _) = tokio_tungstenite::connect_async(format!("ws://{}", server_addr))
404            .await
405            .unwrap();
406        let (mut sink, mut stream) = ws_stream.split();
407
408        sink.send(Message::Binary(b"hello websocket".to_vec().into()))
409            .await
410            .unwrap();
411
412        let msg = stream.next().await.unwrap().unwrap();
413        match msg {
414            Message::Binary(data) => assert_eq!(&*data, b"hello websocket"),
415            _ => panic!("expected binary frame"),
416        }
417
418        sink.send(Message::Close(None)).await.unwrap();
419        server_handle.await.unwrap();
420    }
421
422    #[tokio::test]
423    async fn test_max_message_size_enforced() {
424        let server_listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
425        let server_addr = server_listener.local_addr().unwrap();
426
427        let server_handle = tokio::spawn(async move {
428            let (stream, _) = server_listener.accept().await.unwrap();
429            let server = WebSocketTunnelServer::new(1024);
430            let mut bs = server.accept_upgrade(stream).await.unwrap();
431            let mut buf = [0u8; 2048];
432            let result = bs.read_exact(&mut buf).await;
433            assert!(result.is_err());
434        });
435
436        let (ws_stream, _) = tokio_tungstenite::connect_async(format!("ws://{}", server_addr))
437            .await
438            .unwrap();
439        let (mut sink, _stream) = ws_stream.split();
440
441        let large_msg = vec![0u8; 2048];
442        sink.send(Message::Binary(large_msg.into())).await.unwrap();
443
444        server_handle.await.unwrap();
445    }
446
447    #[tokio::test]
448    async fn test_close_frame_yields_eof() {
449        let server_listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
450        let server_addr = server_listener.local_addr().unwrap();
451
452        let server_handle = tokio::spawn(async move {
453            let (stream, _) = server_listener.accept().await.unwrap();
454            let server = WebSocketTunnelServer::with_default_config();
455            let mut bs = server.accept_upgrade(stream).await.unwrap();
456            let mut buf = [0u8; 1];
457            let result = bs.read_exact(&mut buf).await;
458            assert!(result.is_err());
459        });
460
461        let (ws_stream, _) = tokio_tungstenite::connect_async(format!("ws://{}", server_addr))
462            .await
463            .unwrap();
464        let (mut sink, _stream) = ws_stream.split();
465
466        sink.send(Message::Close(None)).await.unwrap();
467
468        server_handle.await.unwrap();
469    }
470
471    #[tokio::test]
472    async fn test_ping_pong_skipped() {
473        let server_listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
474        let server_addr = server_listener.local_addr().unwrap();
475
476        let server_handle = tokio::spawn(async move {
477            let (stream, _) = server_listener.accept().await.unwrap();
478            let server = WebSocketTunnelServer::with_default_config();
479            let mut bs = server.accept_upgrade(stream).await.unwrap();
480            let mut buf = [0u8; 10];
481            bs.read_exact(&mut buf).await.unwrap();
482            assert_eq!(&buf, b"after-ping");
483        });
484
485        let (ws_stream, _) = tokio_tungstenite::connect_async(format!("ws://{}", server_addr))
486            .await
487            .unwrap();
488        let (mut sink, _stream) = ws_stream.split();
489
490        sink.send(Message::Ping(b"ping-data".to_vec().into()))
491            .await
492            .unwrap();
493        sink.send(Message::Pong(b"pong-data".to_vec().into()))
494            .await
495            .unwrap();
496        sink.send(Message::Binary(b"after-ping".to_vec().into()))
497            .await
498            .unwrap();
499        sink.send(Message::Close(None)).await.unwrap();
500
501        server_handle.await.unwrap();
502    }
503
504    #[tokio::test]
505    async fn test_text_frame_skipped() {
506        let server_listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
507        let server_addr = server_listener.local_addr().unwrap();
508
509        let server_handle = tokio::spawn(async move {
510            let (stream, _) = server_listener.accept().await.unwrap();
511            let server = WebSocketTunnelServer::with_default_config();
512            let mut bs = server.accept_upgrade(stream).await.unwrap();
513            let mut buf = [0u8; 4];
514            bs.read_exact(&mut buf).await.unwrap();
515            assert_eq!(&buf, b"data");
516        });
517
518        let (ws_stream, _) = tokio_tungstenite::connect_async(format!("ws://{}", server_addr))
519            .await
520            .unwrap();
521        let (mut sink, _stream) = ws_stream.split();
522
523        sink.send(Message::Text("skipped-text".into()))
524            .await
525            .unwrap();
526        sink.send(Message::Binary(b"data".to_vec().into()))
527            .await
528            .unwrap();
529        sink.send(Message::Close(None)).await.unwrap();
530
531        server_handle.await.unwrap();
532    }
533
534    #[tokio::test]
535    async fn test_partial_read_buffering() {
536        let server_listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
537        let server_addr = server_listener.local_addr().unwrap();
538
539        let server_handle = tokio::spawn(async move {
540            let (stream, _) = server_listener.accept().await.unwrap();
541            let server = WebSocketTunnelServer::with_default_config();
542            let mut bs = server.accept_upgrade(stream).await.unwrap();
543
544            let mut buf1 = [0u8; 5];
545            bs.read_exact(&mut buf1).await.unwrap();
546            assert_eq!(&buf1, b"hello");
547
548            let mut buf2 = [0u8; 5];
549            bs.read_exact(&mut buf2).await.unwrap();
550            assert_eq!(&buf2, b"world");
551
552            let mut buf3 = [0u8; 4];
553            bs.read_exact(&mut buf3).await.unwrap();
554            assert_eq!(&buf3, b"done");
555        });
556
557        let (ws_stream, _) = tokio_tungstenite::connect_async(format!("ws://{}", server_addr))
558            .await
559            .unwrap();
560        let (mut sink, _stream) = ws_stream.split();
561
562        sink.send(Message::Binary(b"helloworld".to_vec().into()))
563            .await
564            .unwrap();
565        sink.send(Message::Binary(b"done".to_vec().into()))
566            .await
567            .unwrap();
568        sink.send(Message::Close(None)).await.unwrap();
569
570        server_handle.await.unwrap();
571    }
572
573    #[tokio::test]
574    async fn test_accept_upgrade_with_config() {
575        let server_listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
576        let server_addr = server_listener.local_addr().unwrap();
577
578        let server_handle = tokio::spawn(async move {
579            let (stream, _) = server_listener.accept().await.unwrap();
580            let server = WebSocketTunnelServer::with_default_config();
581            let config = tokio_tungstenite::tungstenite::protocol::WebSocketConfig::default()
582                .max_message_size(Some(8192));
583            let mut bs = server
584                .accept_upgrade_with_config(stream, config)
585                .await
586                .unwrap();
587            let mut buf = [0u8; 6];
588            bs.read_exact(&mut buf).await.unwrap();
589            assert_eq!(&buf, b"config");
590        });
591
592        let (ws_stream, _) = tokio_tungstenite::connect_async(format!("ws://{}", server_addr))
593            .await
594            .unwrap();
595        let (mut sink, _stream) = ws_stream.split();
596
597        sink.send(Message::Binary(b"config".to_vec().into()))
598            .await
599            .unwrap();
600        sink.send(Message::Close(None)).await.unwrap();
601
602        server_handle.await.unwrap();
603    }
604
605    #[tokio::test]
606    async fn test_bidirectional_large_payload() {
607        let server_listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
608        let server_addr = server_listener.local_addr().unwrap();
609
610        let server_handle = tokio::spawn(async move {
611            let (stream, _) = server_listener.accept().await.unwrap();
612            let server = WebSocketTunnelServer::with_default_config();
613            let mut bs = server.accept_upgrade(stream).await.unwrap();
614
615            let mut buf = [0u8; 65536];
616            bs.read_exact(&mut buf).await.unwrap();
617
618            bs.write_all(&buf).await.unwrap();
619            bs.shutdown().await.unwrap();
620        });
621
622        let (ws_stream, _) = tokio_tungstenite::connect_async(format!("ws://{}", server_addr))
623            .await
624            .unwrap();
625        let (mut sink, mut stream) = ws_stream.split();
626
627        let payload: Vec<u8> = (0..65536).map(|i| (i % 256) as u8).collect();
628        sink.send(Message::Binary(payload.clone().into()))
629            .await
630            .unwrap();
631
632        let mut received = Vec::new();
633        loop {
634            match stream.next().await {
635                Some(Ok(Message::Binary(data))) => {
636                    received.extend_from_slice(&data);
637                    if received.len() >= 65536 {
638                        break;
639                    }
640                }
641                Some(Ok(Message::Close(_))) => break,
642                _ => break,
643            }
644        }
645        assert_eq!(&received, &payload);
646
647        sink.send(Message::Close(None)).await.unwrap();
648        server_handle.await.unwrap();
649    }
650
651    #[tokio::test]
652    async fn test_websocket_error_display() {
653        let err = WebSocketError::Handshake("test handshake".into());
654        assert!(err.to_string().contains("test handshake"));
655
656        let err = WebSocketError::Connect("test connect".into());
657        assert!(err.to_string().contains("test connect"));
658
659        let err = WebSocketError::Protocol("test protocol".into());
660        assert!(err.to_string().contains("test protocol"));
661
662        let err = WebSocketError::MessageTooLarge {
663            size: 2048,
664            max: 1024,
665        };
666        assert!(err.to_string().contains("2048"));
667        assert!(err.to_string().contains("1024"));
668    }
669
670    #[tokio::test]
671    async fn test_websocket_client_new() {
672        let client = WebSocketTunnelClient::new(4096);
673        assert_eq!(client.max_message_size, 4096);
674
675        let client = WebSocketTunnelClient::with_default_config();
676        assert_eq!(client.max_message_size, DEFAULT_MAX_MESSAGE_SIZE);
677    }
678
679    #[tokio::test]
680    async fn test_websocket_client_connect() {
681        let server_listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
682        let server_addr = server_listener.local_addr().unwrap();
683
684        let server_handle = tokio::spawn(async move {
685            let (stream, _) = server_listener.accept().await.unwrap();
686            let server = WebSocketTunnelServer::with_default_config();
687            let mut bs = server.accept_upgrade(stream).await.unwrap();
688            let mut buf = [0u8; 5];
689            bs.read_exact(&mut buf).await.unwrap();
690            bs.write_all(b"reply").await.unwrap();
691            bs.shutdown().await.unwrap();
692        });
693
694        let client = WebSocketTunnelClient::with_default_config();
695        let mut bs = client
696            .connect(&format!("ws://{}", server_addr))
697            .await
698            .unwrap();
699        bs.write_all(b"hello").await.unwrap();
700        let mut reply = [0u8; 5];
701        bs.read_exact(&mut reply).await.unwrap();
702        assert_eq!(&reply, b"reply");
703
704        server_handle.await.unwrap();
705    }
706
707    #[tokio::test]
708    async fn test_connect_over_stream() {
709        let server_listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
710        let server_addr = server_listener.local_addr().unwrap();
711
712        let server_handle = tokio::spawn(async move {
713            let (stream, _) = server_listener.accept().await.unwrap();
714            let server = WebSocketTunnelServer::with_default_config();
715            let mut bs = server.accept_upgrade(stream).await.unwrap();
716            let mut buf = [0u8; 12];
717            bs.read_exact(&mut buf).await.unwrap();
718            bs.write_all(b"over-stream!").await.unwrap();
719            bs.shutdown().await.unwrap();
720        });
721
722        let tcp_stream = tokio::net::TcpStream::connect(server_addr).await.unwrap();
723
724        let client = WebSocketTunnelClient::with_default_config();
725        let mut bs = client
726            .connect_over_stream(&format!("ws://{}", server_addr), tcp_stream)
727            .await
728            .unwrap();
729        bs.write_all(b"hello stream").await.unwrap();
730        let mut reply = [0u8; 12];
731        bs.read_exact(&mut reply).await.unwrap();
732        assert_eq!(&reply, b"over-stream!");
733
734        server_handle.await.unwrap();
735    }
736}