1use 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 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
115pub 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 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 use super::*;
247
248 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 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 let register_result = rpc_server
297 .register_endpoint(origin_filter.as_ref(), resource_id, request_handler.clone())
298 .await;
299 assert!(register_result.is_ok());
301 assert!(rpc_server.contains_endpoint(resource_id).await);
302
303 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 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 let register_result = rpc_server
327 .register_endpoint(origin_filter.as_ref(), resource_id, request_handler.clone())
328 .await;
329 assert!(register_result.is_err_and(|e| matches!(e, RegistrationError::InvalidFilter(_v))));
331 assert!(!rpc_server.contains_endpoint(resource_id).await);
332
333 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 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 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 assert!(result.is_err_and(|e| matches!(e, RegistrationError::MaxListenersExceeded)));
364 assert!(rpc_server.contains_endpoint(0x5000).await);
366 }
367
368 #[tokio::test]
369 async fn test_unregister_endpoint_fails_for_non_existing_endpoint() {
370 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 assert!(!rpc_server.contains_endpoint(0x5000).await);
380 let result = rpc_server
381 .unregister_endpoint(None, 0x5000, request_handler)
382 .await;
383
384 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 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 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 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}