// 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))
}