Skip to main content

eggress_protocol_websocket/
lib.rs

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