1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
use std::convert::TryFrom;
use std::net::{SocketAddr, TcpStream};
use std::time::Duration;
use anyhow::{bail, Context, Result};
use async_native_tls_alpn::{Protocol, TlsConnector, TlsStream};
use async_trait::async_trait;
use byteorder::{BigEndian, ByteOrder};
use bytes::{BufMut, BytesMut};
use futures_lite::{AsyncReadExt, AsyncWriteExt};
use rand::Rng;
use smol::Async;
use tracing::debug;
use crate::client::DnsClient;
use crate::codec::{decoder::DNSMessageDecoder, message};
use crate::specs::message::Message;
use crate::timeout;
/// TCP size header is 16 bits, so max theoretical size is 64k
static MAX_TCP_BYTES: u16 = 65535;
pub struct Client {
dns_server: SocketAddr,
conn: Option<TlsStream<Async<TcpStream>>>,
response_buffer: BytesMut,
timeout: Duration,
}
/// DNS Client that queries a server over TLS
impl Client {
/// Constructs a new `Client` that will query the specified `dns_server`.
pub fn new(dns_server: SocketAddr, timeout: Duration) -> Self {
Client {
dns_server,
conn: None,
response_buffer: BytesMut::with_capacity(MAX_TCP_BYTES as usize),
timeout,
}
}
async fn connect(&mut self) -> Result<()> {
let stream = timeout::timeout(
Async::<TcpStream>::connect(self.dns_server.clone()),
&self.timeout,
)
.await
.context("TLS TCP connect timed out")?
.context("TLS TCP connect failed")?;
// Min protocol: If things don't have at least TLS1.2 by now, we should just name and shame.
let connector = TlsConnector::new().min_protocol_version(Some(Protocol::Tlsv12));
let stream = timeout::timeout(
connector.connect(format!("{}", self.dns_server.ip()), stream),
&self.timeout,
)
.await
.context("TLS session connect timed out")?
.context("TLS session connect failed")?;
self.conn = Some(stream);
Ok(())
}
}
#[async_trait]
impl DnsClient for Client {
async fn query(
&mut self,
request: &Message,
query_buffer: &mut BytesMut,
) -> Result<Option<Message>> {
// Reserve 2 bytes for the TLS-specific length prefix
query_buffer.reserve(2);
query_buffer.put_u16(0);
// Encode the message, along with any necessary padding
super::add_request_padding(request, MAX_TCP_BYTES, query_buffer)?;
// Insert the resulting encoded size of the message
// into those leading two bytes that we'd reserved.
let message_len = u16::try_from(query_buffer.len() - 2).with_context(|| {
format!(
"Encoded request size {} exceeds {} limit: {}",
query_buffer.len() - 2,
MAX_TCP_BYTES,
request
)
})?;
query_buffer[0] = ((message_len & 0xFF00) >> 8) as u8;
query_buffer[1] = (message_len & 0xFF) as u8;
// Query is constructed, now let's do the request.
if self.conn.is_none() {
self.connect().await.with_context(|| {
format!("Failed to connect with TLS upstream {:?}", self.dns_server)
})?;
}
let request_id = rand::thread_rng().gen::<u16>();
// For TCP, the size header means that the message actually starts at byte 2
message::update_message_id(request_id, query_buffer, 2)?;
debug!(
"Raw request to {:?} ({}b): {:02X?}",
self.dns_server,
query_buffer.len(),
&query_buffer[..]
);
// NOTE: async is useful here, since it allows us to ensure that the entire write completes within the timeout.
// If we used sync APIs, we would risk a malicious upstream slowly allowing one byte at a time. (sync write_all loops over writes internally)
match timeout::timeout(
self.conn
.as_mut()
.expect("missing connection")
.write_all(query_buffer.as_ref()),
&self.timeout,
)
.await
{
Some(Ok(())) => {}
Some(Err(e)) => {
bail!(
"Failed to write to TLS upstream {:?}: {}",
self.dns_server,
e
)
}
None => {
// Mark connection as dead, reconnect again on next query
self.conn = None;
bail!("Timed out writing to TLS upstream {:?}", self.dns_server)
}
}
// Read first two bytes to get expected response size
let mut response_size_bytes: [u8; 2] = [0, 0];
match timeout::timeout(
self.conn
.as_mut()
.expect("missing connection")
.read_exact(&mut response_size_bytes),
&self.timeout,
)
.await
{
Some(Ok(())) => {}
Some(Err(e)) => {
bail!(
"Failed to read header from TLS upstream {:?}: {}",
self.dns_server,
e
)
}
None => {
// Mark connection as dead, reconnect again on next query
self.conn = None;
bail!(
"Timed out reading header from TLS upstream {:?}",
self.dns_server
)
}
}
// big endian
let response_size = BigEndian::read_u16(&response_size_bytes);
// Read remaining bytes to get response
self.response_buffer.resize(
usize::try_from(response_size).with_context(|| "couldn't convert u16 to usize")?,
0,
);
// NOTE: async is useful here, since it allows us to ensure that the entire read completes within the timeout.
// If we used sync APIs, we would risk a malicious upstream slowly allowing one byte at a time. (sync read_all loops over writes internally)
match timeout::timeout(
self.conn
.as_mut()
.expect("missing connection")
.read_exact(&mut self.response_buffer),
&self.timeout,
)
.await
{
Some(Ok(())) => {}
Some(Err(e)) => {
bail!(
"Failed to read payload from TLS upstream {:?}: {}",
self.dns_server,
e
)
}
None => {
// Mark connection as dead, reconnect again on next query
self.conn = None;
bail!(
"Timed out reading payload from TLS upstream {:?}",
self.dns_server
)
}
}
debug!(
"Raw response from {:?} ({}b): {:02X?}",
self.dns_server,
self.response_buffer.len(),
&self.response_buffer[..]
);
match DNSMessageDecoder::new().decode(&self.response_buffer[..]) {
Ok(Some(mut response)) => {
debug!(
"Untrimmed response from {:?}: {}",
self.dns_server, response
);
if response.header.truncated {
// Message claims to be truncated, shouldn't happen but let's bail anyway
return Ok(None);
}
if response.header.id != request_id {
bail!(
"Returned transaction id {:?} doesn't match sent {:?}",
response.header.id,
request_id
);
}
// Filter out any EDNS PADDING from the response before returning it.
// The PADDING is common/best practice for DoT servers.
super::remove_response_padding(&mut response);
Ok(Some(response))
}
Ok(None) => {
// Message was likely truncated, despite us receiving all the data in the payload
debug!(
"Unable to parse response from server={} to request={:02X?}: {:02X?}",
self.dns_server,
&query_buffer[..],
&self.response_buffer[..],
);
Ok(None)
}
Err(e) => {
// Other parse error
Err(e).context(format!(
"Failed to parse response from server={} to request={:02X?}: {:02X?}",
self.dns_server,
&query_buffer[..],
&self.response_buffer[..],
))
}
}
}
}