Skip to main content

nautilus_network/transport/
tungstenite.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//! `tokio-tungstenite` backend for the transport abstraction.
17//!
18//! Provides `From` conversions between the neutral [`Message`] and
19//! [`TransportError`] types and tungstenite's native types, plus the
20//! [`TungsteniteTransport<S>`] adapter that lifts a tungstenite
21//! `WebSocketStream<S>` into a backend-agnostic [`WsTransport`].
22//!
23//! The message conversions are structural (no payload copies): tungstenite
24//! stores payloads in `Bytes` and `Utf8Bytes`, which we re-wrap directly.
25
26use std::{
27    pin::Pin,
28    task::{Context, Poll},
29};
30
31use bytes::Bytes;
32use futures_util::{Sink, Stream};
33use tokio::io::{AsyncRead, AsyncWrite};
34use tokio_tungstenite::{
35    WebSocketStream,
36    tungstenite::{
37        self, Utf8Bytes,
38        protocol::{CloseFrame as TgCloseFrame, frame::coding::CloseCode},
39    },
40};
41
42use super::{
43    error::TransportError,
44    message::{CloseFrame, Message},
45    stream::WsTransport,
46};
47
48impl From<tungstenite::Message> for Message {
49    fn from(value: tungstenite::Message) -> Self {
50        match value {
51            tungstenite::Message::Text(text) => Self::Text(Bytes::from(text)),
52            tungstenite::Message::Binary(data) => Self::Binary(data),
53            tungstenite::Message::Ping(data) => Self::Ping(data),
54            tungstenite::Message::Pong(data) => Self::Pong(data),
55            tungstenite::Message::Close(frame) => Self::Close(frame.map(Into::into)),
56
57            // Tungstenite only emits Frame when explicitly constructed; treat as binary
58            tungstenite::Message::Frame(frame) => Self::Binary(frame.into_payload()),
59        }
60    }
61}
62
63impl TryFrom<Message> for tungstenite::Message {
64    type Error = TransportError;
65
66    /// Converts a neutral [`Message`] into a tungstenite [`tungstenite::Message`].
67    ///
68    /// Validates the `Text` payload as UTF-8 because tungstenite refuses to
69    /// transmit a Text frame whose body is not valid UTF-8. Other variants
70    /// are infallible.
71    ///
72    /// # Errors
73    ///
74    /// Returns [`TransportError::InvalidUtf8`] if a `Text` payload is not
75    /// valid UTF-8.
76    fn try_from(value: Message) -> Result<Self, Self::Error> {
77        Ok(match value {
78            Message::Text(bytes) => match Utf8Bytes::try_from(bytes) {
79                Ok(text) => Self::Text(text),
80                Err(_) => return Err(TransportError::InvalidUtf8),
81            },
82            Message::Binary(bytes) => Self::Binary(bytes),
83            Message::Ping(bytes) => Self::Ping(bytes),
84            Message::Pong(bytes) => Self::Pong(bytes),
85            Message::Close(frame) => Self::Close(frame.map(Into::into)),
86        })
87    }
88}
89
90impl From<TgCloseFrame> for CloseFrame {
91    fn from(value: TgCloseFrame) -> Self {
92        Self {
93            code: u16::from(value.code),
94            reason: value.reason.as_str().to_owned(),
95        }
96    }
97}
98
99impl From<CloseFrame> for TgCloseFrame {
100    fn from(value: CloseFrame) -> Self {
101        Self {
102            code: CloseCode::from(value.code),
103            reason: value.reason.into(),
104        }
105    }
106}
107
108impl From<tungstenite::Error> for TransportError {
109    fn from(value: tungstenite::Error) -> Self {
110        match value {
111            tungstenite::Error::ConnectionClosed | tungstenite::Error::AlreadyClosed => {
112                Self::ConnectionClosed
113            }
114            tungstenite::Error::Io(e) => Self::Io(e),
115            tungstenite::Error::Tls(e) => Self::Tls(e.to_string()),
116            tungstenite::Error::Capacity(e) => match e {
117                tungstenite::error::CapacityError::MessageTooLong { .. } => Self::MessageTooLarge,
118                e @ tungstenite::error::CapacityError::TooManyHeaders => Self::Other(e.to_string()),
119            },
120            tungstenite::Error::Protocol(
121                tungstenite::error::ProtocolError::ResetWithoutClosingHandshake,
122            ) => Self::ConnectionReset,
123            tungstenite::Error::Protocol(e) => Self::Protocol(e.to_string()),
124            tungstenite::Error::Utf8(_) => Self::InvalidUtf8,
125            tungstenite::Error::Url(e) => Self::InvalidUrl(e.to_string()),
126            tungstenite::Error::Http(resp) => Self::UpgradeRejected(resp.status().as_u16()),
127            tungstenite::Error::HttpFormat(e) => Self::Handshake(e.to_string()),
128            other => Self::Other(other.to_string()),
129        }
130    }
131}
132
133/// Adapter that lifts a `tokio-tungstenite` [`WebSocketStream<S>`] into a
134/// backend-agnostic [`WsTransport`].
135///
136/// Translates messages and errors to the neutral types on the way through
137/// `Stream::poll_next` and `Sink<Message>::start_send` / `poll_*`. The
138/// underlying stream is owned and forwarded to via pin projection.
139#[derive(Debug)]
140pub struct TungsteniteTransport<S> {
141    inner: WebSocketStream<S>,
142}
143
144impl<S> TungsteniteTransport<S> {
145    /// Wraps an established Tungstenite WebSocket stream.
146    #[inline]
147    #[must_use]
148    pub const fn new(inner: WebSocketStream<S>) -> Self {
149        Self { inner }
150    }
151
152    /// Consumes the adapter and returns the underlying stream.
153    #[inline]
154    pub fn into_inner(self) -> WebSocketStream<S> {
155        self.inner
156    }
157
158    /// Borrows the underlying stream.
159    #[inline]
160    pub const fn get_ref(&self) -> &WebSocketStream<S> {
161        &self.inner
162    }
163}
164
165impl<S> Stream for TungsteniteTransport<S>
166where
167    S: AsyncRead + AsyncWrite + Unpin,
168{
169    type Item = Result<Message, TransportError>;
170
171    fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
172        match Pin::new(&mut self.inner).poll_next(cx) {
173            Poll::Ready(Some(Ok(msg))) => Poll::Ready(Some(Ok(Message::from(msg)))),
174            Poll::Ready(Some(Err(e))) => Poll::Ready(Some(Err(TransportError::from(e)))),
175            Poll::Ready(None) => Poll::Ready(None),
176            Poll::Pending => Poll::Pending,
177        }
178    }
179}
180
181impl<S> Sink<Message> for TungsteniteTransport<S>
182where
183    S: AsyncRead + AsyncWrite + Unpin,
184{
185    type Error = TransportError;
186
187    fn poll_ready(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
188        Pin::new(&mut self.inner)
189            .poll_ready(cx)
190            .map_err(TransportError::from)
191    }
192
193    fn start_send(mut self: Pin<&mut Self>, item: Message) -> Result<(), Self::Error> {
194        let native = tungstenite::Message::try_from(item)?;
195        Pin::new(&mut self.inner)
196            .start_send(native)
197            .map_err(TransportError::from)
198    }
199
200    fn poll_flush(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
201        Pin::new(&mut self.inner)
202            .poll_flush(cx)
203            .map_err(TransportError::from)
204    }
205
206    fn poll_close(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
207        Pin::new(&mut self.inner)
208            .poll_close(cx)
209            .map_err(TransportError::from)
210    }
211}
212
213const _: fn() = || {
214    fn assert_ws_transport<T: WsTransport>() {}
215    assert_ws_transport::<TungsteniteTransport<tokio::net::TcpStream>>();
216};
217
218#[cfg(test)]
219mod tests {
220    use bytes::Bytes;
221    use rstest::rstest;
222    use tokio_tungstenite::tungstenite::{self, Utf8Bytes};
223
224    use super::*;
225
226    #[rstest]
227    fn round_trip_text() {
228        let original = tungstenite::Message::Text(Utf8Bytes::from("hello"));
229        let neutral: Message = original.into();
230        assert!(neutral.is_text());
231        assert_eq!(neutral.as_bytes(), b"hello");
232
233        let back = tungstenite::Message::try_from(neutral).unwrap();
234        match back {
235            tungstenite::Message::Text(t) => assert_eq!(t.as_str(), "hello"),
236            other => panic!("expected text, was {other:?}"),
237        }
238    }
239
240    #[rstest]
241    fn try_from_text_rejects_invalid_utf8() {
242        let neutral = Message::Text(Bytes::from_static(&[0xFF, 0xFE]));
243        let err = tungstenite::Message::try_from(neutral).unwrap_err();
244        assert!(matches!(err, TransportError::InvalidUtf8));
245    }
246
247    #[rstest]
248    fn round_trip_binary() {
249        let original = tungstenite::Message::Binary(Bytes::from_static(&[1, 2, 3]));
250        let neutral: Message = original.into();
251        assert_eq!(neutral.as_bytes(), &[1, 2, 3]);
252
253        let back = tungstenite::Message::try_from(neutral).unwrap();
254        match back {
255            tungstenite::Message::Binary(b) => assert_eq!(&b[..], &[1, 2, 3]),
256            other => panic!("expected binary, was {other:?}"),
257        }
258    }
259
260    #[rstest]
261    fn round_trip_ping_pong() {
262        let ping = tungstenite::Message::Ping(Bytes::from_static(b"p"));
263        let neutral: Message = ping.into();
264        assert!(neutral.is_ping());
265
266        let pong = tungstenite::Message::Pong(Bytes::from_static(b"q"));
267        let neutral: Message = pong.into();
268        assert!(neutral.is_pong());
269    }
270
271    #[rstest]
272    fn close_frame_round_trip() {
273        let original = tungstenite::Message::Close(Some(TgCloseFrame {
274            code: CloseCode::Normal,
275            reason: "bye".into(),
276        }));
277        let neutral: Message = original.into();
278        let Message::Close(Some(frame)) = &neutral else {
279            panic!("expected close frame");
280        };
281        assert_eq!(frame.code, 1000);
282        assert_eq!(frame.reason, "bye");
283
284        let back = tungstenite::Message::try_from(neutral).unwrap();
285        let tungstenite::Message::Close(Some(frame)) = back else {
286            panic!("expected close frame");
287        };
288        assert_eq!(u16::from(frame.code), 1000);
289        assert_eq!(frame.reason.as_str(), "bye");
290    }
291
292    #[rstest]
293    fn error_translation_closed() {
294        let err: TransportError = tungstenite::Error::ConnectionClosed.into();
295        assert!(matches!(err, TransportError::ConnectionClosed));
296    }
297
298    #[rstest]
299    fn error_translation_reset_without_closing_handshake() {
300        let err: TransportError = tungstenite::Error::Protocol(
301            tungstenite::error::ProtocolError::ResetWithoutClosingHandshake,
302        )
303        .into();
304        assert!(matches!(err, TransportError::ConnectionReset));
305    }
306
307    #[rstest]
308    fn error_translation_utf8() {
309        let err: TransportError = tungstenite::Error::Utf8(String::from("bad")).into();
310        assert!(matches!(err, TransportError::InvalidUtf8));
311    }
312
313    #[rstest]
314    fn error_translation_message_too_long() {
315        let err: TransportError =
316            tungstenite::Error::Capacity(tungstenite::error::CapacityError::MessageTooLong {
317                size: 65,
318                max_size: 64,
319            })
320            .into();
321
322        assert!(matches!(err, TransportError::MessageTooLarge));
323    }
324
325    #[rstest]
326    fn error_translation_too_many_headers() {
327        let err: TransportError =
328            tungstenite::Error::Capacity(tungstenite::error::CapacityError::TooManyHeaders).into();
329
330        let TransportError::Other(message) = err else {
331            panic!("expected other error, was {err:?}");
332        };
333        assert_eq!(message, "Too many headers");
334    }
335
336    #[rstest]
337    fn error_translation_tls() {
338        let err: TransportError =
339            tungstenite::Error::Tls(tungstenite::error::TlsError::InvalidDnsName).into();
340
341        let TransportError::Tls(message) = err else {
342            panic!("expected TLS error, was {err:?}");
343        };
344        assert_eq!(message, "Invalid DNS name");
345    }
346
347    #[rstest]
348    fn error_translation_url() {
349        let err: TransportError =
350            tungstenite::Error::Url(tungstenite::error::UrlError::UnsupportedUrlScheme).into();
351
352        let TransportError::InvalidUrl(message) = err else {
353            panic!("expected invalid URL error, was {err:?}");
354        };
355        assert_eq!(message, "URL scheme not supported");
356    }
357}