Skip to content
Draft
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
92 changes: 69 additions & 23 deletions src/websocket.rs
Original file line number Diff line number Diff line change
Expand Up @@ -5,7 +5,7 @@ use crate::{
protobuf::Message,
socket_client::split_host_port,
sodiumoxide::crypto::secretbox::Key,
tcp::Encrypt,
tcp::{DynTcpStream, Encrypt},
tls::{get_cached_tls_accept_invalid_cert, get_cached_tls_type, upsert_tls_cache, TlsType},
ResultType,
};
Expand All @@ -22,14 +22,14 @@ use std::{
use tokio::{net::TcpStream, time::timeout};
use tokio_native_tls::native_tls::TlsConnector;
use tokio_tungstenite::{
connect_async_tls_with_config, tungstenite::protocol::Message as WsMessage, Connector,
client_async_tls_with_config, tungstenite::protocol::Message as WsMessage, Connector,
MaybeTlsStream, WebSocketStream,
};
use tungstenite::client::IntoClientRequest;
use tungstenite::protocol::Role;

pub struct WsFramedStream {
stream: WebSocketStream<MaybeTlsStream<TcpStream>>,
stream: WebSocketStream<MaybeTlsStream<DynTcpStream>>,
addr: SocketAddr,
encrypt: Option<Encrypt>,
send_timeout: u64,
Expand Down Expand Up @@ -67,16 +67,18 @@ impl WsFramedStream {

async fn connect(
url: &str,
local_addr: Option<SocketAddr>,
proxy_conf: Option<&Socks5Server>,
ms_timeout: u64,
) -> ResultType<WebSocketStream<MaybeTlsStream<TcpStream>>> {
// to-do: websocket proxy.

) -> ResultType<(WebSocketStream<MaybeTlsStream<DynTcpStream>>, SocketAddr)> {
let tls_type = get_cached_tls_type(url);
let is_tls_type_cached = tls_type.is_some();
let tls_type = tls_type.unwrap_or(TlsType::Rustls);
let danger_accept_invalid_cert = get_cached_tls_accept_invalid_cert(&url);
Self::try_connect(
url,
local_addr,
proxy_conf,
ms_timeout,
tls_type,
is_tls_type_cached,
Expand All @@ -89,28 +91,60 @@ impl WsFramedStream {
#[async_recursion]
async fn try_connect(
url: &str,
local_addr: Option<SocketAddr>,
proxy_conf: Option<&Socks5Server>,
ms_timeout: u64,
tls_type: TlsType,
is_tls_type_cached: bool,
danger_accept_invalid_cert: Option<bool>,
original_danger_accept_invalid_certs: Option<bool>,
) -> ResultType<WebSocketStream<MaybeTlsStream<TcpStream>>> {
) -> ResultType<(WebSocketStream<MaybeTlsStream<DynTcpStream>>, SocketAddr)> {
let ws_config = None;
let disable_nagle = false;
let request = url
.into_client_request()
.map_err(|e| Error::new(ErrorKind::Other, e))?;
let connector =
Self::get_connector(&tls_type, danger_accept_invalid_cert.unwrap_or(false))?;
match timeout(
Duration::from_millis(ms_timeout),
connect_async_tls_with_config(request, ws_config, disable_nagle, connector),
)
.await?
let host = request
.uri()
.host()
.ok_or_else(|| Error::new(ErrorKind::InvalidInput, "Invalid WebSocket URL: no host"))?;
let port = request
.uri()
.port_u16()
.or_else(|| match request.uri().scheme_str() {
Some("wss") => Some(443),
Some("ws") => Some(80),
_ => None,
})
.ok_or_else(|| Error::new(ErrorKind::InvalidInput, "Invalid WebSocket URL scheme"))?;
let target = format!("{host}:{port}");
match timeout(Duration::from_millis(ms_timeout), async {
let (stream, addr) = if let Some(proxy_conf) = proxy_conf {
let stream = crate::tcp::FramedStream::connect(
target.as_str(),
local_addr,
proxy_conf,
ms_timeout,
)
.await?;
(stream.0.into_inner(), stream.1)
} else {
let stream = TcpStream::connect(&target).await?;
let addr = stream.peer_addr()?;
(DynTcpStream(Box::new(stream)), addr)
};
Ok::<_, anyhow::Error>(
client_async_tls_with_config(request, stream, ws_config, connector)
.await
.map(|(stream, _)| (stream, addr)),
)
})
.await??
{
Ok((ws_stream, _)) => {
Ok((ws_stream, addr)) => {
upsert_tls_cache(url, tls_type, danger_accept_invalid_cert.unwrap_or(false));
Ok(ws_stream)
Ok((ws_stream, addr))
}
Err(e) => match (tls_type, is_tls_type_cached, danger_accept_invalid_cert) {
(TlsType::Rustls, _, None) => {
Expand All @@ -121,6 +155,8 @@ impl WsFramedStream {
);
Self::try_connect(
url,
local_addr,
proxy_conf,
ms_timeout,
tls_type,
is_tls_type_cached,
Expand All @@ -137,6 +173,8 @@ impl WsFramedStream {
);
Self::try_connect(
url,
local_addr,
proxy_conf,
ms_timeout,
TlsType::NativeTls,
is_tls_type_cached,
Expand All @@ -153,6 +191,8 @@ impl WsFramedStream {
);
Self::try_connect(
url,
local_addr,
proxy_conf,
ms_timeout,
tls_type,
is_tls_type_cached,
Expand Down Expand Up @@ -184,17 +224,22 @@ impl WsFramedStream {

pub async fn new<T: AsRef<str>>(
url: T,
_local_addr: Option<SocketAddr>,
_proxy_conf: Option<&Socks5Server>,
local_addr: Option<SocketAddr>,
proxy_conf: Option<&Socks5Server>,
ms_timeout: u64,
) -> ResultType<Self> {
let stream = Self::connect(url.as_ref(), ms_timeout).await?;
let addr = match stream.get_ref() {
MaybeTlsStream::Plain(tcp) => tcp.peer_addr()?,
MaybeTlsStream::NativeTls(tls) => tls.get_ref().get_ref().get_ref().peer_addr()?,
MaybeTlsStream::Rustls(tls) => tls.get_ref().0.peer_addr()?,
_ => return Err(Error::new(ErrorKind::Other, "Unsupported stream type").into()),
let configured_proxy = if proxy_conf.is_none() {
Config::get_socks()
} else {
None
};
let (stream, addr) = Self::connect(
url.as_ref(),
local_addr,
proxy_conf.or(configured_proxy.as_ref()),
ms_timeout,
)
.await?;

let ws = Self {
stream,
Expand All @@ -213,6 +258,7 @@ impl WsFramedStream {

#[inline]
pub async fn from_tcp_stream(stream: TcpStream, addr: SocketAddr) -> ResultType<Self> {
let stream = DynTcpStream(Box::new(stream));
let ws_stream =
WebSocketStream::from_raw_socket(MaybeTlsStream::Plain(stream), Role::Client, None)
.await;
Expand Down