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