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_server.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;
19use std::time::Duration;
20
21use async_trait::async_trait;
22use tracing::{debug, info};
23
24use crate::{
25    communication::{
26        build_message, RegistrationError, RequestHandler, RpcServer, ServiceInvocationError,
27        UPayload,
28    },
29    LocalUriProvider, UListener, UMessage, UMessageBuilder, UStatus, UTransport, UUri,
30};
31
32struct RequestListener<T: UTransport> {
33    request_handler: Arc<dyn RequestHandler>,
34    transport: Arc<T>,
35}
36
37impl<T: UTransport> RequestListener<T> {
38    async fn process_valid_request(&self, resource_id: u16, request_message: UMessage) {
39        let transport_clone = self.transport.clone();
40        let request_handler_clone = self.request_handler.clone();
41        let mut response_builder =
42            UMessageBuilder::response_for_request(request_message.attributes());
43
44        let request_message_id = request_message.id().to_hyphenated_string();
45        let request_timeout = request_message.ttl_unchecked();
46        let payload_format = request_message
47            .payload_format()
48            .unwrap_or(crate::UPayloadFormat::Unspecified);
49        let request_payload = request_message
50            .payload()
51            .map(|data| UPayload::new(data, payload_format));
52
53        debug!(
54            ttl = request_timeout,
55            id = request_message_id,
56            resource_id = resource_id,
57            "processing RPC request"
58        );
59
60        let invocation_result_future = request_handler_clone.handle_request(
61            resource_id,
62            request_message.attributes(),
63            request_payload,
64        );
65        let outcome = tokio::time::timeout(
66            Duration::from_millis(request_timeout as u64),
67            invocation_result_future,
68        )
69        .await
70        .map_err(|_e| {
71            info!(ttl = request_timeout, "request handler timed out");
72            ServiceInvocationError::DeadlineExceeded
73        })
74        .and_then(|v| v);
75
76        let response = match outcome {
77            Ok(response_payload) => build_message(&mut response_builder, response_payload),
78            Err(e) => {
79                let error = UStatus::from(e);
80                response_builder
81                    .with_comm_status(error.code())
82                    .build_with_protobuf_payload(&error)
83            }
84        };
85
86        match response {
87            Ok(response_message) => {
88                if let Err(e) = transport_clone.send(response_message).await {
89                    info!(ucode = e.code().value(), "failed to send response message");
90                }
91            }
92            Err(e) => {
93                info!("failed to create response message: {}", e);
94            }
95        }
96    }
97}
98
99#[async_trait]
100impl<T: UTransport> UListener for RequestListener<T> {
101    async fn on_receive(&self, msg: UMessage) {
102        if msg.is_request() {
103            // cannot fail because inbound messages are validated at the transport layer already
104            let method_id = msg.sink_unchecked().resource_id();
105            self.process_valid_request(method_id, msg).await;
106        } else {
107            debug!(
108                message_type = msg.type_().to_cloudevent_type(),
109                "ignoring non-request message received by RPC server"
110            );
111        }
112    }
113}
114
115/// An [`RpcServer`] which keeps all information about registered endpoints in memory.
116///
117/// The server requires an implementations of [`UTransport`] for receiving RPC Request messages
118/// from clients and sending back RPC Response messages.
119///
120/// For each [endpoint being registered](`Self::register_endpoint`), a [`UListener`] is created for
121/// the given request handler and registered with the underlying transport. The listener is also
122/// mapped to the endpoint's method resource ID in order to prevent registration of multiple
123/// request handlers for the same method.
124pub struct InMemoryRpcServer<T, P> {
125    transport: Arc<T>,
126    uri_provider: Arc<P>,
127    request_listeners: tokio::sync::Mutex<HashMap<u16, Arc<dyn UListener>>>,
128}
129
130impl<T, P> core::fmt::Debug for InMemoryRpcServer<T, P> {
131    fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
132        f.debug_struct("InMemoryRpcServer").finish_non_exhaustive()
133    }
134}
135
136impl<T: UTransport, P: LocalUriProvider> InMemoryRpcServer<T, P> {
137    /// Creates a new RPC server for a given transport.
138    pub fn new(transport: Arc<T>, uri_provider: Arc<P>) -> Self {
139        InMemoryRpcServer {
140            transport,
141            uri_provider,
142            request_listeners: tokio::sync::Mutex::new(HashMap::new()),
143        }
144    }
145
146    fn validate_sink_filter(filter: &UUri) -> Result<(), RegistrationError> {
147        if !filter.is_rpc_method() {
148            return Err(RegistrationError::InvalidFilter(
149                "RPC endpoint's resource ID must be in range [0x0001, 0x7FFF]".to_string(),
150            ));
151        }
152        Ok(())
153    }
154
155    fn validate_origin_filter(filter: Option<&UUri>) -> Result<(), RegistrationError> {
156        if let Some(uri) = filter {
157            if !uri.is_rpc_response() {
158                return Err(RegistrationError::InvalidFilter(
159                    "origin filter's resource ID must be 0".to_string(),
160                ));
161            }
162        }
163        Ok(())
164    }
165
166    #[cfg(test)]
167    async fn contains_endpoint(&self, resource_id: u16) -> bool {
168        let listener_map = self.request_listeners.lock().await;
169        listener_map.contains_key(&resource_id)
170    }
171}
172
173#[async_trait]
174impl<T: UTransport + 'static, P: LocalUriProvider> RpcServer for InMemoryRpcServer<T, P> {
175    async fn register_endpoint(
176        &self,
177        origin_filter: Option<&UUri>,
178        resource_id: u16,
179        request_handler: Arc<dyn RequestHandler>,
180    ) -> Result<(), RegistrationError> {
181        Self::validate_origin_filter(origin_filter)?;
182        let sink_filter = self.uri_provider.get_resource_uri(resource_id);
183        Self::validate_sink_filter(&sink_filter)?;
184
185        let mut listener_map = self.request_listeners.lock().await;
186        if let Entry::Vacant(e) = listener_map.entry(resource_id) {
187            let listener = Arc::new(RequestListener {
188                request_handler,
189                transport: self.transport.clone(),
190            });
191            self.transport
192                .register_listener(
193                    origin_filter.unwrap_or(&UUri::any_with_resource_id(
194                        crate::uri::RESOURCE_ID_RESPONSE,
195                    )),
196                    Some(&sink_filter),
197                    listener.clone(),
198                )
199                .await
200                .map(|_| {
201                    e.insert(listener);
202                })
203                .map_err(RegistrationError::from)
204        } else {
205            Err(RegistrationError::MaxListenersExceeded)
206        }
207    }
208
209    async fn unregister_endpoint(
210        &self,
211        origin_filter: Option<&UUri>,
212        resource_id: u16,
213        _request_handler: Arc<dyn RequestHandler>,
214    ) -> Result<(), RegistrationError> {
215        Self::validate_origin_filter(origin_filter)?;
216        let sink_filter = self.uri_provider.get_resource_uri(resource_id);
217        Self::validate_sink_filter(&sink_filter)?;
218
219        let mut listener_map = self.request_listeners.lock().await;
220        if let Entry::Occupied(entry) = listener_map.entry(resource_id) {
221            let listener = entry.get().to_owned();
222            self.transport
223                .unregister_listener(
224                    origin_filter.unwrap_or(&UUri::any_with_resource_id(
225                        crate::uri::RESOURCE_ID_RESPONSE,
226                    )),
227                    Some(&sink_filter),
228                    listener,
229                )
230                .await
231                .map(|_| {
232                    entry.remove();
233                })
234                .map_err(RegistrationError::from)
235        } else {
236            Err(RegistrationError::NoSuchListener)
237        }
238    }
239}
240
241#[cfg(test)]
242mod tests {
243
244    // [utest->dsn~communication-layer-impl-default~1]
245
246    use super::*;
247
248    // use protobuf::well_known_types::wrappers::StringValue;
249    use test_case::test_case;
250    use tokio::sync::Notify;
251
252    use crate::{
253        communication::rpc::MockRequestHandler, utransport::MockTransport, StaticUriProvider,
254        UAttributes, UCode, UUri, UUID,
255    };
256
257    fn new_uri_provider() -> Arc<StaticUriProvider> {
258        Arc::new(
259            StaticUriProvider::new("", 0x0005, 0x02)
260                .expect("should have been able to create URI provider"),
261        )
262    }
263
264    #[test_case(None, 0x4A10; "for empty origin filter")]
265    #[test_case(Some(UUri::try_from_parts("authority", 0xBF1A, 0x01, 0x0000).unwrap()), 0x4A10; "for specific origin filter")]
266    #[test_case(Some(UUri::try_from_parts("*", 0xFFFF, 0x01, 0x0000).unwrap()), 0x7091; "for wildcard origin filter")]
267    #[tokio::test]
268    async fn test_register_endpoint_succeeds(origin_filter: Option<UUri>, resource_id: u16) {
269        // GIVEN an RpcServer for a transport
270        let request_handler = Arc::new(MockRequestHandler::new());
271        let mut transport = MockTransport::new();
272        let uri_provider = new_uri_provider();
273        let expected_source_filter = origin_filter
274            .clone()
275            .unwrap_or(UUri::any_with_resource_id(0));
276        let param_check = move |source_filter: &UUri,
277                                sink_filter: &Option<&UUri>,
278                                _listener: &Arc<dyn UListener>| {
279            source_filter == &expected_source_filter
280                && sink_filter.is_some_and(|uri| uri.resource_id() == resource_id)
281        };
282        transport
283            .expect_do_register_listener()
284            .once()
285            .withf(param_check.clone())
286            .returning(|_source_filter, _sink_filter, _listener| Ok(()));
287        transport
288            .expect_do_unregister_listener()
289            .once()
290            .withf(param_check)
291            .returning(|_source_filter, _sink_filter, _listener| Ok(()));
292
293        let rpc_server = InMemoryRpcServer::new(Arc::new(transport), uri_provider);
294
295        // WHEN registering a request handler
296        let register_result = rpc_server
297            .register_endpoint(origin_filter.as_ref(), resource_id, request_handler.clone())
298            .await;
299        // THEN registration succeeds
300        assert!(register_result.is_ok());
301        assert!(rpc_server.contains_endpoint(resource_id).await);
302
303        // and the handler can be unregistered again
304        let unregister_result = rpc_server
305            .unregister_endpoint(origin_filter.as_ref(), resource_id, request_handler)
306            .await;
307        assert!(unregister_result.is_ok());
308        assert!(!rpc_server.contains_endpoint(resource_id).await);
309    }
310
311    #[test_case(None, 0x0000; "for resource ID 0")]
312    #[test_case(None, 0x8000; "for resource ID out of range")]
313    #[test_case(Some(UUri::try_from_parts("*", 0xFFFF, 0xFF, 0x0001).unwrap()), 0x4A10; "for source filter with invalid resource ID")]
314    #[tokio::test]
315    async fn test_register_endpoint_fails(origin_filter: Option<UUri>, resource_id: u16) {
316        // GIVEN an RpcServer for a transport
317        let request_handler = Arc::new(MockRequestHandler::new());
318        let mut transport = MockTransport::new();
319        let uri_provider = new_uri_provider();
320        transport.expect_do_register_listener().never();
321        transport.expect_do_unregister_listener().never();
322
323        let rpc_server = InMemoryRpcServer::new(Arc::new(transport), uri_provider);
324
325        // WHEN registering a request handler using invalid parameters
326        let register_result = rpc_server
327            .register_endpoint(origin_filter.as_ref(), resource_id, request_handler.clone())
328            .await;
329        // THEN registration fails
330        assert!(register_result.is_err_and(|e| matches!(e, RegistrationError::InvalidFilter(_v))));
331        assert!(!rpc_server.contains_endpoint(resource_id).await);
332
333        // and an attempt to unregister the handler using the same invalid parameters also fails with the same error
334        let unregister_result = rpc_server
335            .unregister_endpoint(origin_filter.as_ref(), resource_id, request_handler)
336            .await;
337        assert!(unregister_result.is_err_and(|e| matches!(e, RegistrationError::InvalidFilter(_v))));
338    }
339
340    #[tokio::test]
341    async fn test_register_endpoint_fails_for_duplicate_endpoint() {
342        // GIVEN an RpcServer for a transport
343        let request_handler = Arc::new(MockRequestHandler::new());
344        let mut transport = MockTransport::new();
345        let uri_provider = new_uri_provider();
346        transport
347            .expect_do_register_listener()
348            .once()
349            .return_const(Ok(()));
350
351        let rpc_server = InMemoryRpcServer::new(Arc::new(transport), uri_provider);
352
353        // WHEN registering a request handler for an already existing endpoint
354        assert!(rpc_server
355            .register_endpoint(None, 0x5000, request_handler.clone())
356            .await
357            .is_ok());
358        let result = rpc_server
359            .register_endpoint(None, 0x5000, request_handler)
360            .await;
361
362        // THEN registration of the additional handler fails
363        assert!(result.is_err_and(|e| matches!(e, RegistrationError::MaxListenersExceeded)));
364        // but the original endpoint is still registered
365        assert!(rpc_server.contains_endpoint(0x5000).await);
366    }
367
368    #[tokio::test]
369    async fn test_unregister_endpoint_fails_for_non_existing_endpoint() {
370        // GIVEN an RpcServer for a transport
371        let request_handler = Arc::new(MockRequestHandler::new());
372        let mut transport = MockTransport::new();
373        let uri_provider = new_uri_provider();
374        transport.expect_do_unregister_listener().never();
375
376        let rpc_server = InMemoryRpcServer::new(Arc::new(transport), uri_provider);
377
378        // WHEN trying to unregister a non existing endpoint
379        assert!(!rpc_server.contains_endpoint(0x5000).await);
380        let result = rpc_server
381            .unregister_endpoint(None, 0x5000, request_handler)
382            .await;
383
384        // THEN registration fails
385        assert!(result.is_err_and(|e| matches!(e, RegistrationError::NoSuchListener)));
386    }
387
388    #[tokio::test]
389    async fn test_request_listener_invokes_operation_successfully() {
390        let mut request_handler = MockRequestHandler::new();
391        let mut transport = MockTransport::new();
392        let notify = Arc::new(Notify::new());
393        let notify_clone = notify.clone();
394        let value = b"Hello";
395        let message_id = UUID::build();
396        let message_id_clone = message_id.clone();
397        let message_source = UUri::try_from("up://localhost/A100/1/0").unwrap();
398        let message_source_clone = message_source.clone();
399
400        request_handler
401            .expect_handle_request()
402            .once()
403            .withf(move |resource_id, message_attributes, request_payload| {
404                request_payload.as_ref().is_some_and(|pl| {
405                    let message_source = message_attributes.source();
406                    pl.payload().to_vec().as_slice() == value.as_slice()
407                        && *resource_id == 0x7000_u16
408                        && *message_source == message_source_clone
409                })
410            })
411            .returning(|_resource_id, _message_attributes, _request_payload| {
412                let response_payload = UPayload::new(value.as_slice(), crate::UPayloadFormat::Raw);
413                Ok(Some(response_payload))
414            });
415        transport
416            .expect_do_send()
417            .once()
418            .withf(move |response_message| {
419                response_message.payload() == Some(value.as_slice().into())
420                    && response_message.is_response()
421                    && response_message
422                        .commstatus()
423                        .is_none_or(|code| code == UCode::Ok)
424                    && response_message.request_id_unchecked() == &message_id_clone
425            })
426            .returning(move |_msg| {
427                notify_clone.notify_one();
428                Ok(())
429            });
430        let request_message = UMessageBuilder::request(
431            UUri::try_from("up://localhost/A200/1/7000").unwrap(),
432            message_source,
433            5_000,
434        )
435        .with_message_id(message_id)
436        .build_with_payload(value.as_slice(), crate::UPayloadFormat::Raw)
437        .unwrap();
438
439        let request_listener = RequestListener {
440            request_handler: Arc::new(request_handler),
441            transport: Arc::new(transport),
442        };
443        request_listener.on_receive(request_message).await;
444        let result = tokio::time::timeout(Duration::from_secs(2), notify.notified()).await;
445        assert!(result.is_ok());
446    }
447
448    #[tokio::test]
449    async fn test_request_listener_invokes_operation_erroneously() {
450        let mut request_handler = MockRequestHandler::new();
451        let mut transport = MockTransport::new();
452        let notify = Arc::new(Notify::new());
453        let notify_clone = notify.clone();
454        let message_id = UUID::build();
455        let message_id_clone = message_id.clone();
456
457        request_handler
458            .expect_handle_request()
459            .once()
460            .withf(|resource_id, _message_attributes, _request_payload| *resource_id == 0x7000_u16)
461            .returning(|_resource_id, _message_attributes, _request_payload| {
462                Err(ServiceInvocationError::NotFound(
463                    "no such object".to_string(),
464                ))
465            });
466        transport
467            .expect_do_send()
468            .once()
469            .withf(move |response_message| {
470                let error: UStatus = response_message.extract_protobuf().unwrap();
471                error.code() == UCode::NotFound
472                    && response_message.is_response()
473                    && response_message.commstatus_unchecked() == error.code()
474                    && response_message.request_id_unchecked() == &message_id_clone
475            })
476            .returning(move |_msg| {
477                notify_clone.notify_one();
478                Ok(())
479            });
480        let request_message = UMessageBuilder::request(
481            UUri::try_from("up://localhost/A200/1/7000").unwrap(),
482            UUri::try_from("up://localhost/A100/1/0").unwrap(),
483            5_000,
484        )
485        .with_message_id(message_id)
486        .build()
487        .unwrap();
488
489        let request_listener = RequestListener {
490            request_handler: Arc::new(request_handler),
491            transport: Arc::new(transport),
492        };
493        request_listener.on_receive(request_message).await;
494        let result = tokio::time::timeout(Duration::from_secs(2), notify.notified()).await;
495        assert!(result.is_ok());
496    }
497
498    #[tokio::test]
499    async fn test_request_listener_times_out() {
500        // we need to manually implement the RequestHandler
501        // because from within the MockRequestHandler's expectation
502        // we cannot yield the current task (we can only use the blocking
503        // thread::sleep function)
504        struct NonRespondingHandler;
505        #[async_trait]
506        impl RequestHandler for NonRespondingHandler {
507            async fn handle_request(
508                &self,
509                resource_id: u16,
510                _message_attributes: &UAttributes,
511                _request_payload: Option<UPayload>,
512            ) -> Result<Option<UPayload>, ServiceInvocationError> {
513                assert_eq!(resource_id, 0x7000);
514                // this will yield the current task and allow the
515                // RequestListener to run into the timeout
516                tokio::time::sleep(Duration::from_millis(2000)).await;
517                Ok(None)
518            }
519        }
520
521        let request_handler = NonRespondingHandler {};
522        let mut transport = MockTransport::new();
523        let notify = Arc::new(Notify::new());
524        let notify_clone = notify.clone();
525        let message_id = UUID::build();
526        let message_id_clone = message_id.clone();
527
528        transport
529            .expect_do_send()
530            .once()
531            .withf(move |response_message| {
532                let error: UStatus = response_message.extract_protobuf().unwrap();
533                error.code() == UCode::DeadlineExceeded
534                    && response_message.is_response()
535                    && response_message.commstatus_unchecked() == error.code()
536                    && response_message.request_id_unchecked() == &message_id_clone
537            })
538            .returning(move |_msg| {
539                notify_clone.notify_one();
540                Ok(())
541            });
542        let request_message = UMessageBuilder::request(
543            UUri::try_from("up://localhost/A200/1/7000").unwrap(),
544            UUri::try_from("up://localhost/A100/1/0").unwrap(),
545            // make sure this request times out very quickly
546            100,
547        )
548        .with_message_id(message_id)
549        .build()
550        .expect("should have been able to create RPC Request message");
551
552        let request_listener = RequestListener {
553            request_handler: Arc::new(request_handler),
554            transport: Arc::new(transport),
555        };
556        request_listener.on_receive(request_message).await;
557        let result = tokio::time::timeout(Duration::from_secs(2), notify.notified()).await;
558        assert!(result.is_ok());
559    }
560}