1use 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
33fn handle_response_message(response: UMessage) -> Result<Option<UPayload>, ServiceInvocationError> {
35 match response.commstatus() {
36 Some(UCode::Ok) | None => {
37 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 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 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 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 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 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 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
147pub 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 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 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 let mut mock_transport = MockTransport::default();
303 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 let creation_attempt =
316 InMemoryRpcClient::new(Arc::new(mock_transport), new_uri_provider()).await;
317
318 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 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 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 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 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 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 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 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 let (response_listener_result, _) = join!(captured_listener_rx, request_sent.notified());
435 let response_listener = response_listener_result.unwrap();
436
437 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 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 response_listener.on_receive(response_message).await;
453 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 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 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 first_request_sent.notified().await;
510
511 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 let response = second_request_handle.await.unwrap();
528 assert!(response.is_err_and(|e| matches!(e, ServiceInvocationError::AlreadyExists(_))));
529 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 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 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 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 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 let mut mock_transport = MockTransport::default();
582 mock_transport
583 .expect_do_register_listener()
584 .returning(|_source_filter, _sink_filter, _listener| Ok(()));
585 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 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 assert!(response.is_err_and(|e| { matches!(e, ServiceInvocationError::DeadlineExceeded) }));
603 assert!(!client.contains_pending_request(&message_id));
604 }
605}