// Licensed to the LF AI & Data foundation under one
// or more contributor license agreements. See the NOTICE file
// distributed with this work for additional information
// regarding copyright ownership. The ASF licenses this file
// to you under the Apache License, Version 2.0 (the
// "License"); you may not use this file except in compliance
// with the License. You may obtain a copy of the License at
//
//     http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
//
// 移植自 milvus-io/milvus client/milvusclient/{read.go,read_options.go}
// 的 `Search` / `AnnRequest.searchRequest` 那一路(Apache-2.0)。
//
// 上游把 `SearchResults.results`(`schema.SearchResultData`)按 `topks`
// 切成每个 query 一个 `ResultSet`,每个 `ResultSet` 里再按输出字段摊成
// 列。本移植保留「按 query 切」这一层,但不把自己伪装成带 schema 的
// `ResultSet` —— 输出字段的列直接从响应里按名字取。

///|
/// `search_params` 里服务端认的键。字面量与上游逐条对齐。
pub let search_param_anns_field : String = "anns_field"

///|
pub let search_param_topk : String = "topk"

///|
pub let search_param_offset : String = "offset"

///|
pub let search_param_metric_type : String = "metric_type"

///|
pub let search_param_round_decimal : String = "round_decimal"

///|
pub let search_param_ignore_growing : String = "ignore_growing"

///|
pub let search_param_params : String = "params"

///|
/// 一条命中。`id` 是主键的文本形式:整数主键与字符串主键都收敛到
/// `String`,因为 `Hit` 的消费方几乎总是拿它去查行,而不是做算术。
///
/// `string_id` 只对字符串主键为真,用来区分「主键就是 `"12"`」和
/// 「整数主键 12」。
pub(all) struct SearchHit {
  id : String
  string_id : Bool
  score : Float
  fields : Map[String, @column.ColumnValue]
} derive(Debug)

///|
/// 一次 query 的全部命中,按分数降序(服务端保证的距离序)。
pub(all) struct SearchResult {
  hits : Array[SearchHit]
} derive(Debug)

///|
/// 检索的入参。对应上游 `searchOption` + `AnnRequest` 里本移植覆盖的那部分。
pub(all) struct SearchOption {
  collection_name : String
  partition_names : Array[String]
  anns_field : String
  limit : Int
  vectors : Array[@column.ColumnValue]
  expr : String
  output_fields : Array[String]
  metric_type : @index.MetricType?
  search_params : Array[(String, String)]
  offset : Int
  ignore_growing : Bool
  round_decimal : Int
  consistency_level : ConsistencyLevel?
}

///|
/// 记一个查询向量。每个元素是「一行」,也就是说:
/// `FloatVector(dim, [v1, v2])` 表示 nq = 2、每行 dim 维。
pub fn new_search_option(
  collection_name : String,
  limit : Int,
  vectors : Array[@column.ColumnValue],
) -> SearchOption {
  {
    collection_name,
    partition_names: [],
    anns_field: "",
    limit,
    vectors,
    expr: "",
    output_fields: [],
    metric_type: None,
    search_params: [],
    offset: 0,
    ignore_growing: false,
    round_decimal: -1,
    consistency_level: None,
  }
}

///|
pub fn SearchOption::with_anns_field(
  self : SearchOption,
  anns_field : String,
) -> SearchOption {
  { ..self, anns_field, }
}

///|
pub fn SearchOption::with_filter(
  self : SearchOption,
  expr : String,
) -> SearchOption {
  { ..self, expr, }
}

///|
pub fn SearchOption::with_output_fields(
  self : SearchOption,
  output_fields : Array[String],
) -> SearchOption {
  { ..self, output_fields, }
}

///|
pub fn SearchOption::with_partitions(
  self : SearchOption,
  partition_names : Array[String],
) -> SearchOption {
  { ..self, partition_names, }
}

///|
pub fn SearchOption::with_metric_type(
  self : SearchOption,
  metric_type : @index.MetricType,
) -> SearchOption {
  { ..self, metric_type: Some(metric_type), }
}

