Skip to main content

nautilus_network/transport/
sockudo.rs

1// -------------------------------------------------------------------------------------------------
2//  Copyright (C) 2015-2026 Nautech Systems Pty Ltd. All rights reserved.
3//  https://nautechsystems.io
4//
5//  Licensed under the GNU Lesser General Public License Version 3.0 (the "License");
6//  You may not use this file except in compliance with the License.
7//  You may obtain a copy of the License at https://www.gnu.org/licenses/lgpl-3.0.en.html
8//
9//  Unless required by applicable law or agreed to in writing, software
10//  distributed under the License is distributed on an "AS IS" BASIS,
11//  WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12//  See the License for the specific language governing permissions and
13//  limitations under the License.
14// -------------------------------------------------------------------------------------------------
15
16//! `sockudo-ws` backend for the transport abstraction.
17//!
18//! Mirrors the layout of the [`tungstenite`](super::tungstenite) module: provides
19//! `From`/`TryFrom` conversions between the neutral [`Message`] / [`TransportError`]
20//! and sockudo's native types, plus a [`SockudoTransport<S>`] adapter that lifts a
21//! sockudo [`WebSocketStream<S>`] into the backend-agnostic [`WsTransport`] trait.
22//!
23//! The `Message` enums are structurally identical: both carry payloads as `bytes::Bytes`
24//! across all five variants, so conversions are zero-copy and infallible.
25//!
26//! Request bytes come from sockudo's header-aware builder; the response loop stays
27//! local so upgrade rejections map to [`TransportError::UpgradeRejected`] with
28//! host-context logging instead of collapsing into a handshake failure.
29
30use 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
54// WebSocket upgrade headers we always set, plus body-framing headers that have
55// no place on a GET upgrade.
56const 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
70/// Performs the client handshake, building the upgrade request with sockudo's
71/// header-aware builder and reading the response locally.
72///
73/// The local response loop preserves the rejection status for
74/// [`TransportError::UpgradeRejected`] and the retry-aware log severity; caller
75/// pre-validates `extra_headers` via [`validate_extra_headers`].
76pub(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
154// Gated on the error variant rather than its message. `HandshakeFailed` means the response parsed
155// as HTTP but failed handshake validation, which is the only case where a rejection status is
156// recoverable; `InvalidHttp` stays permanent. Matching the diagnostic text would put an upstream
157// string into transport semantics, and recovering after any error would let a plausible status
158// line with malformed headers pass as a rejection.
159fn 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    // A 101 that still failed validation is a protocol fault, not a rejection.
174    status.parse().ok().filter(|status| *status != 101)
175}
176
177// A rejection the client goes on to retry is reported at WARN: an ERROR arms the
178// `shutdown_on_error` trigger, which stops the node while the reconnect loop is still recovering.
179// Permanent rejections, malformed responses, and failed accept-key checks retain ERROR handling.
180//
181// `host` names the endpoint so a node running several connections shows which one failed. It is
182// the `Host` header built from the URL authority, so it carries no userinfo, path, or signed
183// query; the request path stays out of the message for that reason.
184fn 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
223/// Replay bytes read during the handshake before forwarding to the inner IO.
224pub(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    /// Converts a neutral [`Message`] into a Sockudo [`SockudoMessage`].
290    ///
291    /// Conversion is infallible: both enums carry payloads as `bytes::Bytes` across
292    /// all variants. Sockudo validates UTF-8 on Text frames at parse time, not at
293    /// send time, so feeding it non-UTF-8 bytes via [`Self::Text`] is the caller's
294    /// responsibility.
295    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            // Keepalive and idle deadlines are dead connections, TimedOut takes
342            // the connection-drop warn path.
343            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
351/// Adapter that lifts a `sockudo-ws` [`WebSocketStream<S>`] into a
352/// backend-agnostic [`WsTransport`].
353///
354/// Translates messages and errors to the neutral types on the way through
355/// `Stream::poll_next` and `Sink<Message>::start_send` / `poll_*`. The
356/// underlying stream is owned and forwarded to via pin projection.
357///
358/// If flushing an outbound frame returns `Pending`, the next [`Stream::poll_next`] retries the
359/// flush before reading. This prevents queued control responses from being stranded when write
360/// backpressure coincides with a quiet reader.
361pub struct SockudoTransport<S> {
362    inner: WebSocketStream<S>,
363    pending_flush: bool,
364}
365
366impl<S> SockudoTransport<S> {
367    /// Wraps an established Sockudo WebSocket stream.
368    #[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    /// Consumes the adapter and returns the underlying stream.
378    #[inline]
379    pub fn into_inner(self) -> WebSocketStream<S> {
380        self.inner
381    }
382
383    /// Borrows the underlying stream.
384    #[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        // Drain any flush that returned Pending on a prior poll so queued
405        // control responses (Pong, close reply) reach the peer before the
406        // next read. Errors are dropped here; subsequent writes through the
407        // sink half surface them.
408        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        // Sockudo queues automatic Pong / close-response frames into the
423        // write buffer during poll_next. Nudge a flush so they reach the peer
424        // promptly even on a reader-only client; track a pending flush so the
425        // next poll retries when backpressure stalls the write socket.
426        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    // The log-capture harness is Linux-only for CI stability.
815    #[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}