Skip to main content

ocpp_client/
client.rs

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
29/// The OCPP client engine, generic over one version's protocol error type. `OCPP1_6Client`
30/// and `OCPP2_0_1Client` are just `Client<OCPP1_6Error>` / `Client<OCPP2_0_1Error>` - the
31/// dispatch/timeout/error machinery below is written once and shared by every version.
32pub 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    /// Build a client over any transport - the WebSocket adapter used by `connect_1_6` is
64    /// just one implementation of `TransportSink`/`TransportStream`; tests and non-WebSocket
65    /// transports (an embedded framed link, an in-memory fake for unit tests) construct a
66    /// client the same way. `executor`/`timer` are likewise pluggable: the `tokio-runtime`
67    /// feature provides `TokioExecutor`/`TokioTimer`; embedded users supply their own (e.g.
68    /// backed by `embassy-executor`/`embassy-time`).
69    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    /// Same as [`Client::from_transport`], but with automatic reconnect: when the transport
88    /// closes (`TransportStream::recv` returns `Ok(None)`/`Err(_)`), the background read loop
89    /// calls `reconnector.connect()` (backing off per `reconnect_policy` between failed
90    /// attempts) instead of exiting, and swaps in the new transport once one succeeds.
91    /// `reconnector: None` reproduces `from_transport`'s behavior - the read loop exits on
92    /// disconnect and the client goes quiet. `connect_1_6`/`connect_2_0_1`/`connect_2_1` use
93    /// this constructor with a WebSocket-backed `Reconnector`.
94    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    /// Send a CALL for `A` and wait for the matching CALLRESULT/CALLERROR.
190    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    /// Register a handler for CALLs the other side sends for action `A`. Replaces any
199    /// previously registered handler for the same action.
200    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    /// Wait for exactly one CALL for action `A` (bounded by the client's timeout), answer
234    /// it with `callback`, and return the parsed request. Only useful in tests.
235    #[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    /// Send a `SEND` (OCPP-J 2.1 only) fire-and-forget message: writes the frame and returns as
261    /// soon as the transport accepts it - no waiter, no timeout, since the spec forbids the
262    /// receiver from ever replying to a `SEND`.
263    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    /// Register a handler for `SEND` (OCPP-J 2.1 only) messages of action `A`. Unlike
282    /// [`Client::on`], `callback` returns nothing - the spec forbids replying to a `SEND`, so
283    /// there's no response to send back. Replaces any previously registered handler for the
284    /// same action.
285    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    /// Register a callback that fires every time the background read loop redials
345    /// successfully after a disconnect (see [`Client::from_transport_with_reconnect`]). Never
346    /// fires for the initial connection, only for later reconnects - the initial `Client` is
347    /// already handed back post-connect, so callers run their own post-connect setup (e.g.
348    /// `BootNotification`) right after `connect_1_6`/`from_transport_with_reconnect` returns.
349    /// This is the hook for redoing that setup (or resyncing any other session state) after a
350    /// dropped-and-restored connection; this crate does not re-run `BootNotification` or replay
351    /// any state on its own.
352    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}