///|
pub fn SearchOption::with_search_param(
  self : SearchOption,
  key : String,
  value : String,
) -> SearchOption {
  let search_params : Array[(String, String)] = []
  let mut replaced = false
  for pair in self.search_params {
    if pair.0 == key {
      search_params.push((key, value))
      replaced = true
    } else {
      search_params.push(pair)
    }
  }
  if !replaced {
    search_params.push((key, value))
  }
  { ..self, search_params, }
}

///|
pub fn SearchOption::with_offset(
  self : SearchOption,
  offset : Int,
) -> SearchOption {
  { ..self, offset, }
}

///|
pub fn SearchOption::with_ignore_growing(
  self : SearchOption,
  ignore_growing? : Bool = true,
) -> SearchOption {
  { ..self, ignore_growing, }
}

///|
pub fn SearchOption::with_round_decimal(
  self : SearchOption,
  round_decimal : Int,
) -> SearchOption {
  { ..self, round_decimal, }
}

///|
pub fn SearchOption::with_consistency_level(
  self : SearchOption,
  level : ConsistencyLevel,
) -> SearchOption {
  { ..self, consistency_level: Some(level), }
}

///|
/// 占位符对应的 `common.PlaceholderType`。
///
/// 只有第一行决定类型 —— 与上游 `vector2Placeholder` 一致,
/// 混合类型的一批向量会被服务端按第一批的类型解析,早点在本地报错更好。
fn placeholder_type(
  vectors : Array[@column.ColumnValue],
) -> @common.PlaceholderType raise ClientError {
  if vectors.length() == 0 {
    raise ClientError::Encode("search requires at least one query vector")
  }
  match vectors[0] {
    FloatVector(_, _) => @common.PlaceholderType::FloatVector
    Float16Vector(_, _) => @common.PlaceholderType::Float16Vector
    BFloat16Vector(_, _) => @common.PlaceholderType::BFloat16Vector
    BinaryVector(_, _) => @common.PlaceholderType::BinaryVector
    Int8Vector(_, _) => @common.PlaceholderType::Int8Vector
    SparseFloatVector(_) => @common.PlaceholderType::SparseFloatVector
    VarChar(_) | String(_) | Text(_) => @common.PlaceholderType::VarChar
    _ =>
      raise ClientError::Encode(
        "\{column_value_type_name(vectors[0])} is not a searchable vector type",
      )
  }
}

///|
/// 一行的序列化形态。定长向量(float / fp16 / bf16 / binary / int8)在
/// wire 上是「每行一块字节」,所以这里逐行编码;稀疏向量本来就是
/// 行内不定长,`to_bytes` 已经是行编码。
fn serialize_row(
  column : @column.ColumnValue,
  index : Int,
) -> Bytes raise ClientError {
  match column {
    FloatVector(dim, rows) =>
      floats_row("FloatVector", dim, rows, index, row_to_le_bytes)
    Float16Vector(dim, rows) =>
      floats_row(
        "Float16Vector", dim, rows, index, @entity.float16_vector_bytes,
      )
    BFloat16Vector(dim, rows) =>
      floats_row(
        "BFloat16Vector", dim, rows, index, @entity.bfloat16_vector_bytes,
      )
    BinaryVector(dim, rows) => bytes_row("BinaryVector", dim / 8, rows, index)
    Int8Vector(dim, rows) => bytes_row("Int8Vector", dim, rows, index)
    SparseFloatVector(rows) => {
      if index >= rows.length() {
        raise ClientError::Encode(
          "SparseFloatVector: row \{index} is out of range (have \{rows.length()})",
        )
      }
      rows[index].to_bytes() catch {
        err => raise ClientError::Schema(entity_schema_message(err))
      }
    }
    _ =>
      raise ClientError::Encode(
        "\{column_value_type_name(column)} cannot be used as a query vector",
      )
  }
}

