///|
/// Return whether larger scores are better for a metric.
pub fn higher_score_is_better(metric : DistanceMetric) -> Bool {
  match metric {
    Cosine | DotProduct => true
    Euclidean | Manhattan => false
  }
}

///|
/// Compare two scores according to their metric.
pub fn score_is_better(
  candidate : Double,
  current : Double,
  metric : DistanceMetric,
) -> Bool {
  if higher_score_is_better(metric) {
    candidate > current
  } else {
    candidate < current
  }
}

///|
/// Return ids for a relevance judgement set.
pub fn relevant_ids(results : Array[SearchResult]) -> Map[String, Bool] {
  let ids = Map([])
  for result in results {
    ids.set(result.id, true)
  }
  ids
}

///|
/// Compute precision@k for an approximate ranking against a relevance set.
pub fn precision_at_k(
  results : Array[SearchResult],
  relevant : Array[SearchResult],
  k : Int,
) -> Double {
  if k <= 0 {
    return 0.0
  }
  let ids = relevant_ids(relevant)
  let limit = if k < results.length() { k } else { results.length() }
  if limit == 0 {
    return 0.0
  }
  let mut hits = 0
  for i = 0; i < limit; i = i + 1 {
    if ids.contains(results[i].id) {
      hits = hits + 1
    }
  }
  hits.to_double() / limit.to_double()
}

///|
/// Compute average precision for a ranked result list.
pub fn average_precision(
  results : Array[SearchResult],
  relevant : Array[SearchResult],
) -> Double {
  if relevant.length() == 0 {
    return 0.0
  }
  let ids = relevant_ids(relevant)
  let mut hits = 0
  let mut total = 0.0
  for i = 0; i < results.length(); i = i + 1 {
    if ids.contains(results[i].id) {
      hits = hits + 1
      total = total + hits.to_double() / (i + 1).to_double()
    }
  }
  total / relevant.length().to_double()
}

///|
/// Compute mean average precision for corresponding query batches.
pub fn mean_average_precision(
  approximate : Array[Array[SearchResult]],
  exact : Array[Array[SearchResult]],
) -> Double {
  let count = if approximate.length() < exact.length() {
    approximate.length()
  } else {
    exact.length()
  }
  if count == 0 {
    return 0.0
  }
  let mut total = 0.0
  for i = 0; i < count; i = i + 1 {
    total = total + average_precision(approximate[i], exact[i])
  }
  total / count.to_double()
}

///|
/// Compute discounted cumulative gain at k.
pub fn dcg_at_k(
  results : Array[SearchResult],
  relevant : Array[SearchResult],
  k : Int,
) -> Double {
  if k <= 0 {
    return 0.0
  }
  let ids = relevant_ids(relevant)
  let limit = if k < results.length() { k } else { results.length() }
  let mut total = 0.0
  for i = 0; i < limit; i = i + 1 {
    if ids.contains(results[i].id) {
      total = total + dcg_discount(i + 1)
    }
  }
  total
}

///|
fn dcg_discount(rank : Int) -> Double {
  match rank {
    1 => 1.0
    2 => 0.6309297535714574
    3 => 0.5
    4 => 0.43067655807339306
    5 => 0.38685280723454163
    6 => 0.3562071871080222
    7 => 0.3333333333333333
    8 => 0.3154648767857287
    9 => 0.3010299956639812
    10 => 0.2890648263178878
    _ => 1.0 / (rank + 1).to_double().sqrt()
  }
}

///|
/// Compute normalized discounted cumulative gain at k.
pub fn ndcg_at_k(
  results : Array[SearchResult],
  relevant : Array[SearchResult],
  k : Int,
) -> Double {
  if relevant.length() == 0 || k <= 0 {
    return 0.0
  }
  let actual = dcg_at_k(results, relevant, k)
  let ideal = dcg_at_k(relevant, relevant, k)
  if ideal == 0.0 {
    0.0
  } else {
    actual / ideal
  }
}

///|
/// Count exact matches at every requested cutoff.
pub fn recall_curve(
  approximate : Array[SearchResult],
  exact : Array[SearchResult],
  cutoffs : Array[Int],
) -> Array[(Int, Double)] {
  let curve = []
  for cutoff in cutoffs {
    curve.push((cutoff, recall_at_k(approximate, exact, cutoff)))
  }
  curve
}

///|
/// Return all result scores in ranking order for logging or charting.
pub fn result_scores(results : Array[SearchResult]) -> Array[Double] {
  let scores = []
  for result in results {
    scores.push(result.score)
  }
  scores
}

///|
/// Compute the largest score gap between adjacent results.
pub fn largest_score_gap(results : Array[SearchResult]) -> Double {
  if results.length() < 2 {
    return 0.0
  }
  let mut largest = 0.0
  for i = 1; i < results.length(); i = i + 1 {
    let gap = (results[i - 1].score - results[i].score).abs()
    if gap > largest {
      largest = gap
    }
  }
  largest
}

///|
/// Count distinct relevant documents retrieved by an approximate batch.
pub fn retrieved_relevant_count(
  approximate : Array[SearchResult],
  relevant : Array[SearchResult],
) -> Int {
  let ids = relevant_ids(relevant)
  let seen = Map([])
  for result in approximate {
    if ids.contains(result.id) {
      seen.set(result.id, true)
    }
  }
  seen.length()
}

///|
/// Measure how much of the relevance set appears anywhere in a result list.
pub fn coverage(
  results : Array[SearchResult],
  relevant : Array[SearchResult],
) -> Double {
  if relevant.length() == 0 {
    return 0.0
  }
  retrieved_relevant_count(results, relevant).to_double() /
  relevant.length().to_double()
}

///|
/// Return an aggregate coverage score for corresponding query batches.
pub fn mean_coverage(
  approximate : Array[Array[SearchResult]],
  exact : Array[Array[SearchResult]],
) -> Double {
  let count = if approximate.length() < exact.length() {
    approximate.length()
  } else {
    exact.length()
  }
  if count == 0 {
    return 0.0
  }
  let mut total = 0.0
  for i = 0; i < count; i = i + 1 {
    total = total + coverage(approximate[i], exact[i])
  }
  total / count.to_double()
}