// UDP DNS transport (RFC 1035 section 4.2.1).

///|
pub(all) struct ServerAddress {
  host : String
  port : Int
} derive(Debug, Eq)

///|
fn parse_port(port : StringView) -> Result[Int, TransportError] {
  if port.is_empty() {
    return Err(InvalidServer("missing port"))
  }
  let value = Ref(0)
  for c in port {
    if c < '0' || c > '9' {
      return Err(InvalidServer("port must be decimal"))
    }
    let digit = c.to_int() - '0'.to_int()
    if value.val > (65535 - digit) / 10 {
      return Err(InvalidServer("port must be between 1 and 65535"))
    }
    value.val = value.val * 10 + digit
  }
  if value.val == 0 {
    Err(InvalidServer("port must be between 1 and 65535"))
  } else {
    Ok(value.val)
  }
}

///|
/// Parse an IPv4, hostname, or bracketed IPv6 DNS server address.
/// `host`, `[ipv6]`, and `host:port` use port 53 when omitted.
pub fn parse_server_addr(
  source : String,
) -> Result[ServerAddress, TransportError] {
  let addr = source.trim().to_owned()
  if addr.is_empty() {
    return Err(InvalidServer("server address is empty"))
  }
  if addr.has_prefix("[") {
    let close = match addr.find("]") {
      Some(index) => index
      None => return Err(InvalidServer("IPv6 address is missing ]"))
    }
    let host = addr[1:close].to_owned()
    if host.is_empty() {
      return Err(InvalidServer("IPv6 address is empty"))
    }
    let suffix = addr[close + 1:]
    if suffix.is_empty() {
      return Ok({ host, port: dns_port.to_int() })
    }
    if !suffix.has_prefix(":") {
      return Err(InvalidServer("expected :port after IPv6 address"))
    }
    match parse_port(suffix[1:]) {
      Ok(port) => Ok({ host, port })
      Err(err) => Err(err)
    }
  } else {
    let parts = addr.split(":").collect()
    match parts.length() {
      1 => Ok({ host: parts[0].to_owned(), port: dns_port.to_int() })
      2 => {
        let host = parts[0].to_owned()
        if host.is_empty() {
          return Err(InvalidServer("host is empty"))
        }
        match parse_port(parts[1]) {
          Ok(port) => Ok({ host, port })
          Err(err) => Err(err)
        }
      }
      _ => Err(InvalidServer("IPv6 addresses must use [address]:port"))
    }
  }
}

///|
async fn resolve_server_addr(
  server : ServerAddress,
) -> Result[@socket.Addr, TransportError] {
  try @socket.Addr::resolve(server.host, port=server.port) catch {
    err => Err(ConnectionFailed(err.to_string()))
  } noraise {
    addr => Ok(addr)
  }
}

///|
pub struct UdpTransport {
  server : String
  timeout_ms : Int
  max_retries : Int
  udp_payload_size : UInt16
}

///|
pub fn UdpTransport::new(server : String) -> UdpTransport {
  UdpTransport::with_options(server, default_timeout_ms, default_retries)
}

///|
pub fn UdpTransport::with_options(
  server : String,
  timeout_ms : Int,
  max_retries : Int,
  udp_payload_size? : UInt16 = 1232,
) -> UdpTransport {
  {
    server,
    timeout_ms,
    max_retries: if max_retries < 0 {
      0
    } else {
      max_retries
    },
    udp_payload_size,
  }
}

///|
pub fn UdpTransport::with_timeout(
  server : String,
  timeout_ms : Int,
) -> UdpTransport {
  UdpTransport::with_options(server, timeout_ms, default_retries)
}

///|
pub fn UdpTransport::with_retries(
  server : String,
  timeout_ms : Int,
  max_retries : Int,
) -> UdpTransport {
  UdpTransport::with_options(server, timeout_ms, max_retries)
}

///|
pub fn UdpTransport::max_payload(self : UdpTransport) -> UInt16 {
  self.udp_payload_size
}

///|
fn udp_response_is_truncated(response : Array[Byte]) -> Bool {
  response.length() >= dns_header_size && (response[2].to_int() & 0x02) != 0
}

///|
async fn UdpTransport::send_once(
  self : UdpTransport,
  target : @socket.Addr,
  payload : Array[Byte],
) -> Result[Array[Byte], TransportError] {
  let socket = @socket.UdpClient(target) catch {
    err => return Err(ConnectionFailed(err.to_string()))
  }
  defer socket.close()
  socket.send(array_to_bytes(payload)) catch {
    err => return Err(SendFailed(err.to_string()))
  }
  let received = @async.with_timeout_opt(self.timeout_ms, () => {
    let buffer = FixedArray::make(self.udp_payload_size.to_int(), b'\x00')
    let n = socket.recv(buffer)
    buffer.unsafe_reinterpret_as_bytes()[:n].to_owned()
  }) catch {
    err => return Err(RecvFailed(err.to_string()))
  }
  match received {
    None => Err(Timeout(self.timeout_ms))
    Some(response) => {
      let response = bytes_to_array(response)
      if response.length() < dns_header_size {
        Err(InvalidMessage("UDP response shorter than DNS header"))
      } else if udp_response_is_truncated(response) {
        Err(Truncated)
      } else {
        Ok(response)
      }
    }
  }
}

///|
/// Send one DNS datagram. A connected UDP socket only accepts responses from
/// the configured server; the OS therefore enforces the source-address check.
pub async fn UdpTransport::send(
  self : UdpTransport,
  payload : Array[Byte],
) -> Result[Array[Byte], TransportError] {
  if self.timeout_ms <= 0 {
    return Err(InvalidMessage("timeout_ms must be positive"))
  }
  if self.udp_payload_size < 512 {
    return Err(InvalidMessage("UDP payload size must be at least 512"))
  }
  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(
    "UDP transport did not run",
  )
  for _ in 0..<=self.max_retries {
    let result = self.send_once(target, payload)
    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 UdpTransport::as_transport(self : UdpTransport) -> DnsTransport {
  DnsTransport::from_request(
    payload => self.send(payload),
    udp_payload_size=self.udp_payload_size,
  )
}