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 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 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#[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#[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}