// Response validation, DNS status mapping, and cache TTL helpers used by the
// resolver core.

///|
fn map_transport_error(error : TransportError) -> ResolveError {
  match error {
    Timeout(_) => Timeout
    Truncated => TruncatedAndTcpFailed
    ConnectionFailed(reason) => ConnectionFailed(reason)
    SendFailed(reason) => ConnectionFailed(reason)
    RecvFailed(reason) => ConnectionFailed(reason)
    InvalidServer(reason) => ConnectionFailed(reason)
    InvalidMessage(reason) => FormatError(reason)
  }
}

///|
fn validate_response(
  response : Message,
  query_id : UInt16,
  query_name : String,
  qtype : UInt16,
) -> Result[Unit, ResolveError] {
  if response.header.id != query_id {
    return Err(ResponseMismatch("query ID does not match"))
  }
  if !response.header.is_response() {
    return Err(ResponseMismatch("QR bit is not set"))
  }
  if response.header.opcode() != opcode_query {
    return Err(ResponseMismatch("response opcode does not match QUERY"))
  }
  if response.questions.length() != 1 {
    return Err(ResponseMismatch("response must echo exactly one question"))
  }
  let question = response.questions[0]
  if !dns_name_equal(question.name, query_name) ||
    question.qtype != qtype ||
    question.qclass != qclass_in {
    return Err(ResponseMismatch("response question does not match request"))
  }
  if response.header.is_truncated() {
    return Err(TruncatedAndTcpFailed)
  }
  Ok(())
}

///|
fn ttl_seconds(value : UInt) -> Int64 {
  // DNS TTL is an unsigned 32-bit value. Int64 preserves its entire
  // 0..4_294_967_295 range without reinterpreting the high bit as negative or
  // shortening a legitimate cache lifetime.
  value.to_int64()
}

///|
fn full_rcode(message : Message) -> Int {
  let low = message.header.rcode()
  for rr in message.additionals {
    if rr.rtype == qtype_opt {
      // OPT TTL is EXT-RCODE | VERSION | FLAGS (RFC 6891 section 6.1.3).
      let high = (rr.ttl >> 24).reinterpret_as_int()
      return (high << 4) | low
    }
  }
  low
}

///|
fn rcode_error(message : Message) -> Result[Unit, ResolveError] {
  let rcode = full_rcode(message)
  if rcode > 15 {
    return Err(ExtendedRcode(rcode))
  }
  if rcode == rcode_noerror {
    Ok(())
  } else if rcode == rcode_formerr {
    Err(FormErr)
  } else if rcode == rcode_servfail {
    Err(ServFail)
  } else if rcode == rcode_nxdomain {
    Err(NXDomain)
  } else if rcode == rcode_notimp {
    Err(NotImplemented)
  } else if rcode == rcode_refused {
    Err(Refused)
  } else {
    Err(FormatError("unsupported DNS response code: " + rcode.to_string()))
  }
}

// RFC 2308 negative TTL = min(SOA TTL, SOA.MINIMUM). A response without an SOA
// does not provide a standards-defined negative-cache lifetime and therefore
// must not be cached.

///|
fn extract_negative_ttl(message : Message) -> Int64? {
  for rr in message.authorities {
    match rr.rdata {
      SOA(_, _, _, _, _, _, minimum) =>
        if rr.rtype == qtype_soa && rr.rclass == qclass_in {
          let soa_ttl = ttl_seconds(rr.ttl)
          let minimum_ttl = ttl_seconds(minimum)
          return Some(if soa_ttl < minimum_ttl { soa_ttl } else { minimum_ttl })
        }
      _ => ()
    }
  }
  None
}

///|
fn response_is_referral(message : Message) -> Bool {
  if message.answers.length() > 0 {
    return false
  }
  let has_ns = Ref(false)
  let has_soa = Ref(false)
  for rr in message.authorities {
    if rr.rclass != qclass_in {
      continue
    }
    if rr.rtype == qtype_ns {
      has_ns.val = true
    } else if rr.rtype == qtype_soa {
      has_soa.val = true
    }
  }
  has_ns.val && !has_soa.val
}

///|
fn negative_cache_ttl_for_query(
  message : Message,
  query_name : String,
  max_depth : Int,
) -> Int64? {
  let ttl = match extract_negative_ttl(message) {
    Some(value) => Ref(value)
    None => return None
  }
  let current = Ref(query_name)
  let tracker = CnameTracker::new(query_name, max_depth~)
  for ;; {
    let target = match cname_target_for_name(message.answers, current.val) {
      Some(value) => value
      None => break
    }
    match cname_ttl_for_name(message.answers, current.val) {
      Some(value) => ttl.val = min_ttl(ttl.val, value)
      None => ()
    }
    match tracker.follow(target) {
      Ok(_) => current.val = target
      // Do not cache a malformed/looping alias proof as a stable NXDOMAIN.
      Err(_) => return None
    }
  }
  Some(ttl.val)
}

///|
fn min_answer_ttl(answers : Array[RR]) -> Int64 {
  if answers.length() == 0 {
    return 0L
  }
  let minimum = Ref(ttl_seconds(answers[0].ttl))
  for rr in answers {
    let ttl = ttl_seconds(rr.ttl)
    if ttl < minimum.val {
      minimum.val = ttl
    }
  }
  minimum.val
}

///|
fn min_ttl(left : Int64, right : Int64) -> Int64 {
  if left < right {
    left
  } else {
    right
  }
}

///|
fn records_for_question(
  answers : Array[RR],
  name : String,
  qtype : UInt16,
) -> Array[RR] {
  let records : Array[RR] = Array::new(capacity=answers.length())
  for rr in answers {
    if rr.rtype == qtype &&
      rr.rclass == qclass_in &&
      dns_name_equal(rr.name, name) {
      records.push(rr)
    }
  }
  records
}

///|
fn cname_target_for_name(answers : Array[RR], name : String) -> String? {
  for rr in answers {
    match rr.rdata {
      CNAME(target) if rr.rclass == qclass_in && dns_name_equal(rr.name, name) =>
        return Some(target)
      _ => ()
    }
  }
  None
}

///|
fn cname_ttl_for_name(answers : Array[RR], name : String) -> Int64? {
  for rr in answers {
    if rr.rtype == qtype_cname &&
      rr.rclass == qclass_in &&
      dns_name_equal(rr.name, name) {
      return Some(ttl_seconds(rr.ttl))
    }
  }
  None
}