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