nautilus_network/transport/
tungstenite.rs1use 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::Message::Frame(frame) => Self::Binary(frame.into_payload()),
59 }
60 }
61}
62
63impl TryFrom<Message> for tungstenite::Message {
64 type Error = TransportError;
65
66 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#[derive(Debug)]
140pub struct TungsteniteTransport<S> {
141 inner: WebSocketStream<S>,
142}
143
144impl<S> TungsteniteTransport<S> {
145 #[inline]
147 #[must_use]
148 pub const fn new(inner: WebSocketStream<S>) -> Self {
149 Self { inner }
150 }
151
152 #[inline]
154 pub fn into_inner(self) -> WebSocketStream<S> {
155 self.inner
156 }
157
158 #[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}