hickory_client/client/
memoize_client_handle.rs

1// Copyright 2015-2016 Benjamin Fry <benjaminfry@me.com>
2//
3// Licensed under the Apache License, Version 2.0, <LICENSE-APACHE or
4// https://apache.org/licenses/LICENSE-2.0> or the MIT license <LICENSE-MIT or
5// https://opensource.org/licenses/MIT>, at your option. This file may not be
6// copied, modified, or distributed except according to those terms.
7use 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// TODO: move to proto
24/// A ClientHandle for memoized (cached) responses to queries.
25///
26/// This wraps a ClientHandle, changing the implementation `send()` to store the response against
27///  the Message.Query that was sent. This should reduce network traffic especially during things
28///  like DNSSEC validation. *Warning* this will currently cache for the life of the Client.
29#[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    /// Returns a new handle wrapping the specified client
41    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        // TODO: what if we want to support multiple queries (non-standard)?
54        let query = request.queries().first().expect("no query!").clone();
55
56        // lock all the currently running queries
57        let mut active_queries = active_queries.lock().await;
58
59        // TODO: we need to consider TTL on the records here at some point
60        // If the query is running, grab that existing one...
61        if let Some(rc_stream) = active_queries.get(&query) {
62            return rc_stream.clone();
63        };
64
65        // Otherwise issue a new query and store in the map
66        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        // should get the same result for each...
173        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}