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