1use std::{
31 pin::Pin,
32 task::{Context, Poll},
33};
34
35use bytes::{Bytes, BytesMut};
36use futures_util::{Sink, Stream};
37use sockudo_ws::{
38 HandshakeResult,
39 error::{CloseReason as SockudoCloseReason, Error as SockudoError},
40 handshake,
41 protocol::Message as SockudoMessage,
42 stream::WebSocketStream,
43};
44use tokio::io::{AsyncRead, AsyncWrite, ReadBuf};
45
46use super::{
47 error::{TransportError, retryable_status},
48 message::{CloseFrame, Message},
49 stream::WsTransport,
50};
51
52const MAX_HTTP_HEADER_SIZE: usize = 8192;
53
54const RESERVED_UPGRADE_HEADERS: &[&str] = &[
57 "host",
58 "upgrade",
59 "connection",
60 "sec-websocket-key",
61 "sec-websocket-version",
62 "sec-websocket-protocol",
63 "sec-websocket-extensions",
64 "content-length",
65 "transfer-encoding",
66 "te",
67 "trailer",
68];
69
70pub(crate) async fn client_handshake_with_headers<S>(
77 stream: &mut S,
78 host: &str,
79 path: &str,
80 extra_headers: &[(String, String)],
81) -> Result<HandshakeResult, TransportError>
82where
83 S: AsyncRead + AsyncWrite + Unpin,
84{
85 use tokio::io::{AsyncReadExt, AsyncWriteExt};
86
87 let key = handshake::generate_key();
88 let request =
89 handshake::build_request_with_headers(host, path, &key, None, None, Some(extra_headers))?;
90
91 stream
92 .write_all(&request)
93 .await
94 .map_err(TransportError::from)?;
95 stream.flush().await.map_err(TransportError::from)?;
96
97 let mut buf = BytesMut::with_capacity(4096);
98
99 loop {
100 if buf.len() > MAX_HTTP_HEADER_SIZE {
101 return Err(SockudoError::InvalidHttp("response too large").into());
102 }
103
104 let n = stream
105 .read_buf(&mut buf)
106 .await
107 .map_err(TransportError::from)?;
108 if n == 0 {
109 return Err(TransportError::ConnectionClosed);
110 }
111
112 let parsed = match handshake::parse_response(&buf) {
113 Ok(parsed) => parsed,
114 Err(e) => {
115 let status = rejected_upgrade_status(&e, &buf);
116 log_handshake_response(host, &e, &buf, status);
117 return Err(
118 status.map_or_else(|| TransportError::from(e), TransportError::UpgradeRejected)
119 );
120 }
121 };
122
123 if let Some((res, consumed)) = parsed {
124 let accept = res.accept.ok_or_else(|| {
125 let e = SockudoError::HandshakeFailed("missing Sec-WebSocket-Accept");
126 log_handshake_response(host, &e, &buf, None);
127 TransportError::from(e)
128 })?;
129
130 if !handshake::validate_accept_key(&key, accept) {
131 let e = SockudoError::HandshakeFailed("invalid Sec-WebSocket-Accept");
132 log_handshake_response(host, &e, &buf, None);
133 return Err(e.into());
134 }
135
136 let res_protocol = res.protocol.map(String::from);
137 let res_extensions = res.extensions.map(String::from);
138 let leftover = if consumed < buf.len() {
139 Some(buf.split_off(consumed).freeze())
140 } else {
141 None
142 };
143
144 return Ok(HandshakeResult {
145 path: path.to_string(),
146 protocol: res_protocol,
147 extensions: res_extensions,
148 leftover,
149 });
150 }
151 }
152}
153
154fn rejected_upgrade_status(err: &SockudoError, buf: &[u8]) -> Option<u16> {
160 if !matches!(err, SockudoError::HandshakeFailed(_)) {
161 return None;
162 }
163
164 let status_line_end = buf.windows(2).position(|window| window == b"\r\n")?;
165 let status_line = std::str::from_utf8(&buf[..status_line_end]).ok()?;
166 let mut parts = status_line.split_whitespace();
167 let version = parts.next()?;
168 let status = parts.next()?;
169 if !version.starts_with("HTTP/1.") || status.len() != 3 {
170 return None;
171 }
172
173 status.parse().ok().filter(|status| *status != 101)
175}
176
177fn log_handshake_response(host: &str, err: &SockudoError, buf: &BytesMut, status: Option<u16>) {
185 let response_bytes = buf.len();
186
187 match status {
188 Some(status) if retryable_status(status) => log::warn!(
189 "Sockudo handshake rejected by {host} with retryable status {status}; response bytes={response_bytes}"
190 ),
191 Some(status) => log::error!(
192 "Sockudo handshake rejected by {host} with permanent status {status}; response bytes={response_bytes}"
193 ),
194 None => log::error!(
195 "Sockudo handshake failed for {host}: {err}; response bytes={response_bytes}"
196 ),
197 }
198}
199
200pub(crate) fn validate_extra_headers(headers: &[(String, String)]) -> Result<(), SockudoError> {
201 for (name, value) in headers {
202 validate_extra_header(name, value)?;
203 }
204 Ok(())
205}
206
207fn validate_extra_header(name: &str, value: &str) -> Result<(), SockudoError> {
208 let parsed_name = name
209 .parse::<http::HeaderName>()
210 .map_err(|_| SockudoError::InvalidHttp("invalid header name"))?;
211
212 if RESERVED_UPGRADE_HEADERS.contains(&parsed_name.as_str()) {
213 return Err(SockudoError::InvalidHttp(
214 "reserved upgrade header not allowed in extra_headers",
215 ));
216 }
217
218 http::HeaderValue::from_str(value)
219 .map_err(|_| SockudoError::InvalidHttp("invalid header value"))?;
220 Ok(())
221}
222
223pub(crate) struct PrefixedIo<S> {
225 inner: S,
226 prefix: Bytes,
227}
228
229impl<S> PrefixedIo<S> {
230 pub(crate) const fn new(inner: S, prefix: Bytes) -> Self {
231 Self { inner, prefix }
232 }
233}
234
235impl<S> AsyncRead for PrefixedIo<S>
236where
237 S: AsyncRead + Unpin,
238{
239 fn poll_read(
240 mut self: Pin<&mut Self>,
241 cx: &mut Context<'_>,
242 buf: &mut ReadBuf<'_>,
243 ) -> Poll<std::io::Result<()>> {
244 if !self.prefix.is_empty() {
245 let n = self.prefix.len().min(buf.remaining());
246 let chunk = self.prefix.split_to(n);
247 buf.put_slice(&chunk);
248 return Poll::Ready(Ok(()));
249 }
250
251 Pin::new(&mut self.inner).poll_read(cx, buf)
252 }
253}
254
255impl<S> AsyncWrite for PrefixedIo<S>
256where
257 S: AsyncWrite + Unpin,
258{
259 fn poll_write(
260 mut self: Pin<&mut Self>,
261 cx: &mut Context<'_>,
262 buf: &[u8],
263 ) -> Poll<std::io::Result<usize>> {
264 Pin::new(&mut self.inner).poll_write(cx, buf)
265 }
266
267 fn poll_flush(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<std::io::Result<()>> {
268 Pin::new(&mut self.inner).poll_flush(cx)
269 }
270
271 fn poll_shutdown(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<std::io::Result<()>> {
272 Pin::new(&mut self.inner).poll_shutdown(cx)
273 }
274}
275
276impl From<SockudoMessage> for Message {
277 fn from(value: SockudoMessage) -> Self {
278 match value {
279 SockudoMessage::Text(b) => Self::Text(b),
280 SockudoMessage::Binary(b) => Self::Binary(b),
281 SockudoMessage::Ping(b) => Self::Ping(b),
282 SockudoMessage::Pong(b) => Self::Pong(b),
283 SockudoMessage::Close(reason) => Self::Close(reason.map(Into::into)),
284 }
285 }
286}
287
288impl From<Message> for SockudoMessage {
289 fn from(value: Message) -> Self {
296 match value {
297 Message::Text(b) => Self::Text(b),
298 Message::Binary(b) => Self::Binary(b),
299 Message::Ping(b) => Self::Ping(b),
300 Message::Pong(b) => Self::Pong(b),
301 Message::Close(frame) => Self::Close(frame.map(Into::into)),
302 }
303 }
304}
305
306impl From<SockudoCloseReason> for CloseFrame {
307 fn from(value: SockudoCloseReason) -> Self {
308 Self {
309 code: value.code,
310 reason: value.reason,
311 }
312 }
313}
314
315impl From<CloseFrame> for SockudoCloseReason {
316 fn from(value: CloseFrame) -> Self {
317 Self {
318 code: value.code,
319 reason: value.reason,
320 }
321 }
322}
323
324impl From<SockudoError> for TransportError {
325 fn from(value: SockudoError) -> Self {
326 match value {
327 SockudoError::Io(e) => Self::Io(e),
328 SockudoError::ConnectionClosed => Self::ConnectionClosed,
329 SockudoError::ConnectionReset => Self::ConnectionReset,
330 SockudoError::Closed(reason) => Self::ClosedByPeer(reason.map(Into::into)),
331 SockudoError::MessageTooLarge => Self::MessageTooLarge,
332 SockudoError::FrameTooLarge => Self::FrameTooLarge,
333 SockudoError::InvalidUtf8 => Self::InvalidUtf8,
334 SockudoError::InvalidFrame(msg) | SockudoError::Protocol(msg) => {
335 Self::Protocol(msg.to_string())
336 }
337 SockudoError::InvalidHttp(msg) | SockudoError::HandshakeFailed(msg) => {
338 Self::Handshake(msg.to_string())
339 }
340
341 timeout @ (SockudoError::HeartbeatTimeout | SockudoError::IdleTimeout) => Self::Io(
344 std::io::Error::new(std::io::ErrorKind::TimedOut, timeout.to_string()),
345 ),
346 other => Self::Other(other.to_string()),
347 }
348 }
349}
350
351pub struct SockudoTransport<S> {
362 inner: WebSocketStream<S>,
363 pending_flush: bool,
364}
365
366impl<S> SockudoTransport<S> {
367 #[inline]
369 #[must_use]
370 pub const fn new(inner: WebSocketStream<S>) -> Self {
371 Self {
372 inner,
373 pending_flush: false,
374 }
375 }
376
377 #[inline]
379 pub fn into_inner(self) -> WebSocketStream<S> {
380 self.inner
381 }
382
383 #[inline]
385 pub const fn get_ref(&self) -> &WebSocketStream<S> {
386 &self.inner
387 }
388}
389
390impl<S> std::fmt::Debug for SockudoTransport<S> {
391 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
392 f.debug_struct(stringify!(SockudoTransport))
393 .finish_non_exhaustive()
394 }
395}
396
397impl<S> Stream for SockudoTransport<S>
398where
399 S: AsyncRead + AsyncWrite + Unpin,
400{
401 type Item = Result<Message, TransportError>;
402
403 fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
404 if self.pending_flush {
409 match Pin::new(&mut self.inner).poll_flush(cx) {
410 Poll::Ready(_) => self.pending_flush = false,
411 Poll::Pending => {}
412 }
413 }
414
415 let result = match Pin::new(&mut self.inner).poll_next(cx) {
416 Poll::Ready(Some(Ok(msg))) => Poll::Ready(Some(Ok(Message::from(msg)))),
417 Poll::Ready(Some(Err(e))) => Poll::Ready(Some(Err(TransportError::from(e)))),
418 Poll::Ready(None) => Poll::Ready(None),
419 Poll::Pending => return Poll::Pending,
420 };
421
422 match Pin::new(&mut self.inner).poll_flush(cx) {
427 Poll::Ready(_) => self.pending_flush = false,
428 Poll::Pending => self.pending_flush = true,
429 }
430
431 result
432 }
433}
434
435impl<S> Sink<Message> for SockudoTransport<S>
436where
437 S: AsyncRead + AsyncWrite + Unpin,
438{
439 type Error = TransportError;
440
441 fn poll_ready(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
442 Pin::new(&mut self.inner)
443 .poll_ready(cx)
444 .map_err(TransportError::from)
445 }
446
447 fn start_send(mut self: Pin<&mut Self>, item: Message) -> Result<(), Self::Error> {
448 Pin::new(&mut self.inner)
449 .start_send(SockudoMessage::from(item))
450 .map_err(TransportError::from)
451 }
452
453 fn poll_flush(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
454 Pin::new(&mut self.inner)
455 .poll_flush(cx)
456 .map_err(TransportError::from)
457 }
458
459 fn poll_close(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
460 Pin::new(&mut self.inner)
461 .poll_close(cx)
462 .map_err(TransportError::from)
463 }
464}
465
466const _: fn() = || {
467 fn assert_ws_transport<T: WsTransport>() {}
468 assert_ws_transport::<SockudoTransport<tokio::net::TcpStream>>();
469};
470
471#[cfg(test)]
472mod tests {
473 use bytes::Bytes;
474 use rstest::rstest;
475 #[cfg(not(feature = "turmoil"))]
476 use sockudo_ws::handshake::generate_accept_key;
477 #[cfg(not(feature = "turmoil"))]
478 use tokio::io::{AsyncRead, AsyncReadExt, AsyncWriteExt, duplex};
479
480 use super::*;
481
482 #[cfg(not(feature = "turmoil"))]
483 async fn read_http_request<S>(stream: &mut S) -> Vec<u8>
484 where
485 S: AsyncRead + Unpin,
486 {
487 let mut buf = Vec::new();
488 let mut chunk = [0u8; 256];
489
490 loop {
491 let n = stream.read(&mut chunk).await.unwrap();
492 assert!(n > 0, "HTTP request closed before headers completed");
493 buf.extend_from_slice(&chunk[..n]);
494 if buf.windows(4).any(|window| window == b"\r\n\r\n") {
495 return buf;
496 }
497 }
498 }
499
500 #[cfg(not(feature = "turmoil"))]
501 fn build_test_response(sec_websocket_key: &str, extra_bytes: &[u8]) -> Vec<u8> {
502 let accept = generate_accept_key(sec_websocket_key);
503 let mut response = format!(
504 concat!(
505 "HTTP/1.1 101 Switching Protocols\r\n",
506 "Upgrade: websocket\r\n",
507 "Connection: Upgrade\r\n",
508 "Sec-WebSocket-Accept: {}\r\n",
509 "\r\n",
510 ),
511 accept
512 )
513 .into_bytes();
514 response.extend_from_slice(extra_bytes);
515 response
516 }
517
518 #[cfg(not(feature = "turmoil"))]
519 fn extract_header<'a>(request: &'a str, name: &str) -> Option<&'a str> {
520 request.lines().find_map(|line| {
521 let (header_name, header_value) = line.split_once(':')?;
522 if header_name.eq_ignore_ascii_case(name) {
523 Some(header_value.trim())
524 } else {
525 None
526 }
527 })
528 }
529
530 #[rstest]
531 #[tokio::test]
532 #[cfg(not(feature = "turmoil"))]
533 async fn handshake_rejects_eof_before_response() {
534 let (mut client, mut server) = duplex(4096);
535
536 let peer = tokio::spawn(async move {
537 read_http_request(&mut server).await;
538 });
539
540 let error = client_handshake_with_headers(&mut client, "localhost", "/", &[])
541 .await
542 .unwrap_err();
543 peer.await.unwrap();
544
545 assert!(matches!(error, TransportError::ConnectionClosed));
546 }
547
548 #[rstest]
549 #[case::at_limit(MAX_HTTP_HEADER_SIZE, "connection closed")]
550 #[case::above_limit(MAX_HTTP_HEADER_SIZE + 1, "handshake failed: response too large")]
551 #[tokio::test]
552 #[cfg(not(feature = "turmoil"))]
553 async fn handshake_rejects_incomplete_headers_at_size_boundary(
554 #[case] response_size: usize,
555 #[case] expected: &str,
556 ) {
557 let (mut client, mut server) = duplex(MAX_HTTP_HEADER_SIZE * 2);
558
559 let peer = tokio::spawn(async move {
560 read_http_request(&mut server).await;
561 let mut response = b"HTTP/1.1 101 Switching Protocols\r\nX-Padding: ".to_vec();
562 response.resize(response_size, b'x');
563 server.write_all(&response).await.unwrap();
564 });
565
566 let error = client_handshake_with_headers(&mut client, "localhost", "/", &[])
567 .await
568 .unwrap_err();
569 peer.await.unwrap();
570
571 assert_eq!(error.to_string(), expected);
572 }
573
574 #[tokio::test]
575 #[cfg(not(feature = "turmoil"))]
576 async fn client_handshake_with_headers_sends_custom_headers() {
577 let (mut client, mut server) = duplex(4096);
578 let headers = vec![
579 ("ok-access-key".to_string(), "key-1".to_string()),
580 ("ok-access-passphrase".to_string(), "pass-1".to_string()),
581 ];
582
583 let server_task = tokio::spawn(async move {
584 let request = read_http_request(&mut server).await;
585 let request = String::from_utf8(request).unwrap();
586
587 assert!(request.starts_with("GET /ws/v5/public-sbe?instId=BTC-USDT HTTP/1.1\r\n"));
588 assert_eq!(extract_header(&request, "Host"), Some("ws.okx.com:8443"));
589 assert_eq!(extract_header(&request, "ok-access-key"), Some("key-1"));
590 assert_eq!(
591 extract_header(&request, "ok-access-passphrase"),
592 Some("pass-1")
593 );
594
595 let sec_websocket_key = extract_header(&request, "Sec-WebSocket-Key").unwrap();
596 let response = build_test_response(sec_websocket_key, &[]);
597 server.write_all(&response).await.unwrap();
598 });
599
600 let handshake = client_handshake_with_headers(
601 &mut client,
602 "ws.okx.com:8443",
603 "/ws/v5/public-sbe?instId=BTC-USDT",
604 &headers,
605 )
606 .await
607 .unwrap();
608
609 assert_eq!(handshake.path, "/ws/v5/public-sbe?instId=BTC-USDT");
610 assert!(handshake.leftover.is_none());
611 server_task.await.unwrap();
612 }
613
614 #[rstest]
615 #[cfg(not(feature = "turmoil"))]
616 #[case::host("Host")]
617 #[case::upgrade("Upgrade")]
618 #[case::connection("Connection")]
619 #[case::sec_websocket_key("Sec-WebSocket-Key")]
620 #[case::sec_websocket_version("Sec-WebSocket-Version")]
621 #[case::sec_websocket_protocol("Sec-WebSocket-Protocol")]
622 #[case::sec_websocket_extensions("Sec-WebSocket-Extensions")]
623 #[case::content_length("Content-Length")]
624 #[case::transfer_encoding("Transfer-Encoding")]
625 #[case::te("TE")]
626 #[case::trailer("Trailer")]
627 fn validate_extra_header_rejects_reserved_upgrade_headers(#[case] name: &str) {
628 let err = validate_extra_header(name, "value").unwrap_err();
629
630 assert!(matches!(
631 err,
632 SockudoError::InvalidHttp("reserved upgrade header not allowed in extra_headers")
633 ));
634 }
635
636 #[tokio::test]
637 #[cfg(not(feature = "turmoil"))]
638 async fn client_handshake_with_headers_rejects_missing_accept() {
639 let (mut client, mut server) = duplex(4096);
640
641 let server_task = tokio::spawn(async move {
642 let _request = read_http_request(&mut server).await;
643 server
644 .write_all(
645 b"HTTP/1.1 101 Switching Protocols\r\n\
646 Upgrade: websocket\r\n\
647 Connection: Upgrade\r\n\
648 \r\n",
649 )
650 .await
651 .unwrap();
652 });
653
654 let err = client_handshake_with_headers(&mut client, "example.com", "/ws", &[])
655 .await
656 .unwrap_err();
657
658 assert!(matches!(
659 err,
660 TransportError::Handshake(ref msg) if msg == "missing Sec-WebSocket-Accept"
661 ));
662 server_task.await.unwrap();
663 }
664
665 #[tokio::test]
666 #[cfg(not(feature = "turmoil"))]
667 async fn client_handshake_with_headers_preserves_rejected_status() {
668 let (mut client, mut server) = duplex(4096);
669
670 let server_task = tokio::spawn(async move {
671 let _request = read_http_request(&mut server).await;
672 server.write_all(b"HTTP/1.1 429\r\n\r\n").await.unwrap();
673 });
674
675 let err = client_handshake_with_headers(&mut client, "example.com", "/ws", &[])
676 .await
677 .unwrap_err();
678
679 assert!(matches!(err, TransportError::UpgradeRejected(429)));
680 server_task.await.unwrap();
681 }
682
683 #[tokio::test]
684 #[cfg(not(feature = "turmoil"))]
685 async fn client_handshake_with_headers_returns_leftover_bytes() {
686 let (mut client, mut server) = duplex(4096);
687 let extra = b"\x81\x05hello";
688
689 let server_task = tokio::spawn(async move {
690 let request = read_http_request(&mut server).await;
691 let request = String::from_utf8(request).unwrap();
692 let sec_websocket_key = extract_header(&request, "Sec-WebSocket-Key").unwrap();
693 let response = build_test_response(sec_websocket_key, extra);
694 server.write_all(&response).await.unwrap();
695 });
696
697 let handshake = client_handshake_with_headers(&mut client, "example.com", "/ws", &[])
698 .await
699 .unwrap();
700
701 assert_eq!(handshake.leftover.as_deref(), Some(extra.as_slice()));
702 server_task.await.unwrap();
703 }
704
705 #[tokio::test]
706 #[cfg(not(feature = "turmoil"))]
707 async fn prefixed_io_replays_leftover_before_socket() {
708 let (client, mut server) = duplex(4096);
709 let mut prefixed = PrefixedIo::new(client, Bytes::from_static(b"abc"));
710
711 let server_task = tokio::spawn(async move {
712 server.write_all(b"def").await.unwrap();
713 });
714
715 let mut buf = [0u8; 6];
716 prefixed.read_exact(&mut buf).await.unwrap();
717
718 assert_eq!(&buf, b"abcdef");
719 server_task.await.unwrap();
720 }
721
722 #[rstest]
723 fn round_trip_text() {
724 let original = SockudoMessage::Text(Bytes::from_static(b"hello"));
725 let neutral: Message = original.into();
726 assert!(neutral.is_text());
727 assert_eq!(neutral.as_bytes(), b"hello");
728
729 let back: SockudoMessage = neutral.into();
730 match back {
731 SockudoMessage::Text(b) => assert_eq!(&b[..], b"hello"),
732 other => panic!("expected text, was {other:?}"),
733 }
734 }
735
736 #[rstest]
737 fn round_trip_binary() {
738 let original = SockudoMessage::Binary(Bytes::from_static(&[1, 2, 3]));
739 let neutral: Message = original.into();
740 assert_eq!(neutral.as_bytes(), &[1, 2, 3]);
741
742 let back: SockudoMessage = neutral.into();
743 match back {
744 SockudoMessage::Binary(b) => assert_eq!(&b[..], &[1, 2, 3]),
745 other => panic!("expected binary, was {other:?}"),
746 }
747 }
748
749 #[rstest]
750 fn round_trip_ping_pong() {
751 let neutral: Message = SockudoMessage::Ping(Bytes::from_static(b"p")).into();
752 assert!(neutral.is_ping());
753
754 let neutral: Message = SockudoMessage::Pong(Bytes::from_static(b"q")).into();
755 assert!(neutral.is_pong());
756 }
757
758 #[rstest]
759 fn close_frame_round_trip() {
760 let original = SockudoMessage::Close(Some(SockudoCloseReason {
761 code: 1000,
762 reason: "bye".into(),
763 }));
764 let neutral: Message = original.into();
765 let Message::Close(Some(frame)) = &neutral else {
766 panic!("expected close frame");
767 };
768 assert_eq!(frame.code, 1000);
769 assert_eq!(frame.reason, "bye");
770
771 let back: SockudoMessage = neutral.into();
772 let SockudoMessage::Close(Some(reason)) = back else {
773 panic!("expected close frame");
774 };
775 assert_eq!(reason.code, 1000);
776 assert_eq!(reason.reason, "bye");
777 }
778
779 #[rstest]
780 fn error_translation_closed() {
781 let err: TransportError = SockudoError::ConnectionClosed.into();
782 assert!(matches!(err, TransportError::ConnectionClosed));
783 }
784
785 #[rstest]
786 fn error_translation_utf8() {
787 let err: TransportError = SockudoError::InvalidUtf8.into();
788 assert!(matches!(err, TransportError::InvalidUtf8));
789 }
790
791 #[rstest]
792 fn error_translation_handshake() {
793 let err: TransportError = SockudoError::HandshakeFailed("bad").into();
794 assert!(matches!(err, TransportError::Handshake(_)));
795 }
796
797 #[rstest]
798 #[case(SockudoError::HeartbeatTimeout, "WebSocket Pong deadline expired")]
799 #[case(SockudoError::IdleTimeout, "WebSocket inbound idle deadline expired")]
800 fn error_translation_timeouts_are_timed_out_io(
801 #[case] sockudo: SockudoError,
802 #[case] message: &str,
803 ) {
804 let err: TransportError = sockudo.into();
805 let TransportError::Io(io_err) = &err else {
806 panic!("expected I/O timeout, was: {err:?}");
807 };
808
809 assert_eq!(io_err.kind(), std::io::ErrorKind::TimedOut);
810 assert_eq!(io_err.to_string(), message);
811 assert_eq!(err.to_string(), format!("I/O error: {message}"));
812 }
813
814 #[cfg(not(feature = "turmoil"))]
816 #[cfg(target_os = "linux")]
817 #[cfg(not(all(feature = "simulation", madsim)))]
818 mod handshake_logging {
819 use log::Level;
820 use rstest::rstest;
821 use tokio::io::{AsyncWriteExt, duplex};
822
823 use super::read_http_request;
824 use crate::{
825 logging::tests::capture_logs_for,
826 transport::{error::TransportError, sockudo::client_handshake_with_headers},
827 };
828
829 const LOG_TARGETS: &[&str] = &["nautilus_network::transport::sockudo"];
830 const HOST: &str = "ws.example.com:8443";
831 const PATH_SECRET: &str = "handshake-path-secret";
832
833 #[rstest]
834 #[case::retryable_bad_gateway("HTTP/1.1 502 Bad Gateway\r\n\r\n", Level::Warn, Some(502))]
835 #[case::retryable_rate_limited(
836 "HTTP/1.1 429 Too Many Requests\r\n\r\n",
837 Level::Warn,
838 Some(429)
839 )]
840 #[case::permanent_unauthorized(
841 "HTTP/1.1 401 Unauthorized\r\n\r\n",
842 Level::Error,
843 Some(401)
844 )]
845 #[case::permanent_not_found("HTTP/1.1 404 Not Found\r\n\r\n", Level::Error, Some(404))]
846 #[case::malformed_status_line("NOT-HTTP 502 Bad Gateway\r\n\r\n", Level::Error, None)]
847 #[case::missing_accept_key(
848 "HTTP/1.1 101 Switching Protocols\r\nUpgrade: websocket\r\nConnection: Upgrade\r\n\r\n",
849 Level::Error,
850 None
851 )]
852 #[case::invalid_accept_key(
853 "HTTP/1.1 101 Switching Protocols\r\nUpgrade: websocket\r\nConnection: Upgrade\r\n\
854 Sec-WebSocket-Accept: AAAAAAAAAAAAAAAAAAAAAAAAAAA=\r\n\r\n",
855 Level::Error,
856 None
857 )]
858 #[tokio::test]
859 async fn handshake_failure_log_level_follows_retry_policy(
860 #[case] response: &'static str,
861 #[case] expected_level: Level,
862 #[case] expected_status: Option<u16>,
863 ) {
864 let capture = capture_logs_for(LOG_TARGETS).await;
865 let (mut client, mut server) = duplex(4096);
866
867 let server_task = tokio::spawn(async move {
868 let _request = read_http_request(&mut server).await;
869 server.write_all(response.as_bytes()).await.unwrap();
870 server.flush().await.unwrap();
871 });
872
873 let err = client_handshake_with_headers(
874 &mut client,
875 HOST,
876 &format!("/ws?token={PATH_SECRET}"),
877 &[],
878 )
879 .await
880 .expect_err("handshake should fail");
881 server_task.await.unwrap();
882
883 match expected_status {
884 Some(status) => assert!(
885 matches!(err, TransportError::UpgradeRejected(actual) if actual == status),
886 "expected upgrade rejection {status}, was: {err:?}"
887 ),
888 None => assert!(
889 matches!(err, TransportError::Handshake(_)),
890 "expected a permanent handshake failure, was: {err:?}"
891 ),
892 }
893
894 let messages = capture.messages();
895 assert_eq!(
896 messages.len(),
897 1,
898 "expected exactly one handshake log, was: {messages:?}"
899 );
900
901 let (level, message) = &messages[0];
902 assert_eq!(*level, expected_level, "unexpected level for: {message}");
903 assert!(
904 message.contains(HOST),
905 "log should name the endpoint, was: {message}"
906 );
907
908 if let Some(status) = expected_status {
909 assert!(
910 message.contains(&status.to_string()),
911 "log should name the rejection status, was: {message}"
912 );
913 }
914
915 assert!(
916 !message.contains(PATH_SECRET),
917 "log must not carry the request path, was: {message}"
918 );
919 }
920 }
921}