dns.rsannotateddns.rssource67 lines · 2.6 KB · raw
1//! A DNS server on loopback answering A queries from a fixed table, and
2//! recording every question it was asked. Everything else gets an empty
3//! NOERROR answer.
4
5use std::net::{Ipv4Addr, SocketAddr};
6use std::sync::{Arc, Mutex};
7
8use hickory_proto::op::Message;
9use hickory_proto::rr::rdata::A;
10use hickory_proto::rr::{RData, Record, RecordType};
11use tokio::net::UdpSocket;
12
13pub struct MockDns {
14    pub addr: SocketAddr,
15    queries: Arc<Mutex<Vec<(String, RecordType)>>>,
16    task: tokio::task::JoinHandle<()>,
17}
18
19impl MockDns {
20    pub async fn start(table: &[(&str, Ipv4Addr)]) -> Self {
21        let socket = UdpSocket::bind("127.0.0.1:0").await.unwrap();
22        let addr = socket.local_addr().unwrap();
23        let table: Vec<(String, Ipv4Addr)> = table.iter().map(|(n, ip)| (fqdn(n), *ip)).collect();
24        let queries = Arc::new(Mutex::new(Vec::new()));
25        let task = tokio::spawn({
26            let queries = queries.clone();
27            async move {
28                let mut buf = [0u8; 1500];
29                loop {
30                    let Ok((n, from)) = socket.recv_from(&mut buf).await else { return };
31                    let Ok(request) = Message::from_vec(&buf[..n]) else { continue };
32                    let mut response = Message::response(request.metadata.id, request.metadata.op_code);
33                    response.metadata.recursion_desired = request.metadata.recursion_desired;
34                    response.metadata.recursion_available = true;
35                    for query in &request.queries {
36                        let name = query.name().to_ascii().to_lowercase();
37                        queries.lock().unwrap().push((name.clone(), query.query_type()));
38                        if query.query_type() == RecordType::A {
39                            for (_, ip) in table.iter().filter(|(n, _)| *n == name) {
40                                response.add_answer(Record::from_rdata(query.name().clone(), 60, RData::A(A(*ip))));
41                            }
42                        }
43                        response.add_query(query.clone());
44                    }
45                    let _ = socket.send_to(&response.to_vec().unwrap(), from).await;
46                }
47            }
48        });
49        MockDns { addr, queries, task }
50    }
51
52    /// Every (name, type) asked, names fully qualified and lowercase.
53    pub fn queries(&self) -> Vec<(String, RecordType)> {
54        self.queries.lock().unwrap().clone()
55    }
56}
57
58impl Drop for MockDns {
59    fn drop(&mut self) {
60        self.task.abort();
61    }
62}
63
64fn fqdn(name: &str) -> String {
65    let name = name.to_lowercase();
66    if name.ends_with('.') { name } else { format!("{name}.") }
67}