mod tcp;
use crate::transport::{
error::{TransportError, TransportResult},
protocol::{TransportCommand, TransportEvent},
transport_trait::Transport,
TransportId, TransportIdRef,
};
use lib3h_protocol::DidWork;
use std::{
collections::VecDeque,
io::{Read, Write},
};
type TlsConnectResult<T> = Result<TlsStream<T>, native_tls::HandshakeError<T>>;
type WssHandshakeError<T> = tungstenite::handshake::HandshakeError<
tungstenite::handshake::client::ClientHandshake<TlsStream<T>>,
>;
type WssConnectResult<T> =
Result<(WssStream<T>, tungstenite::handshake::client::Response), WssHandshakeError<T>>;
type BaseStream<T> = T;
type TlsMidHandshake<T> = native_tls::MidHandshakeTlsStream<BaseStream<T>>;
type TlsStream<T> = native_tls::TlsStream<BaseStream<T>>;
type WssMidHandshake<T> =
tungstenite::handshake::MidHandshake<tungstenite::ClientHandshake<TlsStream<T>>>;
type WssStream<T> = tungstenite::protocol::WebSocket<TlsStream<T>>;
type SocketMap<T> = std::collections::HashMap<String, TransportInfo<T>>;
#[derive(Debug)]
enum WssStreamState<T: Read + Write + std::fmt::Debug> {
None,
Connecting(BaseStream<T>),
TlsMidHandshake(TlsMidHandshake<T>),
TlsReady(TlsStream<T>),
WssMidHandshake(WssMidHandshake<T>),
Ready(Box<WssStream<T>>),
}
pub const DEFAULT_HEARTBEAT_MS: usize = 2000;
pub const DEFAULT_HEARTBEAT_WAIT_MS: usize = 5000;
#[derive(Debug)]
struct TransportInfo<T: Read + Write + std::fmt::Debug> {
id: TransportId,
url: url::Url,
last_msg: std::time::Instant,
send_queue: Vec<Vec<u8>>,
stateful_socket: WssStreamState<T>,
}
impl<T: Read + Write + std::fmt::Debug> TransportInfo<T> {
pub fn close(&mut self) -> TransportResult<()> {
if let WssStreamState::Ready(socket) = &mut self.stateful_socket {
socket.close(None)?;
socket.write_pending()?;
}
self.stateful_socket = WssStreamState::None;
Ok(())
}
}
pub type StreamFactory<T> = fn(uri: &str) -> TransportResult<T>;
pub struct TransportWss<T: Read + Write + std::fmt::Debug> {
stream_factory: StreamFactory<T>,
stream_sockets: SocketMap<T>,
event_queue: Vec<TransportEvent>,
n_id: u64,
inbox: VecDeque<TransportCommand>,
}
impl<T: Read + Write + std::fmt::Debug> Transport for TransportWss<T> {
fn connect(&mut self, uri: &str) -> TransportResult<TransportId> {
let uri = url::Url::parse(uri)?;
let host_port = format!(
"{}:{}",
uri.host_str()
.ok_or_else(|| TransportError("bad connect host".into()))?,
uri.port()
.ok_or_else(|| TransportError("bad connect port".into()))?,
);
let socket = (self.stream_factory)(&host_port)?;
let id = self.priv_next_id();
let info = TransportInfo {
id: id.clone(),
url: uri,
last_msg: std::time::Instant::now(),
send_queue: Vec::new(),
stateful_socket: WssStreamState::Connecting(socket),
};
self.stream_sockets.insert(id.clone(), info);
Ok(id)
}
fn close(&mut self, id: &TransportIdRef) -> TransportResult<()> {
if let Some(mut info) = self.stream_sockets.remove(id) {
info.close()?;
}
Ok(())
}
fn close_all(&mut self) -> TransportResult<()> {
let mut errors: Vec<TransportError> = Vec::new();
while !self.stream_sockets.is_empty() {
let key = self
.stream_sockets
.keys()
.next()
.expect("should not be None")
.to_string();
if let Some(mut info) = self.stream_sockets.remove(&key) {
if let Err(e) = info.close() {
errors.push(e);
}
}
}
if errors.is_empty() {
Ok(())
} else {
Err(errors.into())
}
}
fn transport_id_list(&self) -> TransportResult<Vec<TransportId>> {
Ok(self.stream_sockets.keys().map(|k| k.to_string()).collect())
}
fn get_uri(&self, id: &TransportIdRef) -> Option<String> {
let res = self.stream_sockets.get(&id.to_string());
res.map(|info| info.url.as_str().to_string())
}
fn post(&mut self, command: TransportCommand) -> TransportResult<()> {
self.inbox.push_back(command);
Ok(())
}
fn process(&mut self) -> TransportResult<(DidWork, Vec<TransportEvent>)> {
let did_work = self.priv_process_stream_sockets()?;
Ok((did_work, self.event_queue.drain(..).collect()))
}
fn send(&mut self, id_list: &[&TransportIdRef], payload: &[u8]) -> TransportResult<()> {
for id in id_list {
if let Some(info) = self.stream_sockets.get_mut(&id.to_string()) {
info.send_queue.push(payload.to_vec());
}
}
Ok(())
}
fn send_all(&mut self, payload: &[u8]) -> TransportResult<()> {
for info in self.stream_sockets.values_mut() {
info.send_queue.push(payload.to_vec());
}
Ok(())
}
fn bind(&mut self, _url: &str) -> TransportResult<String> {
Ok(String::new())
}
}
impl<T: Read + Write + std::fmt::Debug> TransportWss<T> {
pub fn new(stream_factory: StreamFactory<T>) -> Self {
TransportWss {
stream_factory,
stream_sockets: std::collections::HashMap::new(),
event_queue: Vec::new(),
n_id: 1,
inbox: VecDeque::new(),
}
}
pub fn wait_connect(&mut self, uri: &str) -> TransportResult<TransportId> {
let transport_id = self.connect(&uri)?;
let mut out = Vec::new();
let start = std::time::Instant::now();
while (start.elapsed().as_millis() as usize) < DEFAULT_HEARTBEAT_WAIT_MS {
let (_did_work, evt_lst) = self.process()?;
for evt in evt_lst {
match evt {
TransportEvent::ConnectResult(id) => {
if id == transport_id {
return Ok(id);
}
}
_ => out.push(evt),
}
}
std::thread::sleep(std::time::Duration::from_millis(3));
}
Err(TransportError::new(format!(
"ipc wss connection attempt timed out for '{}'. Received events: {:?}",
transport_id, out
)))
}
fn priv_next_id(&mut self) -> String {
let out = format!("ws{}", self.n_id);
self.n_id += 1;
out
}
fn priv_process_stream_sockets(&mut self) -> TransportResult<bool> {
let mut did_work = false;
let sockets: Vec<(String, TransportInfo<T>)> = self.stream_sockets.drain().collect();
for (id, mut info) in sockets {
if let Err(e) = self.priv_process_socket(&mut did_work, &mut info) {
self.event_queue
.push(TransportEvent::TransportError(info.id.clone(), e));
}
if let WssStreamState::None = info.stateful_socket {
self.event_queue.push(TransportEvent::Closed(info.id));
continue;
}
if info.last_msg.elapsed().as_millis() as usize > DEFAULT_HEARTBEAT_MS {
if let WssStreamState::Ready(socket) = &mut info.stateful_socket {
socket.write_message(tungstenite::Message::Ping(vec![]))?;
}
} else if info.last_msg.elapsed().as_millis() as usize > DEFAULT_HEARTBEAT_WAIT_MS {
self.event_queue.push(TransportEvent::Closed(info.id));
info.stateful_socket = WssStreamState::None;
continue;
}
self.stream_sockets.insert(id, info);
}
Ok(did_work)
}
fn priv_process_socket(
&mut self,
did_work: &mut bool,
info: &mut TransportInfo<T>,
) -> TransportResult<()> {
let socket = std::mem::replace(&mut info.stateful_socket, WssStreamState::None);
match socket {
WssStreamState::None => {
Ok(())
}
WssStreamState::Connecting(socket) => {
info.last_msg = std::time::Instant::now();
*did_work = true;
let connector = native_tls::TlsConnector::builder()
.danger_accept_invalid_certs(true)
.danger_accept_invalid_hostnames(true)
.build()
.expect("failed to build TlsConnector");
info.stateful_socket =
self.priv_tls_handshake(connector.connect(info.url.as_str(), socket))?;
Ok(())
}
WssStreamState::TlsMidHandshake(socket) => {
info.stateful_socket = self.priv_tls_handshake(socket.handshake())?;
Ok(())
}
WssStreamState::TlsReady(socket) => {
info.last_msg = std::time::Instant::now();
*did_work = true;
info.stateful_socket = self
.priv_ws_handshake(&info.id, tungstenite::client(info.url.clone(), socket))?;
Ok(())
}
WssStreamState::WssMidHandshake(socket) => {
info.stateful_socket = self.priv_ws_handshake(&info.id, socket.handshake())?;
Ok(())
}
WssStreamState::Ready(mut socket) => {
let msgs: Vec<Vec<u8>> = info.send_queue.drain(..).collect();
for msg in msgs {
socket.write_message(tungstenite::Message::Binary(msg))?;
}
match socket.read_message() {
Err(tungstenite::error::Error::Io(e)) => {
if e.kind() == std::io::ErrorKind::WouldBlock {
info.stateful_socket = WssStreamState::Ready(socket);
return Ok(());
}
Err(e.into())
}
Err(tungstenite::error::Error::ConnectionClosed(_)) => {
Ok(())
}
Err(e) => Err(e.into()),
Ok(msg) => {
info.last_msg = std::time::Instant::now();
*did_work = true;
let qmsg = match msg {
tungstenite::Message::Text(s) => Some(s.into_bytes()),
tungstenite::Message::Binary(b) => Some(b),
_ => None,
};
if let Some(msg) = qmsg {
self.event_queue
.push(TransportEvent::Received(info.id.clone(), msg));
}
info.stateful_socket = WssStreamState::Ready(socket);
Ok(())
}
}
}
}
}
fn priv_tls_handshake(
&mut self,
res: TlsConnectResult<T>,
) -> TransportResult<WssStreamState<T>> {
match res {
Err(native_tls::HandshakeError::WouldBlock(socket)) => {
Ok(WssStreamState::TlsMidHandshake(socket))
}
Err(e) => Err(e.into()),
Ok(socket) => Ok(WssStreamState::TlsReady(socket)),
}
}
fn priv_ws_handshake(
&mut self,
id: &TransportId,
res: WssConnectResult<T>,
) -> TransportResult<WssStreamState<T>> {
match res {
Err(tungstenite::HandshakeError::Interrupted(socket)) => {
Ok(WssStreamState::WssMidHandshake(socket))
}
Err(e) => Err(e.into()),
Ok((socket, _response)) => {
self.event_queue
.push(TransportEvent::ConnectResult(id.clone()));
Ok(WssStreamState::Ready(Box::new(socket)))
}
}
}
}