///|
fn floats_row(
  name : String,
  dim : Int,
  rows : Array[Array[Float]],
  index : Int,
  encode : (Array[Float]) -> Bytes,
) -> Bytes raise ClientError {
  if index >= rows.length() {
    raise ClientError::Encode(
      "\{name}: vector row \{index} is out of range (have \{rows.length()})",
    )
  }
  let row = rows[index]
  if row.length() != dim {
    raise ClientError::Encode(
      "\{name}: row \{index} has \{row.length()} values, expected \{dim}",
    )
  }
  encode(row)
}

///|
/// 定长字节行:逐行取一段。
fn bytes_row(
  name : String,
  row_bytes : Int,
  rows : Array[Bytes],
  index : Int,
) -> Bytes raise ClientError {
  if index >= rows.length() {
    raise ClientError::Encode(
      "\{name}: vector row \{index} is out of range (have \{rows.length()})",
    )
  }
  let row = rows[index]
  if row.length() != row_bytes {
    raise ClientError::Encode(
      "\{name}: row \{index} has \{row.length()} bytes, expected \{row_bytes}",
    )
  }
  row
}

///|
/// `Array[Float]` 的小端 float32 编码。
fn row_to_le_bytes(row : Array[Float]) -> Bytes {
  let buf = Buffer()
  for value in row {
    let bits = value.reinterpret_as_uint()
    buf.write_byte((bits & 0xFFU).reinterpret_as_int().to_byte())
    buf.write_byte(((bits >> 8) & 0xFFU).reinterpret_as_int().to_byte())
    buf.write_byte(((bits >> 16) & 0xFFU).reinterpret_as_int().to_byte())
    buf.write_byte(((bits >> 24) & 0xFFU).reinterpret_as_int().to_byte())
  }
  buf.to_bytes()
}

///|
/// 编码成 `PlaceholderGroup` 的字节,塞进 `SearchRequest.placeholder_group`。
fn placeholder_group_bytes(
  vectors : Array[@column.ColumnValue],
) -> Bytes raise ClientError {
  let type_ = placeholder_type(vectors)
  let values : Array[Bytes] = []
  for i = 0; i < vectors.length(); i = i + 1 {
    values.push(serialize_row(vectors[i], i))
  }
  let group = @common.PlaceholderGroup::{
    placeholders: [@common.PlaceholderValue::{ tag: "$0", type_, values, }],
  }
  encode_request(group)
}

///|
/// `search_params`。上游总是把 `anns_field` / `topk` / `offset` /
/// `metric_type` / `round_decimal` / `ignore_growing` / `params` 这七个键
/// 写全(`metric_type` 为空时写空串),调用方的 `WithSearchParam`
/// 最后覆盖。这里保持同样的顺序与键集合,方便与服务端的行为对照。
fn search_params_of(option : SearchOption) -> Array[(String, String)] {
  let params : Array[(String, String)] = []
  let seen : Array[String] = []
  let push = (key : String, value : String) => {
    params.push((key, value))
    seen.push(key)
  }
  push(search_param_anns_field, option.anns_field)
  push(search_param_topk, option.limit.to_string())
  push(search_param_offset, option.offset.to_string())
  push(
    search_param_metric_type,
    match option.metric_type {
      Some(m) => m.to_string()
      None => ""
    },
  )
  push(search_param_round_decimal, option.round_decimal.to_string())
  push(search_param_ignore_growing, option.ignore_growing.to_string())
  push(search_param_params, "{}")
  for pair in option.search_params {
    if !seen.contains(pair.0) {
      params.push(pair)
    }
  }
  params
}

///|
/// 检索。
pub async fn Client::search(
  self : Client,
  option : SearchOption,
) -> SearchResult raise ClientError {
  decode_search_results(self.search_call(option).results.unwrap())
}

