// TCP DNS transport (RFC 7766).
///|
pub struct TcpTransport {
server : String
timeout_ms : Int
max_retries : Int
}
///|
pub fn TcpTransport::new(server : String) -> TcpTransport {
TcpTransport::with_options(server, default_timeout_ms, default_retries)
}
///|
pub fn TcpTransport::with_options(
server : String,
timeout_ms : Int,
max_retries : Int,
) -> TcpTransport {
{
server,
timeout_ms,
max_retries: if max_retries < 0 {
0
} else {
max_retries
},
}
}
///|
pub fn TcpTransport::with_timeout(
server : String,
timeout_ms : Int,
) -> TcpTransport {
TcpTransport::with_options(server, timeout_ms, default_retries)
}
///|
pub fn tcp_dns_max_message_size() -> Int {
tcp_dns_max_len
}
///|
/// Prefix a DNS payload with its two-octet RFC 7766 TCP length field.
pub fn encode_tcp_message(
payload : Array[Byte],
) -> Result[Array[Byte], TransportError] {
let length = payload.length()
if length > tcp_dns_max_message_size() {
return Err(InvalidMessage("TCP DNS payload exceeds 65535 bytes"))
}
let frame = Array::make(length + 2, b'\x00')
frame[0] = ((length >> 8) & 0xff).to_byte()
frame[1] = (length & 0xff).to_byte()
for i in 0.. Result[Int, TransportError] {
if offset < 0 || offset + 2 > buf.length() {
Err(InvalidMessage("TCP DNS frame is missing its length prefix"))
} else {
Ok((buf[offset].to_int() << 8) | buf[offset + 1].to_int())
}
}
///|
pub fn validate_tcp_message(buf : Array[Byte]) -> Bool {
match read_tcp_length(buf) {
Ok(length) =>
length > 0 &&
length <= tcp_dns_max_message_size() &&
buf.length() == length + 2
Err(_) => false
}
}
///|
async fn TcpTransport::send_once(
self : TcpTransport,
target : @socket.Addr,
frame : Array[Byte],
) -> Result[Array[Byte], TransportError] {
let response = @async.with_timeout_opt(self.timeout_ms, () => {
let socket = @socket.Tcp::connect(target)
defer socket.close()
socket.write(array_to_bytes(frame))
let prefix = socket.read_exactly(2)
let prefix = bytes_to_array(prefix)
let length = match read_tcp_length(prefix) {
Ok(length) => length
Err(err) => return Err(err)
}
if length == 0 {
return Err(InvalidMessage("TCP DNS response has zero length"))
}
let payload = socket.read_exactly(length)
let payload = bytes_to_array(payload)
if payload.length() < dns_header_size {
Err(InvalidMessage("TCP response shorter than DNS header"))
} else {
Ok(payload)
}
}) catch {
err => return Err(RecvFailed(err.to_string()))
}
match response {
None => Err(Timeout(self.timeout_ms))
Some(result) => result
}
}
///|
/// Send a DNS request over a new TCP connection and read exactly one framed
/// response. Partial reads are handled by `Reader::read_exactly`.
pub async fn TcpTransport::send(
self : TcpTransport,
payload : Array[Byte],
) -> Result[Array[Byte], TransportError] {
if self.timeout_ms <= 0 {
return Err(InvalidMessage("timeout_ms must be positive"))
}
let frame = match encode_tcp_message(payload) {
Ok(frame) => frame
Err(err) => return Err(err)
}
let server = match parse_server_addr(self.server) {
Ok(server) => server
Err(err) => return Err(err)
}
let target = match resolve_server_addr(server) {
Ok(target) => target
Err(err) => return Err(err)
}
let mut last_error = TransportError::ConnectionFailed(
"TCP transport did not run",
)
for _ in 0..<=self.max_retries {
let result = self.send_once(target, frame)
match result {
Ok(response) => return Ok(response)
Err(err) if is_retryable(err) => last_error = err
Err(err) => return Err(err)
}
}
Err(last_error)
}
///|
pub fn TcpTransport::as_transport(self : TcpTransport) -> DnsTransport {
DnsTransport::from_request(payload => self.send(payload))
}