///|
pub struct SearchResult {
  id : String
  score : Double
  metadata : Array[(String, String)]
}

///|
pub impl Show for SearchResult with fn output(self, logger) {
  logger.write_string(
    "SearchResult{id: " + self.id + ", score: " + self.score.to_string() + "}",
  )
}

///|
fn sort_results(results : Array[SearchResult], ascending : Bool) -> Unit {
  let n = results.length()
  for i = 0; i < n; i = i + 1 {
    for j = i + 1; j < n; j = j + 1 {
      let swap = if ascending {
        results[i].score > results[j].score
      } else {
        results[i].score < results[j].score
      }
      if swap {
        let temp = results[i]
        results[i] = results[j]
        results[j] = temp
      }
    }
  }
}

///|
fn matches_filters(
  metadata : Array[(String, String)],
  filters : Array[(String, String)],
) -> Bool {
  if filters.length() == 0 {
    return true
  }
  for f in filters {
    let mut found = false
    for m in metadata {
      if m.0 == f.0 && m.1 == f.1 {
        found = true
        break
      }
    }
    if !found {
      return false
    }
  }
  true
}

// Flat baseline index

///|
pub struct FlatIndex {
  documents : Array[Document]
}

///|
pub fn FlatIndex::new() -> FlatIndex {
  { documents: [] }
}

///|
pub fn FlatIndex::add(self : FlatIndex, doc : Document) -> Unit {
  self.documents.push(doc)
}

///|
pub fn FlatIndex::search(
  self : FlatIndex,
  query : Array[Double],
  top_k : Int,
  metric : DistanceMetric,
  filters : Array[(String, String)],
) -> Array[SearchResult] raise VectorError {
  let results = []
  for doc in self.documents {
    if matches_filters(doc.metadata, filters) {
      let score = calculate_distance(query, doc.vector, metric)
      results.push({ id: doc.id, score, metadata: doc.metadata })
    }
  }

  let ascending = match metric {
    Euclidean | Manhattan => true
    Cosine | DotProduct => false
  }
  sort_results(results, ascending)

  let limit = if top_k < results.length() { top_k } else { results.length() }
  let top = []
  for i = 0; i < limit; i = i + 1 {
    top.push(results[i])
  }
  top
}

// IVF-Flat Index (Inverted File)

///|
pub struct IvfIndex {
  mut centroids : Array[Array[Double]]
  inverted_lists : Array[Array[Document]]
  k : Int
  metric : DistanceMetric
}

///|
pub fn IvfIndex::new(k : Int, metric : DistanceMetric) -> IvfIndex {
  let lists = Array::make(k, [])
  for i = 0; i < k; i = i + 1 {
    lists[i] = []
  }
  { centroids: [], inverted_lists: lists, k, metric }
}

///|
pub fn IvfIndex::build(
  self : IvfIndex,
  docs : Array[Document],
) -> Unit raise VectorError {
  if docs.length() == 0 {
    return
  }
  let vectors = []
  for d in docs {
    vectors.push(d.vector)
  }

  self.centroids = kmeans(vectors, self.k, 20)

  for i = 0; i < self.k; i = i + 1 {
    self.inverted_lists[i].clear()
  }

  for doc in docs {
    let mut min_dist = -1.0
    let mut best_centroid = 0
    for c = 0; c < self.k; c = c + 1 {
      let dist = euclidean_distance(doc.vector, self.centroids[c])
      if min_dist < 0.0 || dist < min_dist {
        min_dist = dist
        best_centroid = c
      }
    }
    self.inverted_lists[best_centroid].push(doc)
  }
}

///|
pub fn IvfIndex::search(
  self : IvfIndex,
  query : Array[Double],
  top_k : Int,
  nprobe : Int,
  filters : Array[(String, String)],
) -> Array[SearchResult] raise VectorError {
  if self.centroids.length() == 0 {
    raise IndexError("IVF index is empty or not built")
  }
  let np = if nprobe > self.k { self.k } else { nprobe }

  let centroid_dists = []
  for c = 0; c < self.k; c = c + 1 {
    let dist = euclidean_distance(query, self.centroids[c])
    centroid_dists.push((c, dist))
  }

  let n_centroids = centroid_dists.length()
  for i = 0; i < n_centroids; i = i + 1 {
    for j = i + 1; j < n_centroids; j = j + 1 {
      if centroid_dists[i].1 > centroid_dists[j].1 {
        let temp = centroid_dists[i]
        centroid_dists[i] = centroid_dists[j]
        centroid_dists[j] = temp
      }
    }
  }

  let results = []
  for i = 0; i < np; i = i + 1 {
    let centroid_idx = centroid_dists[i].0
    let list = self.inverted_lists[centroid_idx]
    for doc in list {
      if matches_filters(doc.metadata, filters) {
        let score = calculate_distance(query, doc.vector, self.metric)
        results.push({ id: doc.id, score, metadata: doc.metadata })
      }
    }
  }

  let ascending = match self.metric {
    Euclidean | Manhattan => true
    Cosine | DotProduct => false
  }
  sort_results(results, ascending)

  let limit = if top_k < results.length() { top_k } else { results.length() }
  let top = []
  for i = 0; i < limit; i = i + 1 {
    top.push(results[i])
  }
  top
}

