1use crate::client::ClientConfig;
2use crate::keepalive::KeepaliveBehavior;
3use crate::reconnect::{ReconnectBehavior, Reconnector};
4use crate::transport::{TransportError, TransportSink, TransportStream};
5use base64::Engine;
6use base64::prelude::BASE64_STANDARD;
7use core::future::Future;
8use core::pin::Pin;
9use std::sync::Arc;
10use std::time::Duration;
11use tokio::net::TcpStream;
12use tokio_tungstenite::tungstenite::client::IntoClientRequest;
13use tokio_tungstenite::tungstenite::http::Request;
14use tokio_tungstenite::tungstenite::http::header::{AUTHORIZATION, SEC_WEBSOCKET_PROTOCOL};
15use tokio_tungstenite::{Connector, MaybeTlsStream, WebSocketStream, client_async_tls_with_config};
16use url::Url;
17
18const DEFAULT_TIMEOUT: Duration = Duration::from_secs(5);
19
20#[derive(Clone)]
21pub struct ConnectOptions<'a> {
22 pub username: Option<&'a str>,
23 pub password: Option<&'a str>,
24 pub timeout: Option<Duration>,
25 pub reconnect: ReconnectBehavior,
29 pub tls_config: Option<Arc<rustls::ClientConfig>>,
38 pub reconnector: Option<Arc<dyn Reconnector>>,
57 pub keepalive: KeepaliveBehavior,
72}
73
74impl Default for ConnectOptions<'_> {
78 fn default() -> Self {
79 Self {
80 username: None,
81 password: None,
82 timeout: None,
83 reconnect: ReconnectBehavior::default(),
84 tls_config: None,
85 reconnector: None,
86 keepalive: KeepaliveBehavior::Enabled(crate::keepalive::KeepalivePolicy::default()),
87 }
88 }
89}
90
91impl core::fmt::Debug for ConnectOptions<'_> {
92 fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
95 f.debug_struct("ConnectOptions")
96 .field("username", &self.username)
97 .field(
98 "password",
99 &self.password.map(|_| "<redacted>").unwrap_or("None"),
100 )
101 .field("timeout", &self.timeout)
102 .field("reconnect", &self.reconnect)
103 .field("tls_config", &self.tls_config.as_ref().map(|_| "<set>"))
104 .field(
105 "reconnector",
106 &self.reconnector.as_ref().map(|_| "<custom>"),
107 )
108 .field("keepalive", &self.keepalive)
109 .finish()
110 }
111}
112
113pub async fn websocket_transport(
139 address: &str,
140 version: OcppVersion,
141 options: Option<ConnectOptions<'_>>,
142) -> Result<(Box<dyn TransportSink>, Box<dyn TransportStream>), TransportError> {
143 let (stream, _protocol) = setup_socket(address, version.protocol(), options).await?;
144 Ok(crate::transport::websocket::split(stream))
145}
146
147pub enum NegotiatedClient {
151 #[cfg(feature = "ocpp_1_6")]
152 V1_6(crate::ocpp_1_6::OCPP1_6Client),
153 #[cfg(feature = "ocpp_2_0_1")]
154 V2_0_1(crate::ocpp_2_0_1::OCPP2_0_1Client),
155 #[cfg(feature = "ocpp_2_1")]
156 V2_1(crate::ocpp_2_1::OCPP2_1Client),
157}
158
159#[derive(Debug, Clone, Copy, PartialEq, Eq)]
162pub enum OcppVersion {
163 #[cfg(feature = "ocpp_1_6")]
164 V1_6,
165 #[cfg(feature = "ocpp_2_0_1")]
166 V2_0_1,
167 #[cfg(feature = "ocpp_2_1")]
168 V2_1,
169}
170
171impl OcppVersion {
172 fn protocol(self) -> &'static str {
173 match self {
174 #[cfg(feature = "ocpp_1_6")]
175 OcppVersion::V1_6 => "ocpp1.6",
176 #[cfg(feature = "ocpp_2_0_1")]
177 OcppVersion::V2_0_1 => "ocpp2.0.1",
178 #[cfg(feature = "ocpp_2_1")]
179 OcppVersion::V2_1 => "ocpp2.1",
180 }
181 }
182
183 #[allow(clippy::vec_init_then_push)]
186 fn all_compiled_in() -> Vec<OcppVersion> {
187 let mut versions = Vec::new();
188 #[cfg(feature = "ocpp_2_1")]
189 versions.push(OcppVersion::V2_1);
190 #[cfg(feature = "ocpp_2_0_1")]
191 versions.push(OcppVersion::V2_0_1);
192 #[cfg(feature = "ocpp_1_6")]
193 versions.push(OcppVersion::V1_6);
194 versions
195 }
196}
197
198pub async fn connect(
206 address: &str,
207 versions: Option<&[OcppVersion]>,
208 options: Option<ConnectOptions<'_>>,
209) -> Result<NegotiatedClient, Box<dyn std::error::Error + Send + Sync>> {
210 let all_compiled_in;
211 let versions = match versions {
212 Some(versions) => versions,
213 None => {
214 all_compiled_in = OcppVersion::all_compiled_in();
215 &all_compiled_in
216 }
217 };
218 let offered = versions
219 .iter()
220 .map(|v| v.protocol())
221 .collect::<Vec<_>>()
222 .join(", ");
223 let (stream, negotiated) = setup_socket(address, &offered, options.clone()).await?;
224 let protocol = versions
225 .iter()
226 .find(|v| v.protocol() == negotiated)
227 .map(|v| v.protocol())
228 .ok_or_else(|| format!("Server negotiated unsupported protocol: {negotiated}"))?;
229 let config = prepare(address, protocol, options);
230 let (sink, source) = crate::transport::websocket::split(stream);
231
232 Ok(match protocol {
233 #[cfg(feature = "ocpp_1_6")]
234 "ocpp1.6" => NegotiatedClient::V1_6(crate::Client::from_transport_with_config(
235 sink,
236 source,
237 Box::new(crate::runtime::tokio::TokioExecutor),
238 Box::new(crate::runtime::tokio::TokioTimer),
239 config,
240 )),
241 #[cfg(feature = "ocpp_2_0_1")]
242 "ocpp2.0.1" => NegotiatedClient::V2_0_1(crate::Client::from_transport_with_config(
243 sink,
244 source,
245 Box::new(crate::runtime::tokio::TokioExecutor),
246 Box::new(crate::runtime::tokio::TokioTimer),
247 config,
248 )),
249 #[cfg(feature = "ocpp_2_1")]
250 "ocpp2.1" => NegotiatedClient::V2_1(crate::Client::from_transport_with_config(
251 sink,
252 source,
253 Box::new(crate::runtime::tokio::TokioExecutor),
254 Box::new(crate::runtime::tokio::TokioTimer),
255 config,
256 )),
257 _ => unreachable!("protocol only ever holds a value returned by OcppVersion::protocol"),
258 })
259}
260
261#[cfg(feature = "ocpp_1_6")]
263pub async fn connect_1_6(
264 address: &str,
265 options: Option<ConnectOptions<'_>>,
266) -> Result<crate::ocpp_1_6::OCPP1_6Client, Box<dyn std::error::Error + Send + Sync>> {
267 let config = prepare(address, "ocpp1.6", options.clone());
268 let (stream, _protocol) = setup_socket(address, "ocpp1.6", options).await?;
269 let (sink, source) = crate::transport::websocket::split(stream);
270 Ok(crate::Client::from_transport_with_config(
271 sink,
272 source,
273 Box::new(crate::runtime::tokio::TokioExecutor),
274 Box::new(crate::runtime::tokio::TokioTimer),
275 config,
276 ))
277}
278
279#[cfg(feature = "ocpp_2_0_1")]
281pub async fn connect_2_0_1(
282 address: &str,
283 options: Option<ConnectOptions<'_>>,
284) -> Result<crate::ocpp_2_0_1::OCPP2_0_1Client, Box<dyn std::error::Error + Send + Sync>> {
285 let config = prepare(address, "ocpp2.0.1", options.clone());
286 let (stream, _protocol) = setup_socket(address, "ocpp2.0.1", options).await?;
287 let (sink, source) = crate::transport::websocket::split(stream);
288 Ok(crate::Client::from_transport_with_config(
289 sink,
290 source,
291 Box::new(crate::runtime::tokio::TokioExecutor),
292 Box::new(crate::runtime::tokio::TokioTimer),
293 config,
294 ))
295}
296
297#[cfg(feature = "ocpp_2_1")]
299pub async fn connect_2_1(
300 address: &str,
301 options: Option<ConnectOptions<'_>>,
302) -> Result<crate::ocpp_2_1::OCPP2_1Client, Box<dyn std::error::Error + Send + Sync>> {
303 let config = prepare(address, "ocpp2.1", options.clone());
304 let (stream, _protocol) = setup_socket(address, "ocpp2.1", options).await?;
305 let (sink, source) = crate::transport::websocket::split(stream);
306 Ok(crate::Client::from_transport_with_config(
307 sink,
308 source,
309 Box::new(crate::runtime::tokio::TokioExecutor),
310 Box::new(crate::runtime::tokio::TokioTimer),
311 config,
312 ))
313}
314
315fn prepare(
323 address: &str,
324 protocol: &'static str,
325 options: Option<ConnectOptions<'_>>,
326) -> ClientConfig {
327 let defaults = ConnectOptions::default();
328 let timeout = options
329 .as_ref()
330 .and_then(|o| o.timeout)
331 .unwrap_or(DEFAULT_TIMEOUT);
332 let keepalive = options
333 .as_ref()
334 .map(|o| o.keepalive)
335 .unwrap_or(defaults.keepalive);
336 let reconnect = options.as_ref().map(|o| o.reconnect).unwrap_or_default();
337 let username = options
338 .as_ref()
339 .and_then(|o| o.username)
340 .map(str::to_string);
341 let password = options
342 .as_ref()
343 .and_then(|o| o.password)
344 .map(str::to_string);
345 let tls_config = options.as_ref().and_then(|o| o.tls_config.clone());
346
347 let custom = options.as_ref().and_then(|o| o.reconnector.clone());
348
349 let config = ClientConfig::new(timeout).with_keepalive(keepalive);
350
351 match reconnect {
352 ReconnectBehavior::Disabled => config,
353 ReconnectBehavior::Enabled(policy) if custom.is_some() => {
354 let custom = custom.expect("guarded by the match arm");
355 config.with_reconnect(Box::new(SharedReconnector(custom)), policy)
356 }
357 ReconnectBehavior::Enabled(policy) => config.with_reconnect(
358 Box::new(WebSocketReconnector {
359 address: address.to_string(),
360 protocol,
361 username,
362 password,
363 tls_config,
364 }),
365 policy,
366 ),
367 }
368}
369
370struct SharedReconnector(Arc<dyn Reconnector>);
374
375impl Reconnector for SharedReconnector {
376 fn connect<'a>(
377 &'a self,
378 ) -> Pin<
379 Box<
380 dyn Future<
381 Output = Result<
382 (Box<dyn TransportSink>, Box<dyn TransportStream>),
383 TransportError,
384 >,
385 > + Send
386 + 'a,
387 >,
388 > {
389 self.0.connect()
390 }
391}
392
393struct WebSocketReconnector {
396 address: String,
397 protocol: &'static str,
398 username: Option<String>,
399 password: Option<String>,
400 tls_config: Option<Arc<rustls::ClientConfig>>,
401}
402
403impl Reconnector for WebSocketReconnector {
404 fn connect<'a>(
405 &'a self,
406 ) -> Pin<
407 Box<
408 dyn Future<
409 Output = Result<
410 (Box<dyn TransportSink>, Box<dyn TransportStream>),
411 TransportError,
412 >,
413 > + Send
414 + 'a,
415 >,
416 > {
417 Box::pin(async move {
418 let options = ConnectOptions {
419 username: self.username.as_deref(),
420 password: self.password.as_deref(),
421 timeout: None,
422 reconnect: ReconnectBehavior::Disabled,
423 tls_config: self.tls_config.clone(),
424 reconnector: None,
425 keepalive: KeepaliveBehavior::Disabled,
429 };
430 let (stream, _protocol) =
431 setup_socket(&self.address, self.protocol, Some(options)).await?;
432 Ok(crate::transport::websocket::split(stream))
433 })
434 }
435}
436
437async fn setup_socket(
438 address: &str,
439 protocols: &str,
440 options: Option<ConnectOptions<'_>>,
441) -> Result<
442 (WebSocketStream<MaybeTlsStream<TcpStream>>, String),
443 Box<dyn std::error::Error + Send + Sync>,
444> {
445 let address = Url::parse(address)?;
446
447 let socket_addrs = address.socket_addrs(|| None)?;
448 let stream = TcpStream::connect(&*socket_addrs).await?;
449
450 let mut request: Request<()> = address.to_string().into_client_request()?;
451 request
452 .headers_mut()
453 .insert(SEC_WEBSOCKET_PROTOCOL, protocols.parse()?);
454 let mut tls_config = None;
455 if let Some(options) = options {
456 if let Some(username) = options.username {
457 let data = format!("{}:{}", username, options.password.unwrap_or(""));
458 let encoded = BASE64_STANDARD.encode(data);
459 request
460 .headers_mut()
461 .insert(AUTHORIZATION, format!("Basic {encoded}").parse()?);
462 }
463 tls_config = options.tls_config;
464 }
465
466 let connector = tls_config.map(Connector::Rustls);
467 let (stream, response) = client_async_tls_with_config(request, stream, None, connector).await?;
468
469 let protocol = response
470 .headers()
471 .get(SEC_WEBSOCKET_PROTOCOL)
472 .ok_or("No OCPP protocol negotiated")?;
473
474 Ok((stream, protocol.to_str()?.to_string()))
475}