Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
37 changes: 9 additions & 28 deletions Cargo.lock

Some generated files are not rendered by default. Learn more about how customized files appear on GitHub.

7 changes: 4 additions & 3 deletions rumqttc/Cargo.toml
Original file line number Diff line number Diff line change
Expand Up @@ -20,7 +20,7 @@ default = ["use-rustls"]
use-rustls = ["use-rustls-no-provider", "tokio-rustls/default"]
use-rustls-no-provider = ["dep:tokio-rustls", "dep:rustls-webpki", "dep:rustls-pemfile", "dep:rustls-native-certs"]
use-native-tls = ["dep:tokio-native-tls", "dep:native-tls"]
websocket = ["dep:async-tungstenite", "dep:ws_stream_tungstenite", "dep:http"]
websocket = ["dep:async-tungstenite", "dep:http", "dep:futures-io", "dep:pin-project-lite"]
proxy = ["dep:async-http-proxy"]

[dependencies]
Expand All @@ -39,9 +39,10 @@ rustls-webpki = { version = "0.102.8", optional = true }
rustls-pemfile = { version = "2.2.0", optional = true }
rustls-native-certs = { version = "0.8.1", optional = true }
# websockets
async-tungstenite = { version = "0.29.0", default-features = false, features = ["tokio-rustls-native-certs"], optional = true }
ws_stream_tungstenite = { version= "0.15.0", default-features = false, features = ["tokio_io"], optional = true }
async-tungstenite = { version = "0.32.0", default-features = false, features = ["tokio-rustls-native-certs", "futures-03-sink"], optional = true }
http = { version = "1.0.0", optional = true }
pin-project-lite = { version = "0.2", optional = true }
futures-io = { version = "0.3", optional = true }
# native-tls
tokio-native-tls = { version = "0.3.1", optional = true }
native-tls = { version = "0.2.12", optional = true }
Expand Down
3 changes: 1 addition & 2 deletions rumqttc/src/eventloop.rs
Original file line number Diff line number Diff line change
Expand Up @@ -23,9 +23,8 @@ use crate::tls;

#[cfg(feature = "websocket")]
use {
crate::websockets::{split_url, validate_response_headers, UrlError},
crate::websockets::{split_url, validate_response_headers, UrlError, WsStream},
async_tungstenite::tungstenite::client::IntoClientRequest,
ws_stream_tungstenite::WsStream,
};

#[cfg(feature = "proxy")]
Expand Down
3 changes: 1 addition & 2 deletions rumqttc/src/v5/eventloop.rs
Original file line number Diff line number Diff line change
Expand Up @@ -23,9 +23,8 @@ use {std::path::Path, tokio::net::UnixStream};

#[cfg(feature = "websocket")]
use {
crate::websockets::{split_url, validate_response_headers, UrlError},
crate::websockets::{split_url, validate_response_headers, UrlError, WsStream},
async_tungstenite::tungstenite::client::IntoClientRequest,
ws_stream_tungstenite::WsStream,
};

#[cfg(feature = "proxy")]
Expand Down
80 changes: 80 additions & 0 deletions rumqttc/src/websockets.rs
Original file line number Diff line number Diff line change
@@ -1,4 +1,14 @@
use std::{pin::Pin, task::Context};

use async_tungstenite::{
bytes::Sender,
tungstenite::{Error, Message},
ByteReader, ByteWriter, WebSocketReceiver, WebSocketSender, WebSocketStream,
};
use futures_util::Stream;
use http::{header::ToStrError, Response};
use pin_project_lite::pin_project;
use tokio::io::{AsyncRead, AsyncWrite};

#[derive(Debug, thiserror::Error)]
pub enum UrlError {
Expand Down Expand Up @@ -71,3 +81,73 @@ fn port(uri: &http::Uri) -> Option<u16> {
_ => None,
})
}

pin_project! {

/// Takes a [`WebSocketStream`] and makes it into a byte IO stream
/// compatible with the rest of rumqttc.
pub(crate) struct WsStream<S> {
#[pin]
read_half: ByteReader<WebSocketReceiver<S>>,
#[pin]
write_half: ByteWriter<WebSocketSender<S>>,
}
}

impl<S> WsStream<S>
where
S: Unpin + futures_io::AsyncWrite + futures_io::AsyncRead,
{
pub fn new(stream: WebSocketStream<S>) -> Self {
let (sender, receiver) = stream.split();

Self {
read_half: ByteReader::new(receiver),
write_half: ByteWriter::new(sender),
}
}
}

impl<S> AsyncRead for WsStream<S>
where
WebSocketReceiver<S>: Stream<Item = Result<Message, Error>> + Unpin,
{
fn poll_read(
self: Pin<&mut Self>,
cx: &mut Context<'_>,
buf: &mut tokio::io::ReadBuf<'_>,
) -> std::task::Poll<std::io::Result<()>> {
let this = self.project();
this.read_half.poll_read(cx, buf)
}
}

impl<S> AsyncWrite for WsStream<S>
where
WebSocketSender<S>: Sender + Unpin,
{
fn poll_write(
self: Pin<&mut Self>,
cx: &mut Context<'_>,
buf: &[u8],
) -> std::task::Poll<Result<usize, std::io::Error>> {
let this = self.project();
this.write_half.poll_write(cx, buf)
}

fn poll_flush(
self: Pin<&mut Self>,
cx: &mut Context<'_>,
) -> std::task::Poll<Result<(), std::io::Error>> {
let this = self.project();
this.write_half.poll_flush(cx)
}

fn poll_shutdown(
self: Pin<&mut Self>,
cx: &mut Context<'_>,
) -> std::task::Poll<Result<(), std::io::Error>> {
let this = self.project();
this.write_half.poll_shutdown(cx)
}
}