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//! sockudo's public HTTP/1.1 client API does not expose custom headers, so this
27//! module provides a handshake path for upgrade requests that need them.
28
29use std::{
30    pin::Pin,
31    task::{Context, Poll},
32};
33
34use bytes::{BufMut, Bytes, BytesMut};
35use futures_util::{Sink, Stream};
36use nautilus_core::string::secret::REDACTED;
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,
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/// Mirror of `sockudo_ws::handshake::client_handshake` (2.0.1) with custom headers.
71///
72/// Caller pre-validates `extra_headers` via [`validate_extra_headers`].
73pub(crate) async fn client_handshake_with_headers<S>(
74    stream: &mut S,
75    host: &str,
76    path: &str,
77    protocol: Option<&str>,
78    extra_headers: &[(String, String)],
79) -> Result<HandshakeResult, TransportError>
80where
81    S: AsyncRead + AsyncWrite + Unpin,
82{
83    use tokio::io::{AsyncReadExt, AsyncWriteExt};
84
85    let key = handshake::generate_key();
86    let request = build_request_with_headers(host, path, &key, protocol, None, extra_headers);
87
88    stream
89        .write_all(&request)
90        .await
91        .map_err(TransportError::from)?;
92    stream.flush().await.map_err(TransportError::from)?;
93
94    let mut buf = BytesMut::with_capacity(4096);
95
96    loop {
97        if buf.len() > MAX_HTTP_HEADER_SIZE {
98            return Err(SockudoError::InvalidHttp("response too large").into());
99        }
100
101        let n = stream
102            .read_buf(&mut buf)
103            .await
104            .map_err(TransportError::from)?;
105        if n == 0 {
106            return Err(TransportError::ConnectionClosed);
107        }
108
109        let parsed = match handshake::parse_response(&buf) {
110            Ok(parsed) => parsed,
111            Err(e) => {
112                log_handshake_response(&e, &buf);
113                return Err(rejected_upgrade_status(&e, &buf)
114                    .map_or_else(|| TransportError::from(e), TransportError::UpgradeRejected));
115            }
116        };
117
118        if let Some((res, consumed)) = parsed {
119            let accept = res.accept.ok_or_else(|| {
120                let e = SockudoError::HandshakeFailed("missing Sec-WebSocket-Accept");
121                log_handshake_response(&e, &buf);
122                TransportError::from(e)
123            })?;
124
125            if !handshake::validate_accept_key(&key, accept) {
126                let e = SockudoError::HandshakeFailed("invalid Sec-WebSocket-Accept");
127                log_handshake_response(&e, &buf);
128                return Err(e.into());
129            }
130
131            let res_protocol = res.protocol.map(String::from);
132            let res_extensions = res.extensions.map(String::from);
133            let leftover = if consumed < buf.len() {
134                Some(buf.split_off(consumed).freeze())
135            } else {
136                None
137            };
138
139            return Ok(HandshakeResult {
140                path: path.to_string(),
141                protocol: res_protocol,
142                extensions: res_extensions,
143                leftover,
144            });
145        }
146    }
147}
148
149// Gated on the error variant rather than its message. `HandshakeFailed` means the response parsed
150// as HTTP but failed handshake validation, which is the only case where a rejection status is
151// recoverable; `InvalidHttp` stays permanent. Matching the diagnostic text would put an upstream
152// string into transport semantics, and recovering after any error would let a plausible status
153// line with malformed headers pass as a rejection.
154fn rejected_upgrade_status(err: &SockudoError, buf: &[u8]) -> Option<u16> {
155    if !matches!(err, SockudoError::HandshakeFailed(_)) {
156        return None;
157    }
158
159    let status_line_end = buf.windows(2).position(|window| window == b"\r\n")?;
160    let status_line = std::str::from_utf8(&buf[..status_line_end]).ok()?;
161    let mut parts = status_line.split_whitespace();
162    let version = parts.next()?;
163    let status = parts.next()?;
164    if !version.starts_with("HTTP/1.") || status.len() != 3 {
165        return None;
166    }
167
168    // A 101 that still failed validation is a protocol fault, not a rejection.
169    status.parse().ok().filter(|status| *status != 101)
170}
171
172fn log_handshake_response(err: &SockudoError, buf: &BytesMut) {
173    log::error!(
174        "Sockudo handshake failed for {REDACTED}: {err}; response bytes={}",
175        buf.len()
176    );
177}
178
179// Mirror of `sockudo_ws::handshake::build_request` (2.0.1) with `extra_headers`
180// appended; caller pre-validates.
181fn build_request_with_headers(
182    host: &str,
183    path: &str,
184    key: &str,
185    protocol: Option<&str>,
186    extensions: Option<&str>,
187    extra_headers: &[(String, String)],
188) -> Bytes {
189    let mut buf = BytesMut::with_capacity(512);
190
191    buf.put_slice(b"GET ");
192    buf.put_slice(path.as_bytes());
193    buf.put_slice(b" HTTP/1.1\r\n");
194    buf.put_slice(b"Host: ");
195    buf.put_slice(host.as_bytes());
196    buf.put_slice(b"\r\n");
197    buf.put_slice(b"Upgrade: websocket\r\n");
198    buf.put_slice(b"Connection: Upgrade\r\n");
199    buf.put_slice(b"Sec-WebSocket-Key: ");
200    buf.put_slice(key.as_bytes());
201    buf.put_slice(b"\r\n");
202    buf.put_slice(b"Sec-WebSocket-Version: 13\r\n");
203
204    if let Some(proto) = protocol {
205        buf.put_slice(b"Sec-WebSocket-Protocol: ");
206        buf.put_slice(proto.as_bytes());
207        buf.put_slice(b"\r\n");
208    }
209
210    if let Some(ext) = extensions {
211        buf.put_slice(b"Sec-WebSocket-Extensions: ");
212        buf.put_slice(ext.as_bytes());
213        buf.put_slice(b"\r\n");
214    }
215
216    for (name, value) in extra_headers {
217        buf.put_slice(name.as_bytes());
218        buf.put_slice(b": ");
219        buf.put_slice(value.as_bytes());
220        buf.put_slice(b"\r\n");
221    }
222
223    buf.put_slice(b"\r\n");
224    buf.freeze()
225}
226
227pub(crate) fn validate_extra_headers(headers: &[(String, String)]) -> Result<(), SockudoError> {
228    for (name, value) in headers {
229        validate_extra_header(name, value)?;
230    }
231    Ok(())
232}
233
234fn validate_extra_header(name: &str, value: &str) -> Result<(), SockudoError> {
235    let parsed_name = name
236        .parse::<http::HeaderName>()
237        .map_err(|_| SockudoError::InvalidHttp("invalid header name"))?;
238
239    if RESERVED_UPGRADE_HEADERS.contains(&parsed_name.as_str()) {
240        return Err(SockudoError::InvalidHttp(
241            "reserved upgrade header not allowed in extra_headers",
242        ));
243    }
244
245    http::HeaderValue::from_str(value)
246        .map_err(|_| SockudoError::InvalidHttp("invalid header value"))?;
247    Ok(())
248}
249
250/// Replay bytes read during the handshake before forwarding to the inner IO.
251pub(crate) struct PrefixedIo<S> {
252    inner: S,
253    prefix: Bytes,
254}
255
256impl<S> PrefixedIo<S> {
257    pub(crate) const fn new(inner: S, prefix: Bytes) -> Self {
258        Self { inner, prefix }
259    }
260}
261
262impl<S> AsyncRead for PrefixedIo<S>
263where
264    S: AsyncRead + Unpin,
265{
266    fn poll_read(
267        mut self: Pin<&mut Self>,
268        cx: &mut Context<'_>,
269        buf: &mut ReadBuf<'_>,
270    ) -> Poll<std::io::Result<()>> {
271        if !self.prefix.is_empty() {
272            let n = self.prefix.len().min(buf.remaining());
273            let chunk = self.prefix.split_to(n);
274            buf.put_slice(&chunk);
275            return Poll::Ready(Ok(()));
276        }
277
278        Pin::new(&mut self.inner).poll_read(cx, buf)
279    }
280}
281
282impl<S> AsyncWrite for PrefixedIo<S>
283where
284    S: AsyncWrite + Unpin,
285{
286    fn poll_write(
287        mut self: Pin<&mut Self>,
288        cx: &mut Context<'_>,
289        buf: &[u8],
290    ) -> Poll<std::io::Result<usize>> {
291        Pin::new(&mut self.inner).poll_write(cx, buf)
292    }
293
294    fn poll_flush(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<std::io::Result<()>> {
295        Pin::new(&mut self.inner).poll_flush(cx)
296    }
297
298    fn poll_shutdown(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<std::io::Result<()>> {
299        Pin::new(&mut self.inner).poll_shutdown(cx)
300    }
301}
302
303impl From<SockudoMessage> for Message {
304    fn from(value: SockudoMessage) -> Self {
305        match value {
306            SockudoMessage::Text(b) => Self::Text(b),
307            SockudoMessage::Binary(b) => Self::Binary(b),
308            SockudoMessage::Ping(b) => Self::Ping(b),
309            SockudoMessage::Pong(b) => Self::Pong(b),
310            SockudoMessage::Close(reason) => Self::Close(reason.map(Into::into)),
311        }
312    }
313}
314
315impl From<Message> for SockudoMessage {
316    /// Converts a neutral [`Message`] into a Sockudo [`SockudoMessage`].
317    ///
318    /// Conversion is infallible: both enums carry payloads as `bytes::Bytes` across
319    /// all variants. Sockudo validates UTF-8 on Text frames at parse time, not at
320    /// send time, so feeding it non-UTF-8 bytes via [`Self::Text`] is the caller's
321    /// responsibility.
322    fn from(value: Message) -> Self {
323        match value {
324            Message::Text(b) => Self::Text(b),
325            Message::Binary(b) => Self::Binary(b),
326            Message::Ping(b) => Self::Ping(b),
327            Message::Pong(b) => Self::Pong(b),
328            Message::Close(frame) => Self::Close(frame.map(Into::into)),
329        }
330    }
331}
332
333impl From<SockudoCloseReason> for CloseFrame {
334    fn from(value: SockudoCloseReason) -> Self {
335        Self {
336            code: value.code,
337            reason: value.reason,
338        }
339    }
340}
341
342impl From<CloseFrame> for SockudoCloseReason {
343    fn from(value: CloseFrame) -> Self {
344        Self {
345            code: value.code,
346            reason: value.reason,
347        }
348    }
349}
350
351impl From<SockudoError> for TransportError {
352    fn from(value: SockudoError) -> Self {
353        match value {
354            SockudoError::Io(e) => Self::Io(e),
355            SockudoError::ConnectionClosed => Self::ConnectionClosed,
356            SockudoError::ConnectionReset => Self::ConnectionReset,
357            SockudoError::Closed(reason) => Self::ClosedByPeer(reason.map(Into::into)),
358            SockudoError::MessageTooLarge => Self::MessageTooLarge,
359            SockudoError::FrameTooLarge => Self::FrameTooLarge,
360            SockudoError::InvalidUtf8 => Self::InvalidUtf8,
361            SockudoError::InvalidFrame(msg) | SockudoError::Protocol(msg) => {
362                Self::Protocol(msg.to_string())
363            }
364            SockudoError::InvalidHttp(msg) | SockudoError::HandshakeFailed(msg) => {
365                Self::Handshake(msg.to_string())
366            }
367            other => Self::Other(other.to_string()),
368        }
369    }
370}
371
372/// Adapter that lifts a `sockudo-ws` [`WebSocketStream<S>`] into a
373/// backend-agnostic [`WsTransport`].
374///
375/// Translates messages and errors to the neutral types on the way through
376/// `Stream::poll_next` and `Sink<Message>::start_send` / `poll_*`. The
377/// underlying stream is owned and forwarded to via pin projection.
378///
379/// If flushing an outbound frame returns `Pending`, the next [`Stream::poll_next`] retries the
380/// flush before reading. This prevents queued control responses from being stranded when write
381/// backpressure coincides with a quiet reader.
382pub struct SockudoTransport<S> {
383    inner: WebSocketStream<S>,
384    pending_flush: bool,
385}
386
387impl<S> SockudoTransport<S> {
388    /// Wraps an established Sockudo WebSocket stream.
389    #[inline]
390    #[must_use]
391    pub const fn new(inner: WebSocketStream<S>) -> Self {
392        Self {
393            inner,
394            pending_flush: false,
395        }
396    }
397
398    /// Consumes the adapter and returns the underlying stream.
399    #[inline]
400    pub fn into_inner(self) -> WebSocketStream<S> {
401        self.inner
402    }
403
404    /// Borrows the underlying stream.
405    #[inline]
406    pub const fn get_ref(&self) -> &WebSocketStream<S> {
407        &self.inner
408    }
409}
410
411impl<S> std::fmt::Debug for SockudoTransport<S> {
412    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
413        f.debug_struct(stringify!(SockudoTransport))
414            .finish_non_exhaustive()
415    }
416}
417
418impl<S> Stream for SockudoTransport<S>
419where
420    S: AsyncRead + AsyncWrite + Unpin,
421{
422    type Item = Result<Message, TransportError>;
423
424    fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
425        // Drain any flush that returned Pending on a prior poll so queued
426        // control responses (Pong, close reply) reach the peer before the
427        // next read. Errors are dropped here; subsequent writes through the
428        // sink half surface them.
429        if self.pending_flush {
430            match Pin::new(&mut self.inner).poll_flush(cx) {
431                Poll::Ready(_) => self.pending_flush = false,
432                Poll::Pending => {}
433            }
434        }
435
436        let result = match Pin::new(&mut self.inner).poll_next(cx) {
437            Poll::Ready(Some(Ok(msg))) => Poll::Ready(Some(Ok(Message::from(msg)))),
438            Poll::Ready(Some(Err(e))) => Poll::Ready(Some(Err(TransportError::from(e)))),
439            Poll::Ready(None) => Poll::Ready(None),
440            Poll::Pending => return Poll::Pending,
441        };
442
443        // Sockudo queues automatic Pong / close-response frames into the
444        // write buffer during poll_next. Nudge a flush so they reach the peer
445        // promptly even on a reader-only client; track a pending flush so the
446        // next poll retries when backpressure stalls the write socket.
447        match Pin::new(&mut self.inner).poll_flush(cx) {
448            Poll::Ready(_) => self.pending_flush = false,
449            Poll::Pending => self.pending_flush = true,
450        }
451
452        result
453    }
454}
455
456impl<S> Sink<Message> for SockudoTransport<S>
457where
458    S: AsyncRead + AsyncWrite + Unpin,
459{
460    type Error = TransportError;
461
462    fn poll_ready(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
463        Pin::new(&mut self.inner)
464            .poll_ready(cx)
465            .map_err(TransportError::from)
466    }
467
468    fn start_send(mut self: Pin<&mut Self>, item: Message) -> Result<(), Self::Error> {
469        Pin::new(&mut self.inner)
470            .start_send(SockudoMessage::from(item))
471            .map_err(TransportError::from)
472    }
473
474    fn poll_flush(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
475        Pin::new(&mut self.inner)
476            .poll_flush(cx)
477            .map_err(TransportError::from)
478    }
479
480    fn poll_close(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
481        Pin::new(&mut self.inner)
482            .poll_close(cx)
483            .map_err(TransportError::from)
484    }
485}
486
487const _: fn() = || {
488    fn assert_ws_transport<T: WsTransport>() {}
489    assert_ws_transport::<SockudoTransport<tokio::net::TcpStream>>();
490};
491
492#[cfg(test)]
493mod tests {
494    use bytes::Bytes;
495    use rstest::rstest;
496    #[cfg(not(feature = "turmoil"))]
497    use sockudo_ws::handshake::generate_accept_key;
498    #[cfg(not(feature = "turmoil"))]
499    use tokio::io::{AsyncRead, AsyncReadExt, AsyncWriteExt, duplex};
500
501    use super::*;
502
503    #[cfg(not(feature = "turmoil"))]
504    async fn read_http_request<S>(stream: &mut S) -> Vec<u8>
505    where
506        S: AsyncRead + Unpin,
507    {
508        let mut buf = Vec::new();
509        let mut chunk = [0u8; 256];
510
511        loop {
512            let n = stream.read(&mut chunk).await.unwrap();
513            assert!(n > 0, "HTTP request closed before headers completed");
514            buf.extend_from_slice(&chunk[..n]);
515            if buf.windows(4).any(|window| window == b"\r\n\r\n") {
516                return buf;
517            }
518        }
519    }
520
521    #[cfg(not(feature = "turmoil"))]
522    fn build_test_response(sec_websocket_key: &str, extra_bytes: &[u8]) -> Vec<u8> {
523        let accept = generate_accept_key(sec_websocket_key);
524        let mut response = format!(
525            concat!(
526                "HTTP/1.1 101 Switching Protocols\r\n",
527                "Upgrade: websocket\r\n",
528                "Connection: Upgrade\r\n",
529                "Sec-WebSocket-Accept: {}\r\n",
530                "\r\n",
531            ),
532            accept
533        )
534        .into_bytes();
535        response.extend_from_slice(extra_bytes);
536        response
537    }
538
539    #[cfg(not(feature = "turmoil"))]
540    fn extract_header<'a>(request: &'a str, name: &str) -> Option<&'a str> {
541        request.lines().find_map(|line| {
542            let (header_name, header_value) = line.split_once(':')?;
543            if header_name.eq_ignore_ascii_case(name) {
544                Some(header_value.trim())
545            } else {
546                None
547            }
548        })
549    }
550
551    #[tokio::test]
552    #[cfg(not(feature = "turmoil"))]
553    async fn client_handshake_with_headers_sends_custom_headers() {
554        let (mut client, mut server) = duplex(4096);
555        let headers = vec![
556            ("ok-access-key".to_string(), "key-1".to_string()),
557            ("ok-access-passphrase".to_string(), "pass-1".to_string()),
558        ];
559
560        let server_task = tokio::spawn(async move {
561            let request = read_http_request(&mut server).await;
562            let request = String::from_utf8(request).unwrap();
563
564            assert!(request.starts_with("GET /ws/v5/public-sbe?instId=BTC-USDT HTTP/1.1\r\n"));
565            assert_eq!(extract_header(&request, "Host"), Some("ws.okx.com:8443"));
566            assert_eq!(extract_header(&request, "ok-access-key"), Some("key-1"));
567            assert_eq!(
568                extract_header(&request, "ok-access-passphrase"),
569                Some("pass-1")
570            );
571
572            let sec_websocket_key = extract_header(&request, "Sec-WebSocket-Key").unwrap();
573            let response = build_test_response(sec_websocket_key, &[]);
574            server.write_all(&response).await.unwrap();
575        });
576
577        let handshake = client_handshake_with_headers(
578            &mut client,
579            "ws.okx.com:8443",
580            "/ws/v5/public-sbe?instId=BTC-USDT",
581            None,
582            &headers,
583        )
584        .await
585        .unwrap();
586
587        assert_eq!(handshake.path, "/ws/v5/public-sbe?instId=BTC-USDT");
588        assert!(handshake.leftover.is_none());
589        server_task.await.unwrap();
590    }
591
592    #[rstest]
593    #[cfg(not(feature = "turmoil"))]
594    #[case::host("Host")]
595    #[case::upgrade("Upgrade")]
596    #[case::connection("Connection")]
597    #[case::sec_websocket_key("Sec-WebSocket-Key")]
598    #[case::sec_websocket_version("Sec-WebSocket-Version")]
599    #[case::sec_websocket_protocol("Sec-WebSocket-Protocol")]
600    #[case::sec_websocket_extensions("Sec-WebSocket-Extensions")]
601    #[case::content_length("Content-Length")]
602    #[case::transfer_encoding("Transfer-Encoding")]
603    #[case::te("TE")]
604    #[case::trailer("Trailer")]
605    fn validate_extra_header_rejects_reserved_upgrade_headers(#[case] name: &str) {
606        let err = validate_extra_header(name, "value").unwrap_err();
607
608        assert!(matches!(
609            err,
610            SockudoError::InvalidHttp("reserved upgrade header not allowed in extra_headers")
611        ));
612    }
613
614    #[tokio::test]
615    #[cfg(not(feature = "turmoil"))]
616    async fn client_handshake_with_headers_rejects_missing_accept() {
617        let (mut client, mut server) = duplex(4096);
618
619        let server_task = tokio::spawn(async move {
620            let _request = read_http_request(&mut server).await;
621            server
622                .write_all(
623                    b"HTTP/1.1 101 Switching Protocols\r\n\
624                      Upgrade: websocket\r\n\
625                      Connection: Upgrade\r\n\
626                      \r\n",
627                )
628                .await
629                .unwrap();
630        });
631
632        let err = client_handshake_with_headers(&mut client, "example.com", "/ws", None, &[])
633            .await
634            .unwrap_err();
635
636        assert!(matches!(
637            err,
638            TransportError::Handshake(ref msg) if msg == "missing Sec-WebSocket-Accept"
639        ));
640        server_task.await.unwrap();
641    }
642
643    #[tokio::test]
644    #[cfg(not(feature = "turmoil"))]
645    async fn client_handshake_with_headers_preserves_rejected_status() {
646        let (mut client, mut server) = duplex(4096);
647
648        let server_task = tokio::spawn(async move {
649            let _request = read_http_request(&mut server).await;
650            server.write_all(b"HTTP/1.1 429\r\n\r\n").await.unwrap();
651        });
652
653        let err = client_handshake_with_headers(&mut client, "example.com", "/ws", None, &[])
654            .await
655            .unwrap_err();
656
657        assert!(matches!(err, TransportError::UpgradeRejected(429)));
658        server_task.await.unwrap();
659    }
660
661    #[tokio::test]
662    #[cfg(not(feature = "turmoil"))]
663    async fn client_handshake_with_headers_returns_leftover_bytes() {
664        let (mut client, mut server) = duplex(4096);
665        let extra = b"\x81\x05hello";
666
667        let server_task = tokio::spawn(async move {
668            let request = read_http_request(&mut server).await;
669            let request = String::from_utf8(request).unwrap();
670            let sec_websocket_key = extract_header(&request, "Sec-WebSocket-Key").unwrap();
671            let response = build_test_response(sec_websocket_key, extra);
672            server.write_all(&response).await.unwrap();
673        });
674
675        let handshake = client_handshake_with_headers(&mut client, "example.com", "/ws", None, &[])
676            .await
677            .unwrap();
678
679        assert_eq!(handshake.leftover.as_deref(), Some(extra.as_slice()));
680        server_task.await.unwrap();
681    }
682
683    #[tokio::test]
684    #[cfg(not(feature = "turmoil"))]
685    async fn prefixed_io_replays_leftover_before_socket() {
686        let (client, mut server) = duplex(4096);
687        let mut prefixed = PrefixedIo::new(client, Bytes::from_static(b"abc"));
688
689        let server_task = tokio::spawn(async move {
690            server.write_all(b"def").await.unwrap();
691        });
692
693        let mut buf = [0u8; 6];
694        prefixed.read_exact(&mut buf).await.unwrap();
695
696        assert_eq!(&buf, b"abcdef");
697        server_task.await.unwrap();
698    }
699
700    #[rstest]
701    fn round_trip_text() {
702        let original = SockudoMessage::Text(Bytes::from_static(b"hello"));
703        let neutral: Message = original.into();
704        assert!(neutral.is_text());
705        assert_eq!(neutral.as_bytes(), b"hello");
706
707        let back: SockudoMessage = neutral.into();
708        match back {
709            SockudoMessage::Text(b) => assert_eq!(&b[..], b"hello"),
710            other => panic!("expected text, was {other:?}"),
711        }
712    }
713
714    #[rstest]
715    fn round_trip_binary() {
716        let original = SockudoMessage::Binary(Bytes::from_static(&[1, 2, 3]));
717        let neutral: Message = original.into();
718        assert_eq!(neutral.as_bytes(), &[1, 2, 3]);
719
720        let back: SockudoMessage = neutral.into();
721        match back {
722            SockudoMessage::Binary(b) => assert_eq!(&b[..], &[1, 2, 3]),
723            other => panic!("expected binary, was {other:?}"),
724        }
725    }
726
727    #[rstest]
728    fn round_trip_ping_pong() {
729        let neutral: Message = SockudoMessage::Ping(Bytes::from_static(b"p")).into();
730        assert!(neutral.is_ping());
731
732        let neutral: Message = SockudoMessage::Pong(Bytes::from_static(b"q")).into();
733        assert!(neutral.is_pong());
734    }
735
736    #[rstest]
737    fn close_frame_round_trip() {
738        let original = SockudoMessage::Close(Some(SockudoCloseReason {
739            code: 1000,
740            reason: "bye".into(),
741        }));
742        let neutral: Message = original.into();
743        let Message::Close(Some(frame)) = &neutral else {
744            panic!("expected close frame");
745        };
746        assert_eq!(frame.code, 1000);
747        assert_eq!(frame.reason, "bye");
748
749        let back: SockudoMessage = neutral.into();
750        let SockudoMessage::Close(Some(reason)) = back else {
751            panic!("expected close frame");
752        };
753        assert_eq!(reason.code, 1000);
754        assert_eq!(reason.reason, "bye");
755    }
756
757    #[rstest]
758    fn error_translation_closed() {
759        let err: TransportError = SockudoError::ConnectionClosed.into();
760        assert!(matches!(err, TransportError::ConnectionClosed));
761    }
762
763    #[rstest]
764    fn error_translation_utf8() {
765        let err: TransportError = SockudoError::InvalidUtf8.into();
766        assert!(matches!(err, TransportError::InvalidUtf8));
767    }
768
769    #[rstest]
770    fn error_translation_handshake() {
771        let err: TransportError = SockudoError::HandshakeFailed("bad").into();
772        assert!(matches!(err, TransportError::Handshake(_)));
773    }
774}