// KD-Tree Index

///|
pub enum KdNode {
  Empty
  Node(Int, Document, KdNode, KdNode)
}

///|
pub struct KdTreeIndex {
  mut root : KdNode
  metric : DistanceMetric
}

///|
pub fn KdTreeIndex::new(metric : DistanceMetric) -> KdTreeIndex {
  { root: Empty, metric }
}

///|
fn build_kd_node(docs : Array[Document], depth : Int, dim : Int) -> KdNode {
  if docs.length() == 0 {
    return Empty
  }

  let axis = depth % dim
  let n = docs.length()
  for i = 0; i < n; i = i + 1 {
    for j = i + 1; j < n; j = j + 1 {
      if docs[i].vector[axis] > docs[j].vector[axis] {
        let temp = docs[i]
        docs[i] = docs[j]
        docs[j] = temp
      }
    }
  }

  let median = n / 2
  let doc = docs[median]

  let left_docs = []
  for i = 0; i < median; i = i + 1 {
    left_docs.push(docs[i])
  }

  let right_docs = []
  for i = median + 1; i < n; i = i + 1 {
    right_docs.push(docs[i])
  }

  let left = build_kd_node(left_docs, depth + 1, dim)
  let right = build_kd_node(right_docs, depth + 1, dim)

  Node(axis, doc, left, right)
}

///|
pub fn KdTreeIndex::build(self : KdTreeIndex, docs : Array[Document]) -> Unit {
  if docs.length() == 0 {
    self.root = Empty
    return
  }
  let dim = docs[0].vector.length()
  self.root = build_kd_node(docs, 0, dim)
}

///|
fn add_to_best(
  best : Array[SearchResult],
  doc : Document,
  score : Double,
  top_k : Int,
  metric : DistanceMetric,
) -> Unit {
  let is_better = if best.length() < top_k {
    true
  } else {
    let worst_idx = best.length() - 1
    let worst_score = best[worst_idx].score
    match metric {
      Euclidean | Manhattan => score < worst_score
      Cosine | DotProduct => score > worst_score
    }
  }

  if is_better {
    let result = { id: doc.id, score, metadata: doc.metadata }
    if best.length() < top_k {
      best.push(result)
    } else {
      let worst_idx = best.length() - 1
      best[worst_idx] = result
    }

    let ascending = match metric {
      Euclidean | Manhattan => true
      Cosine | DotProduct => false
    }
    sort_results(best, ascending)
  }
}

///|
fn search_kd_node(
  node : KdNode,
  query : Array[Double],
  top_k : Int,
  metric : DistanceMetric,
  filters : Array[(String, String)],
  best : Array[SearchResult],
) -> Unit raise VectorError {
  match node {
    Empty => ()
    Node(axis, doc, left, right) => {
      let dist = calculate_distance(query, doc.vector, metric)

      if matches_filters(doc.metadata, filters) {
        add_to_best(best, doc, dist, top_k, metric)
      }

      let query_val = query[axis]
      let node_val = doc.vector[axis]

      let (closer_branch, further_branch) = if query_val < node_val {
        (left, right)
      } else {
        (right, left)
      }

      search_kd_node(closer_branch, query, top_k, metric, filters, best)

      let plane_dist = if query_val < node_val {
        node_val - query_val
      } else {
        query_val - node_val
      }

      let search_further = if best.length() < top_k {
        true
      } else {
        let worst_idx = best.length() - 1
        let worst_score = best[worst_idx].score
        match metric {
          Euclidean | Manhattan => plane_dist < worst_score
          Cosine | DotProduct => true
        }
      }

      if search_further {
        search_kd_node(further_branch, query, top_k, metric, filters, best)
      }
    }
  }
}

///|
pub fn KdTreeIndex::search(
  self : KdTreeIndex,
  query : Array[Double],
  top_k : Int,
  filters : Array[(String, String)],
) -> Array[SearchResult] raise VectorError {
  if top_k <= 0 {
    return []
  }
  let best = []
  search_kd_node(self.root, query, top_k, self.metric, filters, best)
  let ascending = match self.metric {
    Euclidean | Manhattan => true
    Cosine | DotProduct => false
  }
  sort_results(best, ascending)
  best
}