1use crate::action::{Action, SendAction};
2use crate::envelope::{
3 MESSAGE_TYPE_CALL, MESSAGE_TYPE_ERROR, MESSAGE_TYPE_RESULT, MESSAGE_TYPE_SEND, RawCall,
4 RawError, RawResult, RawSend,
5};
6use crate::error::{ClientError, ProtocolError};
7use crate::reconnect::{ReconnectPolicy, Reconnector};
8use crate::runtime::{Executor, Timer, with_timeout};
9use crate::sync::{BroadcastRegistry, Chan, OneShot, SharedMutex};
10use crate::transport::{TransportEvent, TransportSink, TransportStream};
11use alloc::borrow::ToOwned;
12use alloc::boxed::Box;
13use alloc::collections::{BTreeMap, VecDeque};
14use alloc::format;
15use alloc::string::{String, ToString};
16use alloc::sync::Arc;
17use core::future::Future;
18use core::time::Duration;
19use serde::Serialize;
20use serde::de::DeserializeOwned;
21use serde_json::Value;
22use uuid::Uuid;
23
24type PendingResponses<E> = Arc<SharedMutex<BTreeMap<Uuid, OneShot<Result<Value, E>>>>>;
25type RequestSenders = Arc<SharedMutex<BTreeMap<String, Chan<(String, Value)>>>>;
26type NotificationSenders = Arc<SharedMutex<BTreeMap<String, Chan<Value>>>>;
27type PongWaiters = Arc<SharedMutex<VecDeque<OneShot<()>>>>;
28
29pub struct Client<E: ProtocolError> {
33 sink: Arc<SharedMutex<Box<dyn TransportSink>>>,
34 pending_responses: PendingResponses<E>,
35 request_senders: RequestSenders,
36 notification_senders: NotificationSenders,
37 pong_waiters: PongWaiters,
38 ping_registry: Arc<BroadcastRegistry>,
39 reconnect_registry: Arc<BroadcastRegistry>,
40 executor: Arc<dyn Executor>,
41 timer: Arc<dyn Timer>,
42 timeout: Duration,
43}
44
45impl<E: ProtocolError> Clone for Client<E> {
46 fn clone(&self) -> Self {
47 Self {
48 sink: self.sink.clone(),
49 pending_responses: self.pending_responses.clone(),
50 request_senders: self.request_senders.clone(),
51 notification_senders: self.notification_senders.clone(),
52 pong_waiters: self.pong_waiters.clone(),
53 ping_registry: self.ping_registry.clone(),
54 reconnect_registry: self.reconnect_registry.clone(),
55 executor: self.executor.clone(),
56 timer: self.timer.clone(),
57 timeout: self.timeout,
58 }
59 }
60}
61
62impl<E: ProtocolError> Client<E> {
63 pub fn from_transport(
70 sink: Box<dyn TransportSink>,
71 stream: Box<dyn TransportStream>,
72 timeout: Duration,
73 executor: Box<dyn Executor>,
74 timer: Box<dyn Timer>,
75 ) -> Self {
76 Self::from_transport_with_reconnect(
77 sink,
78 stream,
79 timeout,
80 executor,
81 timer,
82 None,
83 ReconnectPolicy::default(),
84 )
85 }
86
87 pub fn from_transport_with_reconnect(
95 sink: Box<dyn TransportSink>,
96 mut stream: Box<dyn TransportStream>,
97 timeout: Duration,
98 executor: Box<dyn Executor>,
99 timer: Box<dyn Timer>,
100 reconnector: Option<Box<dyn Reconnector>>,
101 reconnect_policy: ReconnectPolicy,
102 ) -> Self {
103 let sink = Arc::new(SharedMutex::new(sink));
104 let pending_responses: PendingResponses<E> = Arc::new(SharedMutex::new(BTreeMap::new()));
105 let request_senders: RequestSenders = Arc::new(SharedMutex::new(BTreeMap::new()));
106 let notification_senders: NotificationSenders = Arc::new(SharedMutex::new(BTreeMap::new()));
107 let pong_waiters: PongWaiters = Arc::new(SharedMutex::new(VecDeque::new()));
108 let ping_registry = Arc::new(BroadcastRegistry::new());
109 let reconnect_registry = Arc::new(BroadcastRegistry::new());
110 let executor: Arc<dyn Executor> = Arc::from(executor);
111 let timer: Arc<dyn Timer> = Arc::from(timer);
112
113 let read_pending_responses = pending_responses.clone();
114 let read_request_senders = request_senders.clone();
115 let read_notification_senders = notification_senders.clone();
116 let read_pong_waiters = pong_waiters.clone();
117 let read_ping_registry = ping_registry.clone();
118 let read_reconnect_registry = reconnect_registry.clone();
119 let read_sink = sink.clone();
120 let read_timer = timer.clone();
121
122 executor.spawn(Box::pin(async move {
123 loop {
124 loop {
125 match stream.recv().await {
126 Ok(Some(TransportEvent::Frame(frame))) => {
127 handle_frame::<E>(
128 &frame,
129 &read_pending_responses,
130 &read_request_senders,
131 &read_notification_senders,
132 &read_sink,
133 )
134 .await;
135 }
136 Ok(Some(TransportEvent::Ping)) => {
137 read_ping_registry.notify_all().await;
138 let mut lock = read_sink.lock().await;
139 let _ = lock.pong().await;
140 }
141 Ok(Some(TransportEvent::Pong)) => {
142 let mut lock = read_pong_waiters.lock().await;
143 if let Some(waiter) = lock.pop_front() {
144 waiter.send(());
145 }
146 }
147 Ok(None) | Err(_) => break,
148 }
149 }
150
151 let Some(reconnector) = reconnector.as_ref() else {
152 break;
153 };
154
155 let mut attempt = 0u32;
156 loop {
157 match reconnector.connect().await {
158 Ok((new_sink, new_stream)) => {
159 *read_sink.lock().await = new_sink;
160 stream = new_stream;
161 tracing::info!(attempt, "ocpp-client: reconnected");
162 read_reconnect_registry.notify_all().await;
163 break;
164 }
165 Err(err) => {
166 tracing::warn!(attempt, error = %err, "ocpp-client: reconnect attempt failed");
167 read_timer.delay(reconnect_policy.delay_for(attempt)).await;
168 attempt = attempt.saturating_add(1);
169 }
170 }
171 }
172 }
173 }));
174
175 Self {
176 sink,
177 pending_responses,
178 request_senders,
179 notification_senders,
180 pong_waiters,
181 ping_registry,
182 reconnect_registry,
183 executor,
184 timer,
185 timeout,
186 }
187 }
188
189 pub async fn call<A: Action>(
191 &self,
192 request: A::Request,
193 ) -> Result<A::Response, ClientError<E>> {
194 let response = self.do_send_request(request, A::NAME).await?;
195 Ok(response)
196 }
197
198 pub async fn on<A, F, FF>(&self, mut callback: F)
201 where
202 A: Action,
203 F: FnMut(A::Request, Self) -> FF + Send + Sync + 'static,
204 FF: Future<Output = Result<A::Response, E>> + Send,
205 {
206 let chan: Chan<(String, Value)> = Chan::new();
207 {
208 let mut lock = self.request_senders.lock().await;
209 lock.insert(A::NAME.to_string(), chan.clone());
210 }
211
212 let client = self.clone();
213 self.executor.spawn(Box::pin(async move {
214 loop {
215 let (message_id, payload) = chan.recv().await;
216 match serde_json::from_value::<A::Request>(payload) {
217 Ok(request) => {
218 let response = callback(request, client.clone()).await;
219 client.do_send_response(response, &message_id).await;
220 }
221 Err(_) => {
222 let error =
223 E::not_implemented(&format!("Failed to parse payload for {}", A::NAME));
224 client
225 .do_send_response::<A::Response>(Err(error), &message_id)
226 .await;
227 }
228 }
229 }
230 }));
231 }
232
233 #[cfg(feature = "test")]
236 pub async fn wait_for<A, F, FF>(&self, mut callback: F) -> Result<A::Request, ClientError<E>>
237 where
238 A: Action,
239 F: FnMut(A::Request, Self) -> FF + Send + Sync + 'static,
240 FF: Future<Output = Result<A::Response, E>> + Send,
241 {
242 let chan: Chan<(String, Value)> = Chan::new();
243 {
244 let mut lock = self.request_senders.lock().await;
245 lock.insert(A::NAME.to_string(), chan.clone());
246 }
247
248 match with_timeout(self.timer.as_ref(), self.timeout, chan.recv()).await {
249 Ok((message_id, payload)) => {
250 let for_callback: A::Request =
251 serde_json::from_value(payload.clone()).map_err(ClientError::Decode)?;
252 let response = callback(for_callback, self.clone()).await;
253 self.do_send_response(response, &message_id).await;
254 serde_json::from_value(payload).map_err(ClientError::Decode)
255 }
256 Err(_) => Err(ClientError::Timeout),
257 }
258 }
259
260 pub async fn send_notification<A: SendAction>(
264 &self,
265 payload: A::Payload,
266 ) -> Result<(), ClientError<E>> {
267 let message_id = Uuid::new_v4();
268 let payload = serde_json::to_value(&payload).map_err(ClientError::Decode)?;
269 let send = RawSend(
270 MESSAGE_TYPE_SEND,
271 message_id.to_string(),
272 A::NAME.to_string(),
273 payload,
274 );
275 let frame = serde_json::to_string(&send).map_err(ClientError::Decode)?;
276
277 let mut lock = self.sink.lock().await;
278 lock.send(frame).await.map_err(ClientError::Transport)
279 }
280
281 pub async fn on_notification<A, F, FF>(&self, mut callback: F)
286 where
287 A: SendAction,
288 F: FnMut(A::Payload, Self) -> FF + Send + Sync + 'static,
289 FF: Future<Output = ()> + Send,
290 {
291 let chan: Chan<Value> = Chan::new();
292 {
293 let mut lock = self.notification_senders.lock().await;
294 lock.insert(A::NAME.to_string(), chan.clone());
295 }
296
297 let client = self.clone();
298 self.executor.spawn(Box::pin(async move {
299 loop {
300 let payload = chan.recv().await;
301 match serde_json::from_value::<A::Payload>(payload) {
302 Ok(payload) => callback(payload, client.clone()).await,
303 Err(err) => {
304 tracing::warn!(error = %err, action = A::NAME, "ocpp-client: failed to parse SEND payload");
305 }
306 }
307 }
308 }));
309 }
310
311 pub async fn send_ping(&self) -> Result<(), ClientError<E>> {
312 let waiter = OneShot::new();
313 {
314 let mut lock = self.pong_waiters.lock().await;
315 lock.push_back(waiter.clone());
316 }
317 {
318 let mut lock = self.sink.lock().await;
319 lock.ping().await.map_err(ClientError::Transport)?;
320 }
321 with_timeout(self.timer.as_ref(), self.timeout, waiter.wait())
322 .await
323 .map(|_| ())
324 .map_err(|_| ClientError::Timeout)
325 }
326
327 pub async fn on_ping<
328 F: FnMut(Self) -> FF + Send + Sync + 'static,
329 FF: Future<Output = ()> + Send,
330 >(
331 &self,
332 mut callback: F,
333 ) {
334 let signal = self.ping_registry.subscribe().await;
335 let client = self.clone();
336 self.executor.spawn(Box::pin(async move {
337 loop {
338 signal.wait().await;
339 callback(client.clone()).await;
340 }
341 }));
342 }
343
344 pub async fn on_reconnect<
353 F: FnMut(Self) -> FF + Send + Sync + 'static,
354 FF: Future<Output = ()> + Send,
355 >(
356 &self,
357 mut callback: F,
358 ) {
359 let signal = self.reconnect_registry.subscribe().await;
360 let client = self.clone();
361 self.executor.spawn(Box::pin(async move {
362 loop {
363 signal.wait().await;
364 callback(client.clone()).await;
365 }
366 }));
367 }
368
369 pub async fn disconnect(&self) -> Result<(), ClientError<E>> {
370 let mut lock = self.sink.lock().await;
371 lock.close().await.map_err(ClientError::Transport)
372 }
373
374 async fn do_send_request<P: Serialize, R: DeserializeOwned>(
375 &self,
376 request: P,
377 action: &str,
378 ) -> Result<R, ClientError<E>> {
379 let message_id = Uuid::new_v4();
380 let payload = serde_json::to_value(&request).map_err(ClientError::Decode)?;
381 let call = RawCall(
382 MESSAGE_TYPE_CALL,
383 message_id.to_string(),
384 action.to_string(),
385 payload,
386 );
387 let frame = serde_json::to_string(&call).map_err(ClientError::Decode)?;
388
389 let waiter = OneShot::new();
390 {
391 let mut lock = self.pending_responses.lock().await;
392 lock.insert(message_id, waiter.clone());
393 }
394
395 {
396 let mut lock = self.sink.lock().await;
397 lock.send(frame).await.map_err(ClientError::Transport)?;
398 }
399
400 let result = with_timeout(self.timer.as_ref(), self.timeout, waiter.wait())
401 .await
402 .map_err(|_| ClientError::Timeout)?;
403
404 match result {
405 Ok(value) => serde_json::from_value(value).map_err(ClientError::Decode),
406 Err(e) => Err(ClientError::Protocol(e)),
407 }
408 }
409
410 async fn do_send_response<R: Serialize>(&self, response: Result<R, E>, message_id: &str) {
411 let frame = match response {
412 Ok(r) => match serde_json::to_value(r) {
413 Ok(value) => serde_json::to_string(&RawResult(
414 MESSAGE_TYPE_RESULT,
415 message_id.to_string(),
416 value,
417 )),
418 Err(e) => return log_send_error(e),
419 },
420 Err(e) => serde_json::to_string(&RawError(
421 MESSAGE_TYPE_ERROR,
422 message_id.to_string(),
423 e.code().to_string(),
424 e.description().to_string(),
425 e.details().to_owned(),
426 )),
427 };
428
429 match frame {
430 Ok(frame) => {
431 let mut lock = self.sink.lock().await;
432 if let Err(err) = lock.send(frame).await {
433 tracing::warn!(error = %err, "ocpp-client: failed to send response");
434 }
435 }
436 Err(err) => {
437 tracing::error!(error = %err, "ocpp-client: failed to encode response");
438 }
439 }
440 }
441}
442
443fn log_send_error(err: serde_json::Error) {
444 tracing::error!(error = %err, "ocpp-client: failed to encode response payload");
445}
446
447async fn handle_frame<E: ProtocolError>(
448 frame: &str,
449 pending_responses: &PendingResponses<E>,
450 request_senders: &RequestSenders,
451 notification_senders: &NotificationSenders,
452 sink: &Arc<SharedMutex<Box<dyn TransportSink>>>,
453) {
454 let value: Value = match serde_json::from_str(frame) {
455 Ok(v) => v,
456 Err(err) => {
457 tracing::warn!(error = %err, "ocpp-client: received malformed frame");
458 return;
459 }
460 };
461
462 let Value::Array(items) = value else {
463 tracing::warn!("ocpp-client: a message should be a JSON array");
464 return;
465 };
466 let Some(Value::Number(message_type)) = items.first() else {
467 tracing::warn!("ocpp-client: missing message type id");
468 return;
469 };
470 let Some(message_type) = message_type.as_u64() else {
471 tracing::warn!("ocpp-client: message type id must be an integer");
472 return;
473 };
474
475 match message_type {
476 MESSAGE_TYPE_CALL => {
477 let call: RawCall = match serde_json::from_str(frame) {
478 Ok(c) => c,
479 Err(err) => {
480 tracing::warn!(error = %err, "ocpp-client: failed to parse CALL");
481 return;
482 }
483 };
484 let action = &call.2;
485 let sender = {
486 let lock = request_senders.lock().await;
487 lock.get(action).cloned()
488 };
489 match sender {
490 Some(sender) => {
491 sender.send((call.1, call.3)).await;
492 }
493 None => {
494 let error =
495 E::not_implemented(&format!("Action '{action}' is not implemented"));
496 let payload = RawError(
497 MESSAGE_TYPE_ERROR,
498 call.1,
499 error.code().to_string(),
500 error.description().to_string(),
501 error.details().to_owned(),
502 );
503 if let Ok(frame) = serde_json::to_string(&payload) {
504 let mut lock = sink.lock().await;
505 let _ = lock.send(frame).await;
506 }
507 }
508 }
509 }
510 MESSAGE_TYPE_RESULT => {
511 let result: RawResult = match serde_json::from_str(frame) {
512 Ok(r) => r,
513 Err(err) => {
514 tracing::warn!(error = %err, "ocpp-client: failed to parse CALLRESULT");
515 return;
516 }
517 };
518 let Ok(id) = Uuid::parse_str(&result.1) else {
519 return;
520 };
521 let mut lock = pending_responses.lock().await;
522 if let Some(sender) = lock.remove(&id) {
523 sender.send(Ok(result.2));
524 }
525 }
526 MESSAGE_TYPE_ERROR => {
527 let error: RawError = match serde_json::from_str(frame) {
528 Ok(e) => e,
529 Err(err) => {
530 tracing::warn!(error = %err, "ocpp-client: failed to parse CALLERROR");
531 return;
532 }
533 };
534 let Ok(id) = Uuid::parse_str(&error.1) else {
535 return;
536 };
537 let mut lock = pending_responses.lock().await;
538 if let Some(sender) = lock.remove(&id) {
539 sender.send(Err(E::from_wire(&error.2, &error.3, error.4)));
540 }
541 }
542 MESSAGE_TYPE_SEND => {
543 let send: RawSend = match serde_json::from_str(frame) {
544 Ok(s) => s,
545 Err(err) => {
546 tracing::warn!(error = %err, "ocpp-client: failed to parse SEND");
547 return;
548 }
549 };
550 let action = &send.2;
551 let sender = {
552 let lock = notification_senders.lock().await;
553 lock.get(action).cloned()
554 };
555 match sender {
556 Some(sender) => sender.send(send.3).await,
557 None => {
558 tracing::warn!(action = %action, "ocpp-client: SEND for unhandled action");
559 }
560 }
561 }
562 other => {
563 tracing::warn!(message_type = other, "ocpp-client: unknown message type id");
564 }
565 }
566}