bitwarden_ipc/rpc/exec/
handler_registry.rs1use erased_serde::Serialize as ErasedSerialize;
2use tokio::sync::RwLock;
3
4use super::handler::{ErasedRpcHandler, RpcHandler};
5use crate::rpc::{error::RpcError, request::RpcRequest, request_message::RpcRequestPayload};
6
7pub struct RpcHandlerRegistry {
8 handlers: RwLock<std::collections::HashMap<String, Box<dyn ErasedRpcHandler>>>,
9}
10
11impl RpcHandlerRegistry {
12 pub fn new() -> Self {
13 Self {
14 handlers: RwLock::new(std::collections::HashMap::new()),
15 }
16 }
17
18 pub async fn register<H>(&self, handler: H)
19 where
20 H: RpcHandler + ErasedRpcHandler + 'static,
21 {
22 let name = H::Request::NAME.to_owned();
23 self.register_erased(name, Box::new(handler)).await;
24 }
25
26 pub async fn register_erased(&self, name: String, handler: Box<dyn ErasedRpcHandler>) {
27 self.handlers.write().await.insert(name, handler);
28 }
29
30 pub async fn handle(
31 &self,
32 request: &RpcRequestPayload,
33 ) -> Result<Box<dyn ErasedSerialize>, RpcError> {
34 match self.handlers.read().await.get(request.request_type()) {
35 Some(handler) => handler.handle(request).await,
36 None => Err(RpcError::NoHandlerFound),
37 }
38 }
39}
40
41#[cfg(test)]
42mod test {
43 use serde::{Deserialize, Serialize, de::DeserializeOwned};
44
45 use super::*;
46 use crate::{
47 rpc::{request::RpcRequest, request_message::RpcRequestMessage},
48 serde_utils,
49 };
50
51 #[derive(Debug, Clone, Serialize, Deserialize)]
52 struct TestRequest {
53 a: i32,
54 b: i32,
55 }
56
57 #[derive(Debug, Clone, Serialize, Deserialize)]
58 struct TestResponse {
59 result: i32,
60 }
61
62 impl RpcRequest for TestRequest {
63 type Response = TestResponse;
64
65 const NAME: &str = "TestRequest";
66 }
67
68 struct TestHandler;
69
70 impl RpcHandler for TestHandler {
71 type Request = TestRequest;
72
73 async fn handle(&self, request: Self::Request) -> TestResponse {
74 TestResponse {
75 result: request.a + request.b,
76 }
77 }
78 }
79
80 #[tokio::test]
81 async fn handle_returns_error_when_no_handler_can_be_found() {
82 let registry = RpcHandlerRegistry::new();
83
84 let request = TestRequest { a: 1, b: 2 };
85 let message = RpcRequestMessage {
86 request,
87 request_id: "test_id".to_string(),
88 request_type: "TestRequest".to_string(),
89 response_topic: "RpcResponseMessage:test_id".to_string(),
90 };
91 let serialized_request =
92 RpcRequestPayload::from_slice(serde_utils::to_vec(&message).unwrap()).unwrap();
93
94 let result = registry.handle(&serialized_request).await;
95
96 assert!(matches!(result, Err(RpcError::NoHandlerFound)));
97 }
98
99 #[tokio::test]
100 async fn handle_runs_previously_registered_handler() {
101 let registry = RpcHandlerRegistry::new();
102
103 registry.register(TestHandler).await;
104
105 let request = TestRequest { a: 1, b: 2 };
106 let message = RpcRequestMessage {
107 request,
108 request_id: "test_id".to_string(),
109 request_type: "TestRequest".to_string(),
110 response_topic: "RpcResponseMessage:test_id".to_string(),
111 };
112 let serialized_request =
113 RpcRequestPayload::from_slice(serde_utils::to_vec(&message).unwrap()).unwrap();
114
115 let result = registry
116 .handle(&serialized_request)
117 .await
118 .expect("Failed to handle request");
119 let response: TestResponse = deserialize_erased_object(&result);
120
121 assert_eq!(response.result, 3);
122 }
123
124 fn deserialize_erased_object<T, R>(value: &T) -> R
125 where
126 T: Serialize,
127 R: DeserializeOwned,
128 {
129 let serialized = serde_utils::to_vec(value).expect("Failed to serialize erased serialize");
130
131 serde_utils::from_slice(&serialized).expect("Failed to deserialize erased serialize")
132 }
133}