///|
/// 编排一次 `Search`:校验、编占位符、拼 `search_params`、发出去、查 status。
///
/// 返回值留着整个 `SearchResults`,因为 `results` 上有迭代器要的续页凭据。
/// 搜索本身与迭代器共用这一份,免得两处对「七个键的顺序」「use_default_consistency
/// 什么时候为真」各说各话。
async fn Client::search_call(
  self : Client,
  option : SearchOption,
) -> @milvus.SearchResults raise ClientError {
  if option.vectors.length() == 0 {
    raise ClientError::Encode("search requires at least one query vector")
  }
  // `nq` 是 `vectors` 的长度,服务端不接受「nq = 0」。迭代器收尾要发的
  // 那个空检索走下面的 `cancel_search` —— 它另开一个入口,因为它的
  // `search_input` 不是占位符。
  let placeholder = placeholder_group_bytes(option.vectors)
  let search_params = search_params_of(option).map(pair => {
    @common.KeyValuePair::KeyValuePair(pair.0, pair.1)
  })
  let consistency = match option.consistency_level {
    Some(level) => level.to_proto()
    None => @common.ConsistencyLevel::Bounded
  }
  let request = @milvus.SearchRequest::{
    base: None,
    db_name: "",
    collection_name: option.collection_name,
    partition_names: option.partition_names,
    dsl: option.expr,
    search_input: @milvus.SearchRequest_SearchInput::PlaceholderGroup(
      placeholder,
    ),
    dsl_type: @common.DslType::BoolExprV1,
    output_fields: option.output_fields,
    search_params,
    travel_timestamp: 0UL,
    guarantee_timestamp: 0UL,
    nq: option.vectors.length().to_int64(),
    not_return_all_meta: false,
    consistency_level: consistency,
    use_default_consistency: option.consistency_level is None,
    namespace_: None,
  }
  let response : @milvus.SearchResults = self.call_service(search_path, request)
  match check_status(response.status) {
    Some(err) => raise err
    None => response
  }
}

///|
/// 迭代器收尾用的空检索:`nq = 0`,`search_input` 显式留空(`NotSet`)。
///
/// 服务端据 `search_iter_id` 找到那个游标会话,`nq = 0` 就表示「这批不要
/// 结果,把会话收掉」。走不了 `Client::search`:那条路要求至少一个查询向量,
/// 会用占位符填 `search_input`。
///
/// `search_params` 由调用方拼好传进来 —— 迭代器比这一层更清楚要带哪些凭据。
///
/// 包内可见:`SearchIterator` 的收尾用它,不属于门面 API。
async fn Client::cancel_search(
  self : Client,
  option : SearchOption,
) -> Unit raise ClientError {
  let search_params = option.search_params.map(pair => {
    @common.KeyValuePair::KeyValuePair(pair.0, pair.1)
  })
  let consistency = match option.consistency_level {
    Some(level) => level.to_proto()
    None => @common.ConsistencyLevel::Bounded
  }
  let request = @milvus.SearchRequest::{
    base: None,
    db_name: "",
    collection_name: option.collection_name,
    partition_names: option.partition_names,
    dsl: option.expr,
    search_input: @milvus.SearchRequest_SearchInput::NotSet,
    dsl_type: @common.DslType::BoolExprV1,
    output_fields: [],
    search_params,
    travel_timestamp: 0UL,
    guarantee_timestamp: 0UL,
    nq: 0L,
    not_return_all_meta: false,
    consistency_level: consistency,
    use_default_consistency: option.consistency_level is None,
    namespace_: None,
  }
  let response : @milvus.SearchResults = self.call_service(search_path, request)
  match check_status(response.status) {
    Some(err) => raise err
    None => ()
  }
}

///|
/// 一次检索的原始响应 + 解好的命中。
///
/// 迭代器两条都要:命中给调用方,`SearchResultData` 上的
/// `search_iterator_v2_results` 是翻下一页的凭据,`SearchResult` 装不下它。
///
/// 包内可见,不对外 —— 外面拿到的是 `SearchResult`,需要续页凭据的是
/// 迭代器,而它在同一个包里。
struct RawSearchResult {
  data : @schema.SearchResultData
  result : SearchResult
} derive(Debug)

///|
/// 与 `Client::search` 同一条路径,只是把原始 `SearchResultData` 一并带回来。
/// 编排逻辑只有一份,见 `search_call`。包内可见,调用方是 `SearchIterator`。
async fn Client::search_raw(
  self : Client,
  option : SearchOption,
) -> RawSearchResult raise ClientError {
  let response = self.search_call(option)
  let data = response.results.unwrap()
  { data, result: decode_search_results(data), }
}

