///|
fn strategy_name(strategy : NegativeStrategy) -> String {
  match strategy {
    Tail(size) => "tail-\{size}"
    HardWindow(size) => "hard-\{size}"
    Stride(step) => "stride-\{step}"
  }
}

///|
fn pool_map(pools : Array[CandidatePool]) -> Map[String, Array[String]] {
  let mapped : Map[String, Array[String]] = Map([])
  for pool in pools {
    mapped[pool.query_id] = pool.doc_ids.copy()
  }
  mapped
}

///|
fn sampled_doc_ids(
  doc_ids : Array[String],
  strategy : NegativeStrategy,
  limit : Int,
) -> Array[(String, Int)] {
  let picked : Array[(String, Int)] = []
  if limit <= 0 {
    return picked
  }
  match strategy {
    Tail(window) => {
      let start = Int::max(0, doc_ids.length() - Int::max(1, window))
      for index in start..= limit {
          break
        }
        picked.push((doc_ids[index], index + 1))
      }
    }
    HardWindow(window) => {
      let top = Int::min(doc_ids.length(), Int::max(1, window))
      for index in 0..= limit {
          break
        }
        picked.push((doc_ids[index], index + 1))
      }
    }
    Stride(step) => {
      let stride = Int::max(step, 1)
      for index in 0..= limit {
          break
        }
        if index % stride == 0 {
          picked.push((doc_ids[index], index + 1))
        }
      }
    }
  }
  picked
}

///|
pub fn sample_negatives(
  qrels : Array[JudgedDoc],
  run : Array[RetrievedDoc],
  pools? : Array[CandidatePool] = [],
  config? : NegativeSampleConfig = NegativeSampleConfig::default(),
) -> Array[NegativeSample] {
  let qrels_by_query = group_qrels(qrels)
  let run_pools = build_candidate_pools(run)
  let external_pools = pool_map(pools)
  let fallback_pools = pool_map(run_pools)
  let query_ids = sorted_query_ids(qrels_by_query, group_runs(run))
  let samples : Array[NegativeSample] = []
  for query_id in query_ids {
    let judgments = qrels_by_query.get_or_default(query_id, [])
    let blocked : Map[String, Unit] = Map([])
    for item in judgments {
      if item.relevance >= config.relevant_threshold || config.skip_judged {
        blocked[item.doc_id] = ()
      }
    }
    let doc_ids = match external_pools.get(query_id) {
      Some(ids) => ids
      None => fallback_pools.get_or_default(query_id, [])
    }
    let selected = sampled_doc_ids(
      doc_ids,
      config.strategy,
      config.per_query * 3,
    )
    let seen : Map[String, Unit] = Map([])
    for pair in selected {
      let (doc_id, rank) = pair
      if blocked.contains(doc_id) || seen.contains(doc_id) {
        continue
      }
      seen[doc_id] = ()
      samples.push({
        query_id,
        doc_id,
        source_rank: rank,
        strategy: strategy_name(config.strategy),
      })
      if seen.length() >= config.per_query {
        break
      }
    }
  }
  samples.sort_by(fn(a, b) {
    let by_query = a.query_id.compare(b.query_id)
    if by_query == 0 {
      a.source_rank - b.source_rank
    } else {
      by_query
    }
  })
  samples
}

///|
pub fn render_negative_samples_tsv(samples : Array[NegativeSample]) -> String {
  let rows : Array[String] = ["query_id\tdoc_id\tsource_rank\tstrategy"]
  for sample in samples {
    rows.push(
      "\{sample.query_id}\t\{sample.doc_id}\t\{sample.source_rank}\t\{sample.strategy}",
    )
  }
  rows.join("\n")
}