hickory_client/client/
memoize_client_handle.rs1use std::collections::HashMap;
8use std::pin::Pin;
9use std::sync::Arc;
10
11use futures_util::future::FutureExt;
12use futures_util::lock::Mutex;
13use futures_util::stream::Stream;
14use hickory_proto::{
15 ProtoError,
16 op::Query,
17 xfer::{DnsHandle, DnsRequest, DnsResponse},
18};
19
20use crate::client::ClientHandle;
21use crate::client::rc_stream::{RcStream, rc_stream};
22
23#[derive(Clone)]
30#[must_use = "queries can only be sent through a ClientHandle"]
31pub struct MemoizeClientHandle<H: ClientHandle> {
32 client: H,
33 active_queries: Arc<Mutex<HashMap<Query, RcStream<<H as DnsHandle>::Response>>>>,
34}
35
36impl<H> MemoizeClientHandle<H>
37where
38 H: ClientHandle,
39{
40 pub fn new(client: H) -> Self {
42 Self {
43 client,
44 active_queries: Arc::new(Mutex::new(HashMap::new())),
45 }
46 }
47
48 async fn inner_send(
49 request: DnsRequest,
50 active_queries: Arc<Mutex<HashMap<Query, RcStream<<H as DnsHandle>::Response>>>>,
51 client: H,
52 ) -> impl Stream<Item = Result<DnsResponse, ProtoError>> {
53 let query = request.queries().first().expect("no query!").clone();
55
56 let mut active_queries = active_queries.lock().await;
58
59 if let Some(rc_stream) = active_queries.get(&query) {
62 return rc_stream.clone();
63 };
64
65 active_queries
67 .entry(query)
68 .or_insert_with(|| rc_stream(client.send(request)))
69 .clone()
70 }
71}
72
73impl<H> DnsHandle for MemoizeClientHandle<H>
74where
75 H: ClientHandle,
76{
77 type Response = Pin<Box<dyn Stream<Item = Result<DnsResponse, ProtoError>> + Send>>;
78
79 fn send<R: Into<DnsRequest>>(&self, request: R) -> Self::Response {
80 let request = request.into();
81
82 Box::pin(
83 Self::inner_send(
84 request,
85 Arc::clone(&self.active_queries),
86 self.client.clone(),
87 )
88 .flatten_stream(),
89 )
90 }
91}
92
93#[cfg(test)]
94mod test {
95 #![allow(clippy::dbg_macro, clippy::print_stdout)]
96
97 use std::pin::Pin;
98 use std::sync::Arc;
99
100 use futures::lock::Mutex;
101 use futures::*;
102 use hickory_proto::{
103 ProtoError,
104 op::{Message, Query},
105 rr::RecordType,
106 xfer::{DnsHandle, DnsRequest, DnsResponse},
107 };
108 use test_support::subscribe;
109
110 use crate::client::*;
111 use hickory_proto::xfer::FirstAnswer;
112
113 #[derive(Clone)]
114 struct TestClient {
115 i: Arc<Mutex<u16>>,
116 }
117
118 impl DnsHandle for TestClient {
119 type Response = Pin<Box<dyn Stream<Item = Result<DnsResponse, ProtoError>> + Send>>;
120
121 fn send<R: Into<DnsRequest> + Send + 'static>(&self, request: R) -> Self::Response {
122 let i = Arc::clone(&self.i);
123 let future = async {
124 let i = i;
125 let request = request;
126 let mut message = Message::new();
127
128 let mut i = i.lock().await;
129
130 message.set_id(*i);
131 println!(
132 "sending {}: {}",
133 *i,
134 request.into().queries().first().expect("no query!").clone()
135 );
136
137 *i += 1;
138
139 Ok(DnsResponse::from_message(message).unwrap())
140 };
141
142 Box::pin(stream::once(future))
143 }
144 }
145
146 #[test]
147 fn test_memoized() {
148 use futures::executor::block_on;
149
150 subscribe();
151
152 let client = MemoizeClientHandle::new(TestClient {
153 i: Arc::new(Mutex::new(0)),
154 });
155
156 let mut test1 = Message::new();
157 test1.add_query(Query::new().set_query_type(RecordType::A).clone());
158
159 let mut test2 = Message::new();
160 test2.add_query(Query::new().set_query_type(RecordType::AAAA).clone());
161
162 let result = block_on(client.send(test1.clone()).first_answer())
163 .ok()
164 .unwrap();
165 assert_eq!(result.id(), 0);
166
167 let result = block_on(client.send(test2.clone()).first_answer())
168 .ok()
169 .unwrap();
170 assert_eq!(result.id(), 1);
171
172 let result = block_on(client.send(test1).first_answer()).ok().unwrap();
174 assert_eq!(result.id(), 0);
175
176 let result = block_on(client.send(test2).first_answer()).ok().unwrap();
177 assert_eq!(result.id(), 1);
178 }
179}