use crate::client::ClientConfig;
use crate::keepalive::KeepaliveBehavior;
use crate::reconnect::{ReconnectBehavior, Reconnector};
use crate::transport::{TransportError, TransportSink, TransportStream};
use base64::Engine;
use base64::prelude::BASE64_STANDARD;
use core::future::Future;
use core::pin::Pin;
use std::sync::Arc;
use std::time::Duration;
use tokio::net::TcpStream;
use tokio_tungstenite::tungstenite::client::IntoClientRequest;
use tokio_tungstenite::tungstenite::http::Request;
use tokio_tungstenite::tungstenite::http::header::{AUTHORIZATION, SEC_WEBSOCKET_PROTOCOL};
use tokio_tungstenite::{Connector, MaybeTlsStream, WebSocketStream, client_async_tls_with_config};
use url::Url;
const DEFAULT_TIMEOUT: Duration = Duration::from_secs(5);
#[derive(Clone)]
pub struct ConnectOptions<'a> {
pub username: Option<&'a str>,
pub password: Option<&'a str>,
pub timeout: Option<Duration>,
pub reconnect: ReconnectBehavior,
pub tls_config: Option<Arc<rustls::ClientConfig>>,
pub reconnector: Option<Arc<dyn Reconnector>>,
pub keepalive: KeepaliveBehavior,
}
impl Default for ConnectOptions<'_> {
fn default() -> Self {
Self {
username: None,
password: None,
timeout: None,
reconnect: ReconnectBehavior::default(),
tls_config: None,
reconnector: None,
keepalive: KeepaliveBehavior::Enabled(crate::keepalive::KeepalivePolicy::default()),
}
}
}
impl core::fmt::Debug for ConnectOptions<'_> {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
f.debug_struct("ConnectOptions")
.field("username", &self.username)
.field(
"password",
&self.password.map(|_| "<redacted>").unwrap_or("None"),
)
.field("timeout", &self.timeout)
.field("reconnect", &self.reconnect)
.field("tls_config", &self.tls_config.as_ref().map(|_| "<set>"))
.field(
"reconnector",
&self.reconnector.as_ref().map(|_| "<custom>"),
)
.field("keepalive", &self.keepalive)
.finish()
}
}
pub async fn websocket_transport(
address: &str,
version: OcppVersion,
options: Option<ConnectOptions<'_>>,
) -> Result<(Box<dyn TransportSink>, Box<dyn TransportStream>), TransportError> {
let (stream, _protocol) = setup_socket(address, version.protocol(), options).await?;
Ok(crate::transport::websocket::split(stream))
}
pub enum NegotiatedClient {
#[cfg(feature = "ocpp_1_6")]
V1_6(crate::ocpp_1_6::OCPP1_6Client),
#[cfg(feature = "ocpp_2_0_1")]
V2_0_1(crate::ocpp_2_0_1::OCPP2_0_1Client),
#[cfg(feature = "ocpp_2_1")]
V2_1(crate::ocpp_2_1::OCPP2_1Client),
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum OcppVersion {
#[cfg(feature = "ocpp_1_6")]
V1_6,
#[cfg(feature = "ocpp_2_0_1")]
V2_0_1,
#[cfg(feature = "ocpp_2_1")]
V2_1,
}
impl OcppVersion {
fn protocol(self) -> &'static str {
match self {
#[cfg(feature = "ocpp_1_6")]
OcppVersion::V1_6 => "ocpp1.6",
#[cfg(feature = "ocpp_2_0_1")]
OcppVersion::V2_0_1 => "ocpp2.0.1",
#[cfg(feature = "ocpp_2_1")]
OcppVersion::V2_1 => "ocpp2.1",
}
}
#[allow(clippy::vec_init_then_push)]
fn all_compiled_in() -> Vec<OcppVersion> {
let mut versions = Vec::new();
#[cfg(feature = "ocpp_2_1")]
versions.push(OcppVersion::V2_1);
#[cfg(feature = "ocpp_2_0_1")]
versions.push(OcppVersion::V2_0_1);
#[cfg(feature = "ocpp_1_6")]
versions.push(OcppVersion::V1_6);
versions
}
}
pub async fn connect(
address: &str,
versions: Option<&[OcppVersion]>,
options: Option<ConnectOptions<'_>>,
) -> Result<NegotiatedClient, Box<dyn std::error::Error + Send + Sync>> {
let all_compiled_in;
let versions = match versions {
Some(versions) => versions,
None => {
all_compiled_in = OcppVersion::all_compiled_in();
&all_compiled_in
}
};
let offered = versions
.iter()
.map(|v| v.protocol())
.collect::<Vec<_>>()
.join(", ");
let (stream, negotiated) = setup_socket(address, &offered, options.clone()).await?;
let protocol = versions
.iter()
.find(|v| v.protocol() == negotiated)
.map(|v| v.protocol())
.ok_or_else(|| format!("Server negotiated unsupported protocol: {negotiated}"))?;
let config = prepare(address, protocol, options);
let (sink, source) = crate::transport::websocket::split(stream);
Ok(match protocol {
#[cfg(feature = "ocpp_1_6")]
"ocpp1.6" => NegotiatedClient::V1_6(crate::Client::from_transport_with_config(
sink,
source,
Box::new(crate::runtime::tokio::TokioExecutor),
Box::new(crate::runtime::tokio::TokioTimer),
config,
)),
#[cfg(feature = "ocpp_2_0_1")]
"ocpp2.0.1" => NegotiatedClient::V2_0_1(crate::Client::from_transport_with_config(
sink,
source,
Box::new(crate::runtime::tokio::TokioExecutor),
Box::new(crate::runtime::tokio::TokioTimer),
config,
)),
#[cfg(feature = "ocpp_2_1")]
"ocpp2.1" => NegotiatedClient::V2_1(crate::Client::from_transport_with_config(
sink,
source,
Box::new(crate::runtime::tokio::TokioExecutor),
Box::new(crate::runtime::tokio::TokioTimer),
config,
)),
_ => unreachable!("protocol only ever holds a value returned by OcppVersion::protocol"),
})
}
#[cfg(feature = "ocpp_1_6")]
pub async fn connect_1_6(
address: &str,
options: Option<ConnectOptions<'_>>,
) -> Result<crate::ocpp_1_6::OCPP1_6Client, Box<dyn std::error::Error + Send + Sync>> {
let config = prepare(address, "ocpp1.6", options.clone());
let (stream, _protocol) = setup_socket(address, "ocpp1.6", options).await?;
let (sink, source) = crate::transport::websocket::split(stream);
Ok(crate::Client::from_transport_with_config(
sink,
source,
Box::new(crate::runtime::tokio::TokioExecutor),
Box::new(crate::runtime::tokio::TokioTimer),
config,
))
}
#[cfg(feature = "ocpp_2_0_1")]
pub async fn connect_2_0_1(
address: &str,
options: Option<ConnectOptions<'_>>,
) -> Result<crate::ocpp_2_0_1::OCPP2_0_1Client, Box<dyn std::error::Error + Send + Sync>> {
let config = prepare(address, "ocpp2.0.1", options.clone());
let (stream, _protocol) = setup_socket(address, "ocpp2.0.1", options).await?;
let (sink, source) = crate::transport::websocket::split(stream);
Ok(crate::Client::from_transport_with_config(
sink,
source,
Box::new(crate::runtime::tokio::TokioExecutor),
Box::new(crate::runtime::tokio::TokioTimer),
config,
))
}
#[cfg(feature = "ocpp_2_1")]
pub async fn connect_2_1(
address: &str,
options: Option<ConnectOptions<'_>>,
) -> Result<crate::ocpp_2_1::OCPP2_1Client, Box<dyn std::error::Error + Send + Sync>> {
let config = prepare(address, "ocpp2.1", options.clone());
let (stream, _protocol) = setup_socket(address, "ocpp2.1", options).await?;
let (sink, source) = crate::transport::websocket::split(stream);
Ok(crate::Client::from_transport_with_config(
sink,
source,
Box::new(crate::runtime::tokio::TokioExecutor),
Box::new(crate::runtime::tokio::TokioTimer),
config,
))
}
fn prepare(
address: &str,
protocol: &'static str,
options: Option<ConnectOptions<'_>>,
) -> ClientConfig {
let defaults = ConnectOptions::default();
let timeout = options
.as_ref()
.and_then(|o| o.timeout)
.unwrap_or(DEFAULT_TIMEOUT);
let keepalive = options
.as_ref()
.map(|o| o.keepalive)
.unwrap_or(defaults.keepalive);
let reconnect = options.as_ref().map(|o| o.reconnect).unwrap_or_default();
let username = options
.as_ref()
.and_then(|o| o.username)
.map(str::to_string);
let password = options
.as_ref()
.and_then(|o| o.password)
.map(str::to_string);
let tls_config = options.as_ref().and_then(|o| o.tls_config.clone());
let custom = options.as_ref().and_then(|o| o.reconnector.clone());
let config = ClientConfig::new(timeout).with_keepalive(keepalive);
match reconnect {
ReconnectBehavior::Disabled => config,
ReconnectBehavior::Enabled(policy) if custom.is_some() => {
let custom = custom.expect("guarded by the match arm");
config.with_reconnect(Box::new(SharedReconnector(custom)), policy)
}
ReconnectBehavior::Enabled(policy) => config.with_reconnect(
Box::new(WebSocketReconnector {
address: address.to_string(),
protocol,
username,
password,
tls_config,
}),
policy,
),
}
}
struct SharedReconnector(Arc<dyn Reconnector>);
impl Reconnector for SharedReconnector {
fn connect<'a>(
&'a self,
) -> Pin<
Box<
dyn Future<
Output = Result<
(Box<dyn TransportSink>, Box<dyn TransportStream>),
TransportError,
>,
> + Send
+ 'a,
>,
> {
self.0.connect()
}
}
struct WebSocketReconnector {
address: String,
protocol: &'static str,
username: Option<String>,
password: Option<String>,
tls_config: Option<Arc<rustls::ClientConfig>>,
}
impl Reconnector for WebSocketReconnector {
fn connect<'a>(
&'a self,
) -> Pin<
Box<
dyn Future<
Output = Result<
(Box<dyn TransportSink>, Box<dyn TransportStream>),
TransportError,
>,
> + Send
+ 'a,
>,
> {
Box::pin(async move {
let options = ConnectOptions {
username: self.username.as_deref(),
password: self.password.as_deref(),
timeout: None,
reconnect: ReconnectBehavior::Disabled,
tls_config: self.tls_config.clone(),
reconnector: None,
keepalive: KeepaliveBehavior::Disabled,
};
let (stream, _protocol) =
setup_socket(&self.address, self.protocol, Some(options)).await?;
Ok(crate::transport::websocket::split(stream))
})
}
}
async fn setup_socket(
address: &str,
protocols: &str,
options: Option<ConnectOptions<'_>>,
) -> Result<
(WebSocketStream<MaybeTlsStream<TcpStream>>, String),
Box<dyn std::error::Error + Send + Sync>,
> {
let address = Url::parse(address)?;
let socket_addrs = address.socket_addrs(|| None)?;
let stream = TcpStream::connect(&*socket_addrs).await?;
let mut request: Request<()> = address.to_string().into_client_request()?;
request
.headers_mut()
.insert(SEC_WEBSOCKET_PROTOCOL, protocols.parse()?);
let mut tls_config = None;
if let Some(options) = options {
if let Some(username) = options.username {
let data = format!("{}:{}", username, options.password.unwrap_or(""));
let encoded = BASE64_STANDARD.encode(data);
request
.headers_mut()
.insert(AUTHORIZATION, format!("Basic {encoded}").parse()?);
}
tls_config = options.tls_config;
}
let connector = tls_config.map(Connector::Rustls);
let (stream, response) = client_async_tls_with_config(request, stream, None, connector).await?;
let protocol = response
.headers()
.get(SEC_WEBSOCKET_PROTOCOL)
.ok_or("No OCPP protocol negotiated")?;
Ok((stream, protocol.to_str()?.to_string()))
}