use std::future::Future; use std::net::SocketAddr; use std::pin::Pin; use std::sync::Arc; use bytes::Bytes; use futures::stream::{Stream, StreamExt}; use http::{HeaderMap, StatusCode}; use thiserror::Error; use tokio_tungstenite::Connector; use tokio_tungstenite::tungstenite::{ Bytes as WsBytes, Message as TungsteniteMessage, protocol::CloseFrame as TungsteniteClose, protocol::frame::coding::CloseCode as TungsteniteCloseCode, }; use url::Url; #[derive(Debug, Error)] pub enum NetworkError { #[error("connect: {0}")] Connect(String), #[error("timeout: {0}")] Timeout(String), #[error("redirect: {0}")] Redirect(String), #[error("transport: {0}")] Transport(String), #[error("body: {0}")] Body(String), #[error("protocol: {0}")] Protocol(String), } pub struct HttpRequest { pub url: Url, pub headers: HeaderMap, } pub type BodyStream = Pin> + Send + 'static>>; pub struct HttpResponseHead { pub status: StatusCode, pub headers: HeaderMap, pub content_length: Option, pub body: BodyStream, } pub type HttpResult = Result; pub type HttpResponseFuture = Pin + Send + 'static>>; pub trait HttpTransport: Send + Sync + 'static { fn execute(&self, request: HttpRequest) -> HttpResponseFuture; } #[derive(Clone, Debug)] pub struct ReqwestHttp { client: reqwest::Client, } impl ReqwestHttp { pub fn new(client: reqwest::Client) -> Self { Self { client } } pub fn shared(client: reqwest::Client) -> Arc { Arc::new(Self::new(client)) } } impl HttpTransport for ReqwestHttp { fn execute(&self, request: HttpRequest) -> HttpResponseFuture { let client = self.client.clone(); Box::pin(async move { let resp = client .get(request.url) .headers(request.headers) .send() .await .map_err(map_reqwest)?; let status = resp.status(); let headers = resp.headers().clone(); let content_length = resp.content_length(); let body: BodyStream = Box::pin( resp.bytes_stream() .map(|chunk| chunk.map_err(|e| NetworkError::Body(e.to_string()))), ); Ok(HttpResponseHead { status, headers, content_length, body, }) }) } } /// Lets jacquard resolve identities over the workspace reqwest (0.13); jacquard's own /// `HttpClient` impl is against reqwest 0.12, which is built here without TLS. impl jacquard_common::http_client::HttpClient for ReqwestHttp { type Error = reqwest::Error; async fn send_http( &self, request: http::Request>, ) -> Result>, reqwest::Error> { let (parts, body) = request.into_parts(); let mut req = self .client .request(parts.method, parts.uri.to_string()) .body(body); for (name, value) in parts.headers.iter() { req = req.header(name, value); } let resp = req.send().await?; let mut builder = http::Response::builder().status(resp.status()); for (name, value) in resp.headers().iter() { builder = builder.header(name, value); } let body = resp.bytes().await?.to_vec(); Ok(builder .body(body) .expect("response parts came from reqwest")) } } fn map_reqwest(err: reqwest::Error) -> NetworkError { let msg = err.to_string(); if err.is_timeout() { NetworkError::Timeout(msg) } else if err.is_connect() { NetworkError::Connect(msg) } else if err.is_redirect() { NetworkError::Redirect(msg) } else { NetworkError::Transport(msg) } } #[derive(Clone, Debug)] pub enum WsMessage { Text(String), Binary(Bytes), Ping(Bytes), Pong(Bytes), Close { code: u16, reason: String }, } pub type WsSendFuture<'a> = Pin> + Send + 'a>>; pub type WsMessageFuture<'a> = Pin>> + Send + 'a>>; pub trait WsSink: Send + 'static { fn send<'a>(&'a mut self, message: WsMessage) -> WsSendFuture<'a>; } pub trait WsStream: Send + 'static { fn next<'a>(&'a mut self) -> WsMessageFuture<'a>; } pub struct WsConn { pub sink: Box, pub stream: Box, } pub type WsConnectFuture = Pin> + Send + 'static>>; pub trait WsTransport: Send + Sync + 'static { fn connect(&self, url: Url) -> WsConnectFuture; } pub type AddrGuard = Arc Result<(), NetworkError> + Send + Sync>; #[derive(Debug, Error)] pub enum WsTlsError { #[error("no root certificates in the system trust store: {0:?}")] NoRoots(Vec), #[error("all {0} certificates in the system trust store failed to parse")] Unparsable(usize), #[error("the aws-lc-rs provider doesn't support any of the default protocol versions: {0}")] Versions(rustls::Error), } #[derive(Clone)] pub struct WsTls(Arc); impl WsTls { pub fn from_native_roots() -> Result { let loaded = rustls_native_certs::load_native_certs(); if loaded.certs.is_empty() { return Err(WsTlsError::NoRoots(loaded.errors)); } let mut roots = rustls::RootCertStore::empty(); let (accepted, rejected) = roots.add_parsable_certificates(loaded.certs); if accepted == 0 { return Err(WsTlsError::Unparsable(rejected)); } let provider = Arc::new(rustls::crypto::aws_lc_rs::default_provider()); let config = rustls::ClientConfig::builder_with_provider(provider) .with_safe_default_protocol_versions() .map_err(WsTlsError::Versions)? .with_root_certificates(roots) .with_no_client_auth(); Ok(Self(Arc::new(config))) } fn connector(&self) -> Connector { Connector::Rustls(Arc::clone(&self.0)) } } pub struct TungsteniteWs { tls: WsTls, } impl TungsteniteWs { pub fn shared(tls: WsTls) -> Arc { Arc::new(Self { tls }) } } fn wired(ws: TungsteniteWsStream) -> WsConn { let (sink, stream) = futures::StreamExt::split(ws); WsConn { sink: Box::new(TungsteniteSink { inner: sink }), stream: Box::new(TungsteniteStream { inner: stream }), } } impl WsTransport for TungsteniteWs { fn connect(&self, url: Url) -> WsConnectFuture { let connector = self.tls.connector(); Box::pin(async move { let (ws, _resp) = tokio_tungstenite::connect_async_tls_with_config( url.as_str(), None, false, Some(connector), ) .await .map_err(|e| NetworkError::Connect(e.to_string()))?; Ok(wired(ws)) }) } } pub struct GuardedWs { guard: AddrGuard, tls: WsTls, } impl GuardedWs { pub fn shared(guard: AddrGuard, tls: WsTls) -> Arc { Arc::new(Self { guard, tls }) } } impl WsTransport for GuardedWs { fn connect(&self, url: Url) -> WsConnectFuture { let guard = self.guard.clone(); let connector = self.tls.connector(); Box::pin(async move { let host = url .host_str() .ok_or_else(|| NetworkError::Connect("ws url missing host".to_owned()))? .to_owned(); let port = url .port_or_known_default() .ok_or_else(|| NetworkError::Connect("ws url missing port".to_owned()))?; let addrs: Vec = tokio::net::lookup_host((host.as_str(), port)) .await .map_err(|e| NetworkError::Connect(e.to_string()))? .collect(); guard(&addrs)?; let addr = addrs .into_iter() .next() .ok_or_else(|| NetworkError::Connect(format!("no addresses for {host}")))?; let tcp = tokio::net::TcpStream::connect(addr) .await .map_err(|e| NetworkError::Connect(e.to_string()))?; let (ws, _resp) = tokio_tungstenite::client_async_tls_with_config( url.as_str(), tcp, None, Some(connector), ) .await .map_err(|e| NetworkError::Connect(e.to_string()))?; Ok(wired(ws)) }) } } type TungsteniteWsStream = tokio_tungstenite::WebSocketStream>; struct TungsteniteSink { inner: futures::stream::SplitSink, } impl WsSink for TungsteniteSink { fn send<'a>(&'a mut self, message: WsMessage) -> WsSendFuture<'a> { Box::pin(async move { use futures::SinkExt; self.inner .send(message_to_tungstenite(message)) .await .map_err(|e| NetworkError::Transport(e.to_string())) }) } } struct TungsteniteStream { inner: futures::stream::SplitStream, } impl WsStream for TungsteniteStream { fn next<'a>(&'a mut self) -> WsMessageFuture<'a> { Box::pin(async move { let item = StreamExt::next(&mut self.inner).await?; Some( item.map_err(|e| NetworkError::Transport(e.to_string())) .and_then(message_from_tungstenite), ) }) } } fn message_to_tungstenite(message: WsMessage) -> TungsteniteMessage { match message { WsMessage::Text(text) => TungsteniteMessage::Text(text.into()), WsMessage::Binary(bytes) => TungsteniteMessage::Binary(WsBytes::copy_from_slice(&bytes)), WsMessage::Ping(bytes) => TungsteniteMessage::Ping(WsBytes::copy_from_slice(&bytes)), WsMessage::Pong(bytes) => TungsteniteMessage::Pong(WsBytes::copy_from_slice(&bytes)), WsMessage::Close { code, reason } => TungsteniteMessage::Close(Some(TungsteniteClose { code: TungsteniteCloseCode::from(code), reason: reason.into(), })), } } fn message_from_tungstenite(message: TungsteniteMessage) -> Result { match message { TungsteniteMessage::Text(t) => Ok(WsMessage::Text(t.to_string())), TungsteniteMessage::Binary(b) => Ok(WsMessage::Binary(Bytes::copy_from_slice(&b))), TungsteniteMessage::Ping(b) => Ok(WsMessage::Ping(Bytes::copy_from_slice(&b))), TungsteniteMessage::Pong(b) => Ok(WsMessage::Pong(Bytes::copy_from_slice(&b))), TungsteniteMessage::Close(close) => { let (code, reason) = close .map(|c| (u16::from(c.code), c.reason.to_string())) .unwrap_or((1000, String::new())); Ok(WsMessage::Close { code, reason }) } TungsteniteMessage::Frame(_) => Err(NetworkError::Protocol( "tungstenite raw frame surfaced unexpectedly".to_owned(), )), } }