Preview of the proposed up-rust native-frame-model branch (up-rust b6b99c6d, up-spec f0e9b17) — not released documentation. branch · write-up

up_rust/communication/
in_memory_rpc_client.rs

1/********************************************************************************
2 * Copyright (c) 2024 Contributors to the Eclipse Foundation
3 *
4 * See the NOTICE file(s) distributed with this work for additional
5 * information regarding copyright ownership.
6 *
7 * This program and the accompanying materials are made available under the
8 * terms of the Apache License Version 2.0 which is available at
9 * https://www.apache.org/licenses/LICENSE-2.0
10 *
11 * SPDX-License-Identifier: Apache-2.0
12 ********************************************************************************/
13
14// [impl->dsn~communication-layer-impl-default~1]
15
16use std::collections::hash_map::Entry;
17use std::collections::HashMap;
18use std::sync::{Arc, Mutex};
19use std::time::Duration;
20
21use async_trait::async_trait;
22use tokio::sync::oneshot::{Receiver, Sender};
23use tokio::time::timeout;
24use tracing::{debug, info};
25
26use crate::{
27    communication::{
28        build_message, CallOptions, RegistrationError, RpcClient, ServiceInvocationError, UPayload,
29    },
30    LocalUriProvider, UCode, UListener, UMessage, UMessageBuilder, UStatus, UTransport, UUri, UUID,
31};
32
33/// Handles an RPC Response message received from the transport layer.
34fn handle_response_message(response: UMessage) -> Result<Option<UPayload>, ServiceInvocationError> {
35    match response.commstatus() {
36        Some(UCode::Ok) | None => {
37            // successful invocation
38            let payload_format = response
39                .payload_format()
40                .unwrap_or(crate::UPayloadFormat::Unspecified);
41            Ok(response
42                .payload()
43                .map(|payload| UPayload::new(payload, payload_format)))
44        }
45        Some(code) => {
46            // try to extract UStatus from response payload
47            let status = response.extract_protobuf().unwrap_or_else(|_e| {
48                UStatus::fail_with_code(code, "failed to invoke service operation")
49            });
50            Err(ServiceInvocationError::from(status))
51        }
52    }
53}
54
55struct ResponseListener {
56    // request ID -> sender for response message
57    pending_requests: Mutex<HashMap<UUID, Sender<UMessage>>>,
58}
59
60impl ResponseListener {
61    fn try_add_pending_request(
62        &self,
63        reqid: UUID,
64    ) -> Result<Receiver<UMessage>, ServiceInvocationError> {
65        let Ok(mut pending_requests) = self.pending_requests.lock() else {
66            return Err(ServiceInvocationError::Internal(
67                "failed to add response handler".to_string(),
68            ));
69        };
70
71        if let Entry::Vacant(entry) = pending_requests.entry(reqid) {
72            let (tx, rx) = tokio::sync::oneshot::channel();
73            entry.insert(tx);
74            Ok(rx)
75        } else {
76            Err(ServiceInvocationError::AlreadyExists(
77                "RPC request with given ID already pending".to_string(),
78            ))
79        }
80    }
81
82    fn handle_response(&self, response_message: UMessage) {
83        let reqid = response_message.request_id_unchecked().clone();
84        let response_sender = {
85            // drop lock as soon as possible
86            let Ok(mut pending_requests) = self.pending_requests.lock() else {
87                info!(
88                    request_id = reqid.to_hyphenated_string(),
89                    "failed to process response message, cannot acquire lock for pending requests map"
90                );
91                return;
92            };
93            pending_requests.remove(&reqid)
94        };
95        if let Some(sender) = response_sender {
96            if let Err(_e) = sender.send(response_message) {
97                // channel seems to be closed already
98                debug!(
99                    request_id = reqid.to_hyphenated_string(),
100                    "failed to deliver RPC Response message, channel already closed"
101                );
102            } else {
103                debug!(
104                    request_id = reqid.to_hyphenated_string(),
105                    "successfully delivered RPC Response message"
106                )
107            }
108        } else {
109            // we seem to have received a duplicate of the response message, ignoring it ...
110            debug!(
111                request_id = reqid.to_hyphenated_string(),
112                "ignoring (duplicate?) RPC Response message with unknown request ID"
113            );
114        }
115    }
116
117    fn remove_pending_request(&self, reqid: &UUID) -> Option<Sender<UMessage>> {
118        self.pending_requests
119            .lock()
120            .map_or(None, |mut pending_requests| pending_requests.remove(reqid))
121    }
122
123    #[cfg(test)]
124    fn contains(&self, reqid: &UUID) -> bool {
125        self.pending_requests
126            .lock()
127            .is_ok_and(|pending_requests| pending_requests.contains_key(reqid))
128    }
129}
130
131#[async_trait]
132impl UListener for ResponseListener {
133    async fn on_receive(&self, msg: UMessage) {
134        // it is sufficient to check if the message is a response
135        // because the transport implementation forwards valid UMessages only
136        if msg.is_response() {
137            self.handle_response(msg);
138        } else {
139            debug!(
140                message_type = msg.type_().to_cloudevent_type(),
141                "ignoring non-response message received by RPC client"
142            );
143        }
144    }
145}
146
147/// An [`RpcClient`] which keeps all information about pending requests in memory.
148///
149/// The client requires an implementations of [`UTransport`] for sending RPC Request messages
150/// to the service implementation and receiving its RPC Response messages.
151///
152/// During [startup](`Self::new`) the client registers a generic [`UListener`] with the transport
153/// for receiving all kinds of messages with a _sink_ address matching the client. The listener
154/// maintains an in-memory mapping of (pending) request IDs to response message handlers.
155///
156/// When an [`RPC call`](Self::invoke_method) is made, an RPC Request message is sent to the service
157/// implementation and a response handler is created and registered with the listener.
158/// When an RPC Response message arrives from the service, the corresponding handler is being looked
159/// up and invoked.
160pub struct InMemoryRpcClient<T, P> {
161    transport: Arc<T>,
162    uri_provider: Arc<P>,
163    response_listener: Arc<ResponseListener>,
164}
165
166impl<T, P> core::fmt::Debug for InMemoryRpcClient<T, P> {
167    fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
168        f.debug_struct("InMemoryRpcClient").finish_non_exhaustive()
169    }
170}
171
172impl<T: UTransport, P: LocalUriProvider> InMemoryRpcClient<T, P> {
173    /// Creates a new RPC client for a given transport.
174    ///
175    /// # Arguments
176    ///
177    /// * `transport` - The uProtocol Transport Layer implementation to use for invoking service operations.
178    /// * `uri_provider` - The helper for creating URIs that represent local resources.
179    ///
180    /// # Errors
181    ///
182    /// Returns an error if the generic RPC Response listener could not be
183    /// registered with the given transport.
184    pub async fn new(transport: Arc<T>, uri_provider: Arc<P>) -> Result<Self, RegistrationError> {
185        let response_listener = Arc::new(ResponseListener {
186            pending_requests: Mutex::new(HashMap::new()),
187        });
188        transport
189            .register_listener(
190                &UUri::any(),
191                Some(&uri_provider.get_source_uri()),
192                response_listener.clone(),
193            )
194            .await
195            .map_err(RegistrationError::from)?;
196
197        Ok(InMemoryRpcClient {
198            transport,
199            uri_provider,
200            response_listener,
201        })
202    }
203
204    #[cfg(test)]
205    fn contains_pending_request(&self, reqid: &UUID) -> bool {
206        self.response_listener.contains(reqid)
207    }
208}
209
210#[async_trait]
211impl<T: UTransport, P: LocalUriProvider> RpcClient for InMemoryRpcClient<T, P> {
212    async fn invoke_method(
213        &self,
214        method: UUri,
215        call_options: CallOptions,
216        payload: Option<UPayload>,
217    ) -> Result<Option<UPayload>, ServiceInvocationError> {
218        let message_id = call_options
219            .message_id()
220            .map_or_else(UUID::build, |id| id.to_owned());
221
222        let mut builder = UMessageBuilder::request(
223            method.clone(),
224            self.uri_provider.get_source_uri(),
225            call_options.ttl(),
226        );
227        builder.with_message_id(message_id.clone());
228        if let Some(token) = call_options.token() {
229            builder.with_token(token.to_owned());
230        }
231        if let Some(priority) = call_options.priority() {
232            builder.with_priority(priority);
233        }
234        let rpc_request_message = build_message(&mut builder, payload)
235            .map_err(|e| ServiceInvocationError::InvalidArgument(e.to_string()))?;
236
237        let receiver = self
238            .response_listener
239            .try_add_pending_request(message_id.clone())?;
240        self.transport
241            .send(rpc_request_message)
242            .await
243            .inspect_err(|_e| {
244                self.response_listener.remove_pending_request(&message_id);
245            })?;
246        debug!(
247            request_id = message_id.to_hyphenated_string(),
248            ttl = call_options.ttl(),
249            "successfully sent RPC Request message"
250        );
251
252        match timeout(Duration::from_millis(call_options.ttl() as u64), receiver).await {
253            Err(_) => {
254                debug!(
255                    request_id = message_id.to_hyphenated_string(),
256                    ttl = call_options.ttl(),
257                    "invocation of service operation has timed out"
258                );
259                self.response_listener.remove_pending_request(&message_id);
260                Err(ServiceInvocationError::DeadlineExceeded)
261            }
262            Ok(result) => match result {
263                Ok(response_message) => handle_response_message(response_message),
264                Err(_e) => {
265                    debug!(
266                        request_id = message_id.to_hyphenated_string(),
267                        "response listener failed to forward response message"
268                    );
269                    self.response_listener.remove_pending_request(&message_id);
270                    Err(ServiceInvocationError::Internal(
271                        "error receiving response message".to_string(),
272                    ))
273                }
274            },
275        }
276    }
277}
278
279#[cfg(test)]
280mod tests {
281
282    // [utest->dsn~communication-layer-impl-default~1]
283
284    use super::*;
285
286    use protobuf::well_known_types::wrappers::StringValue;
287    use tokio::{join, sync::Notify};
288
289    use crate::{utransport::MockTransport, StaticUriProvider, UMessageBuilder, UPriority, UUri};
290
291    fn new_uri_provider() -> Arc<StaticUriProvider> {
292        Arc::new(StaticUriProvider::new("", 0x0005, 0x02).expect("failed to create URI provider"))
293    }
294
295    fn service_method_uri() -> UUri {
296        UUri::try_from_parts("", 0x0001, 0x01, 0x1000).expect("failed to create service method URI")
297    }
298
299    #[tokio::test]
300    async fn test_registration_of_response_listener_fails() {
301        // GIVEN a transport
302        let mut mock_transport = MockTransport::default();
303        // with the maximum number of listeners already registered
304        mock_transport
305            .expect_do_register_listener()
306            .once()
307            .returning(|_source_filter, _sink_filter, _listener| {
308                Err(UStatus::fail_with_code(
309                    UCode::ResourceExhausted,
310                    "max number of listeners exceeded",
311                ))
312            });
313
314        // WHEN trying to create an RpcClient for the transport
315        let creation_attempt =
316            InMemoryRpcClient::new(Arc::new(mock_transport), new_uri_provider()).await;
317
318        // THEN the attempt fails with a MaxListenersExceeded error
319        assert!(
320            creation_attempt.is_err_and(|e| matches!(e, RegistrationError::MaxListenersExceeded))
321        );
322    }
323
324    #[tokio::test]
325    async fn test_invoke_method_fails_with_transport_error() {
326        // GIVEN an RPC client
327        let mut mock_transport = MockTransport::default();
328        mock_transport
329            .expect_do_register_listener()
330            .once()
331            .returning(|_source_filter, _sink_filter, _listener| Ok(()));
332        // with a transport that fails with an error when invoking a method
333        mock_transport
334            .expect_do_send()
335            .returning(|_request_message| {
336                Err(UStatus::fail_with_code(
337                    UCode::Unavailable,
338                    "transport not available",
339                ))
340            });
341        let client = InMemoryRpcClient::new(Arc::new(mock_transport), new_uri_provider())
342            .await
343            .unwrap();
344
345        // WHEN invoking a remote service operation
346        let message_id = UUID::build();
347        let call_options =
348            CallOptions::for_rpc_request(5_000, Some(message_id.clone()), None, None);
349        let response = client
350            .invoke_method(service_method_uri(), call_options, None)
351            .await;
352
353        // THEN the invocation fails with the error caused at the Transport Layer
354        assert!(response.is_err_and(|e| matches!(e, ServiceInvocationError::Unavailable(_msg))));
355        assert!(!client.contains_pending_request(&message_id));
356    }
357
358    #[tokio::test]
359    async fn test_invoke_method_succeeds() {
360        let message_id = UUID::build();
361        let call_options = CallOptions::for_rpc_request(
362            5_000,
363            Some(message_id.clone()),
364            Some("my_token".to_string()),
365            Some(crate::UPriority::CS6),
366        );
367
368        let (captured_listener_tx, captured_listener_rx) = tokio::sync::oneshot::channel();
369        let request_sent = Arc::new(Notify::new());
370        let request_sent_clone = request_sent.clone();
371
372        // GIVEN an RPC client
373        let mut mock_transport = MockTransport::default();
374        mock_transport
375            .expect_do_register_listener()
376            .once()
377            .return_once(move |_source_filter, _sink_filter, listener| {
378                captured_listener_tx.send(listener).map_err(|_e| {
379                    UStatus::fail_with_code(UCode::Internal, "cannot capture listener")
380                })
381            });
382        let expected_message_id = message_id.clone();
383        mock_transport
384            .expect_do_send()
385            .once()
386            .withf(move |request_message| {
387                request_message.id() == &expected_message_id
388                    && request_message.priority_unchecked() == UPriority::CS6
389                    && request_message.ttl_unchecked() == 5_000
390                    && request_message.token() == Some(&String::from("my_token"))
391            })
392            .returning(move |_request_message| {
393                request_sent_clone.notify_one();
394                Ok(())
395            });
396
397        let uri_provider = new_uri_provider();
398        let rpc_client = Arc::new(
399            InMemoryRpcClient::new(Arc::new(mock_transport), uri_provider.clone())
400                .await
401                .unwrap(),
402        );
403        let client: Arc<dyn RpcClient> = rpc_client.clone();
404
405        // WHEN invoking a remote service operation
406        let response_handle = tokio::spawn(async move {
407            let request_payload = StringValue {
408                value: "World".to_string(),
409                ..Default::default()
410            };
411            client
412                .invoke_proto_method::<_, StringValue>(
413                    service_method_uri(),
414                    call_options,
415                    request_payload,
416                )
417                .await
418        });
419
420        // AND the remote service sends the corresponding RPC Response message
421        let response_payload = StringValue {
422            value: "Hello World".to_string(),
423            ..Default::default()
424        };
425        let response_message = UMessageBuilder::response(
426            uri_provider.get_source_uri(),
427            message_id.clone(),
428            service_method_uri(),
429        )
430        .build_with_protobuf_payload(&response_payload)
431        .unwrap();
432
433        // wait for the RPC Request message having been sent
434        let (response_listener_result, _) = join!(captured_listener_rx, request_sent.notified());
435        let response_listener = response_listener_result.unwrap();
436
437        // send the RPC Response message which completes the request
438        let cloned_response_message = response_message.clone();
439        let cloned_response_listener = response_listener.clone();
440        tokio::spawn(async move {
441            cloned_response_listener
442                .on_receive(cloned_response_message)
443                .await
444        });
445
446        // THEN the response contains the expected payload
447        let response = response_handle.await.unwrap();
448        assert!(response.is_ok_and(|payload| payload.value == *"Hello World"));
449        assert!(!rpc_client.contains_pending_request(&message_id));
450
451        // AND if the remote service sends its response message again
452        response_listener.on_receive(response_message).await;
453        // the duplicate response is silently ignored
454        assert!(!rpc_client.contains_pending_request(&message_id));
455    }
456
457    #[tokio::test]
458    async fn test_invoke_method_fails_on_repeated_invocation() {
459        let message_id = UUID::build();
460        let first_request_sent = Arc::new(Notify::new());
461        let first_request_sent_clone = first_request_sent.clone();
462
463        // GIVEN an RPC client
464        let mut mock_transport = MockTransport::default();
465        mock_transport
466            .expect_do_register_listener()
467            .once()
468            .return_const(Ok(()));
469        let expected_message_id = message_id.clone();
470        mock_transport
471            .expect_do_send()
472            .once()
473            .withf(move |request_message| request_message.id() == &expected_message_id)
474            .returning(move |_request_message| {
475                first_request_sent_clone.notify_one();
476                Ok(())
477            });
478
479        let in_memory_rpc_client = Arc::new(
480            InMemoryRpcClient::new(Arc::new(mock_transport), new_uri_provider())
481                .await
482                .unwrap(),
483        );
484        let rpc_client: Arc<dyn RpcClient> = in_memory_rpc_client.clone();
485
486        // WHEN invoking a remote service operation
487        let call_options =
488            CallOptions::for_rpc_request(5_000, Some(message_id.clone()), None, None);
489        let cloned_call_options = call_options.clone();
490        let cloned_rpc_client = rpc_client.clone();
491
492        tokio::spawn(async move {
493            let request_payload = StringValue {
494                value: "World".to_string(),
495                ..Default::default()
496            };
497            cloned_rpc_client
498                .invoke_proto_method::<_, StringValue>(
499                    service_method_uri(),
500                    cloned_call_options,
501                    request_payload,
502                )
503                .await
504        });
505
506        // we wait for the first request message having been sent via the transport
507        // in order to be sure that the pending request has been added to the client's
508        // internal state
509        first_request_sent.notified().await;
510
511        // AND invoking the same operation before the response to the first request has arrived
512        let request_payload = StringValue {
513            value: "World".to_string(),
514            ..Default::default()
515        };
516        let second_request_handle = tokio::spawn(async move {
517            rpc_client
518                .invoke_proto_method::<_, StringValue>(
519                    service_method_uri(),
520                    call_options,
521                    request_payload,
522                )
523                .await
524        });
525
526        // THEN the second invocation fails
527        let response = second_request_handle.await.unwrap();
528        assert!(response.is_err_and(|e| matches!(e, ServiceInvocationError::AlreadyExists(_))));
529        // because there is a pending request for the message ID used in both requests
530        assert!(in_memory_rpc_client.contains_pending_request(&message_id));
531    }
532
533    #[tokio::test]
534    async fn test_invoke_method_fails_with_remote_error() {
535        let (captured_listener_tx, captured_listener_rx) = std::sync::mpsc::channel();
536
537        // GIVEN an RPC client
538        let mut mock_transport = MockTransport::default();
539        mock_transport.expect_do_register_listener().returning(
540            move |_source_filter, _sink_filter, listener| {
541                captured_listener_tx.send(listener).map_err(|_e| {
542                    UStatus::fail_with_code(UCode::Internal, "cannot capture listener")
543                })
544            },
545        );
546        // and a remote service operation that returns an error
547        mock_transport
548            .expect_do_send()
549            .returning(move |request_message| {
550                let error = UStatus::fail_with_code(UCode::NotFound, "no such object");
551                let response_message =
552                    UMessageBuilder::response_for_request(request_message.attributes())
553                        .with_comm_status(UCode::NotFound)
554                        .build_with_protobuf_payload(&error)
555                        .unwrap();
556                let captured_listener = captured_listener_rx.recv().unwrap().to_owned();
557                tokio::spawn(async move { captured_listener.on_receive(response_message).await });
558                Ok(())
559            });
560
561        let client = InMemoryRpcClient::new(Arc::new(mock_transport), new_uri_provider())
562            .await
563            .unwrap();
564
565        // WHEN invoking the remote service operation
566        let message_id = UUID::build();
567        let call_options =
568            CallOptions::for_rpc_request(5_000, Some(message_id.clone()), None, None);
569        let response = client
570            .invoke_method(service_method_uri(), call_options, None)
571            .await;
572
573        // THEN the invocation has failed with the error returned from the service
574        assert!(response.is_err_and(|e| { matches!(e, ServiceInvocationError::NotFound(_msg)) }));
575        assert!(!client.contains_pending_request(&message_id));
576    }
577
578    #[tokio::test]
579    async fn test_invoke_method_times_out() {
580        // GIVEN an RPC client
581        let mut mock_transport = MockTransport::default();
582        mock_transport
583            .expect_do_register_listener()
584            .returning(|_source_filter, _sink_filter, _listener| Ok(()));
585        // and a remote service operation that does not return a response
586        mock_transport
587            .expect_do_send()
588            .returning(|_request_message| Ok(()));
589
590        let client = InMemoryRpcClient::new(Arc::new(mock_transport), new_uri_provider())
591            .await
592            .unwrap();
593
594        // WHEN invoking the remote service operation
595        let message_id = UUID::build();
596        let call_options = CallOptions::for_rpc_request(20, Some(message_id.clone()), None, None);
597        let response = client
598            .invoke_method(service_method_uri(), call_options, None)
599            .await;
600
601        // THEN the invocation times out
602        assert!(response.is_err_and(|e| { matches!(e, ServiceInvocationError::DeadlineExceeded) }));
603        assert!(!client.contains_pending_request(&message_id));
604    }
605}