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
use std::convert::TryFrom;
use std::net::SocketAddr;
use std::time::Duration;
use anyhow::{bail, Context, Result};
use async_net::TcpStream;
use async_trait::async_trait;
use byteorder::{BigEndian, ByteOrder};
use bytes::{BufMut, BytesMut};
use futures_lite::{AsyncReadExt, AsyncWriteExt};
use rand::Rng;
use tracing::debug;
use crate::client::DnsClient;
use crate::codec::{decoder::DNSMessageDecoder, encoder::ENCODER, 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<TcpStream>,
response_buffer: BytesMut,
timeout: Duration,
}
/// DNS Client that queries a server over TCP
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 conn = timeout::timeout(TcpStream::connect(self.dns_server), &self.timeout)
.await
.context("TCP connect timed out")?
.context("TCP connect failed")?;
self.conn = Some(conn);
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 TCP-specific length prefix
query_buffer.reserve(2);
query_buffer.put_u16(0);
// Just use our max for the "udp size"
ENCODER.encode(request, Some(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 TCP 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 TCP upstream {:?}: {}",
self.dns_server,
e
)
}
None => {
// Mark connection as dead, reconnect again on next query
self.conn = None;
bail!("Timed out writing to TCP 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 TCP upstream {:?}: {}",
self.dns_server,
e
)
}
None => {
// Mark connection as dead, reconnect again on next query
self.conn = None;
bail!(
"Timed out reading header from TCP 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 TCP upstream {:?}: {}",
self.dns_server,
e
)
}
None => {
// Mark connection as dead, reconnect again on next query
self.conn = None;
bail!(
"Timed out reading payload from TCP 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(response)) => {
debug!("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
);
}
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[..],
))
}
}
}
}