1use 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
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>(
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
149fn 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 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
179fn 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
250pub(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 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
372pub struct SockudoTransport<S> {
383 inner: WebSocketStream<S>,
384 pending_flush: bool,
385}
386
387impl<S> SockudoTransport<S> {
388 #[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 #[inline]
400 pub fn into_inner(self) -> WebSocketStream<S> {
401 self.inner
402 }
403
404 #[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 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 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}