///|
/// `SearchResultData` → `SearchResult`。
///
/// 上游按 `topks` 把扁平的命中切成 nq 段;本移植把 nq 段再摊平回一条
/// `hits`,因为 `SearchByIDsOption` 之类「一次多查」的用例在本地几乎总是
/// 逐条处理。切段信息(`topks`)保留在响应里,需要时再切。
fn decode_search_results(
  data : @schema.SearchResultData,
) -> SearchResult raise ClientError {
  let total = match data.topks {
    [] => 0
    topks => {
      let mut sum = 0
      for k in topks {
        if k < 0L {
          raise ClientError::Decode("server returned a negative topk \{k}")
        }
        sum += k.to_int()
      }
      sum
    }
  }
  if data.scores.length() < total {
    raise ClientError::Decode(
      "server returned \{data.scores.length()} scores for \{total} hits",
    )
  }
  let ids = id_strings(data.ids, total)
  let score_fields : Map[String, @column.ColumnValue] = Map([])
  for field in data.fields_data {
    let column = @column.from_field_data(field) catch {
      err => raise ClientError::Decode(column_error_message(err))
    }
    score_fields[field.field_name] = column.column
  }
  let hits : Array[SearchHit] = []
  for i = 0; i < total; i = i + 1 {
    let (id, string_id) = ids[i]
    let fields : Map[String, @column.ColumnValue] = Map([])
    for pair in score_fields {
      fields[pair.0] = pair.1
    }
    hits.push({ id, string_id, score: data.scores[i], fields, })
  }
  { hits, }
}

///|
/// `schema.IDs` → 每行的 `(文本, 是否字符串主键)`。
///
/// 字符串主键的 `string_id` 为真,好让调用方区分 `"12"` 与 `12`。
fn id_strings(
  ids : @schema.IDs?,
  total : Int,
) -> Array[(String, Bool)] raise ClientError {
  let out : Array[(String, Bool)] = []
  match ids {
    None => ()
    Some(ids) =>
      match ids.id_field {
        @schema.IDs_IdField::IntId(a) =>
          for i = 0; i < total; i = i + 1 {
            if i >= a.data.length() {
              raise ClientError::Decode(
                "server returned \{a.data.length()} int ids for \{total} hits",
              )
            }
            out.push((a.data[i].to_string(), false))
          }
        @schema.IDs_IdField::StrId(a) =>
          for i = 0; i < total; i = i + 1 {
            if i >= a.data.length() {
              raise ClientError::Decode(
                "server returned \{a.data.length()} string ids for \{total} hits",
              )
            }
            out.push((a.data[i], true))
          }
        @schema.IDs_IdField::UuidId(a) =>
          for i = 0; i < total; i = i + 1 {
            if i >= a.data.length() {
              raise ClientError::Decode(
                "server returned \{a.data.length()} uuid ids for \{total} hits",
              )
            }
            out.push((bytes_to_hex(a.data[i]), true))
          }
        @schema.IDs_IdField::NotSet => ()
      }
  }
  if out.length() < total {
    raise ClientError::Decode(
      "server returned \{out.length()} ids for \{total} hits",
    )
  }
  out
}

///|
/// 16 字节 UUID 的十六进制写法。`@column` 回读时按字节保留,
/// 这里也保持同样的形态,不做 8-4-4-4-12 的加连字符美化 ——
/// 两种写法服务端都收,但一致比好看重要。
fn bytes_to_hex(bytes : Bytes) -> String {
  let digits = "0123456789abcdef"
  let parts : Array[String] = []
  for byte in bytes {
    let v = byte.to_uint()
    let high = (v >> 4).reinterpret_as_int()
    let low = (v & 0xFU).reinterpret_as_int()
    parts.push(digits[high:high + 1].to_owned())
    parts.push(digits[low:low + 1].to_owned())
  }
  parts.join("")
}