use super::query_contract::{query_metric, validate_query_contract};
use crate::doc::{Doc, DocumentMap, VectorValue};
use crate::error::{Error, Result};
use crate::filter::FilterExpr;
use crate::index::{CandidateSelection, IndexRegistry, OrdinalScores};
use crate::query::{FtsDefaultOperator, SearchQuery};
use crate::schema::{CollectionSchema, IndexParams};
use crate::text::{
bm25_term_score, contains_ordered_phrase, parse_fts_query, text_value, FtsEvalContext,
Tokenizer,
};
use crate::types::{DataType, MetricType};
use serde_json::Value;
use std::cmp::Ordering;
use std::collections::{BTreeMap, BinaryHeap};
struct ScoredDoc {
exact_score: f64,
doc: Doc,
}
struct ScoredCandidate<'a> {
exact_score: f64,
doc: &'a Doc,
}
enum ResolvedQueryVector {
Dense { values: Vec<f64>, norm: f64 },
Binary(Vec<u8>),
Sparse(BTreeMap<u32, f64>),
}
struct TopKCollector<'a> {
limit: usize,
candidates: BinaryHeap<ScoredCandidate<'a>>,
}
impl ScoredDoc {
fn new(exact_score: f64, doc: Doc) -> Result<Self> {
if !exact_score.is_finite() {
return Err(Error::resource_exhausted("query score is not finite"));
}
Ok(Self { exact_score, doc })
}
}
impl ScoredCandidate<'_> {
fn new(exact_score: f64, doc: &Doc) -> Result<ScoredCandidate<'_>> {
if !exact_score.is_finite() {
return Err(Error::resource_exhausted("query score is not finite"));
}
Ok(ScoredCandidate { exact_score, doc })
}
fn id(&self) -> &str {
self.doc.get_pk().unwrap_or_default()
}
}
impl PartialEq for ScoredCandidate<'_> {
fn eq(&self, other: &Self) -> bool {
self.cmp(other) == Ordering::Equal
}
}
impl Eq for ScoredCandidate<'_> {}
impl PartialOrd for ScoredCandidate<'_> {
fn partial_cmp(&self, other: &Self) -> Option<Ordering> {
Some(self.cmp(other))
}
}
impl Ord for ScoredCandidate<'_> {
fn cmp(&self, other: &Self) -> Ordering {
other
.exact_score
.total_cmp(&self.exact_score)
.then_with(|| self.id().cmp(other.id()))
}
}
impl<'a> TopKCollector<'a> {
fn new(limit: usize) -> Self {
Self {
limit,
candidates: BinaryHeap::new(),
}
}
fn push(&mut self, exact_score: f64, doc: &'a Doc) -> Result<()> {
let candidate = ScoredCandidate::new(exact_score, doc)?;
if self.limit == 0 {
return Ok(());
}
if self.candidates.len() < self.limit {
self.candidates.push(candidate);
return Ok(());
}
if self
.candidates
.peek()
.is_some_and(|worst| candidate.cmp(worst) == Ordering::Less)
{
self.candidates.pop();
self.candidates.push(candidate);
}
Ok(())
}
fn into_scored_docs(self) -> Result<Vec<ScoredDoc>> {
self.candidates
.into_iter()
.map(|candidate| ScoredDoc::new(candidate.exact_score, candidate.doc.clone()))
.collect()
}
}
#[allow(clippy::cast_possible_truncation)]
pub(super) fn score_to_f32(value: f64) -> Result<f32> {
if !value.is_finite() || value < f64::from(f32::MIN) || value > f64::from(f32::MAX) {
return Err(Error::resource_exhausted(
"query score cannot be represented as f32",
));
}
Ok(value as f32)
}
#[allow(clippy::cast_precision_loss)]
pub(super) fn count_to_f64(value: usize) -> f64 {
value as f64
}
fn binary_score(query: &[u8], vector: &VectorValue) -> Option<f64> {
let (VectorValue::Binary32(stored) | VectorValue::Binary64(stored)) = vector else {
return None;
};
if stored.len() != query.len() {
return None;
}
Some(
-query
.iter()
.zip(stored)
.map(|(left, right)| f64::from((left ^ right).count_ones()))
.sum::<f64>(),
)
}
pub(super) fn sort_docs(docs: &mut [Doc]) {
docs.sort_by(|left, right| {
right
.get_score()
.partial_cmp(&left.get_score())
.unwrap_or(Ordering::Equal)
.then_with(|| {
left.get_pk()
.unwrap_or_default()
.cmp(right.get_pk().unwrap_or_default())
})
});
}
fn sort_scored_docs(docs: &mut [ScoredDoc]) {
docs.sort_by(|left, right| {
right
.exact_score
.partial_cmp(&left.exact_score)
.unwrap_or(Ordering::Equal)
.then_with(|| {
left.doc
.get_pk()
.unwrap_or_default()
.cmp(right.doc.get_pk().unwrap_or_default())
})
});
}
pub(super) fn parse_filter_expression(expression: &str) -> Result<FilterExpr> {
if expression.trim().is_empty() {
return Err(Error::invalid_argument(
"filter expression must not be empty",
));
}
crate::filter::parse_filter(expression)
}
pub(super) fn parse_optional_filter(expression: Option<&str>) -> Result<Option<FilterExpr>> {
expression.map(parse_filter_expression).transpose()
}
pub(super) fn matches_filter(doc: &Doc, filter: Option<&FilterExpr>) -> bool {
filter.map_or(true, |filter| filter.matches(doc))
}
pub(super) fn execute_query_with_candidates(
schema: &CollectionSchema,
docs: &DocumentMap,
indexes: &IndexRegistry,
query: &SearchQuery,
candidate_ids: Option<&CandidateSelection>,
fts_scores: Option<&OrdinalScores>,
filter: Option<&FilterExpr>,
) -> Result<Vec<Doc>> {
let field = validate_query_contract(schema, query)?;
let metric = query_metric(&field, query)?;
let topk = usize::try_from(query.topk)
.map_err(|_| Error::invalid_argument("query topk must be positive"))?;
let mut scored = if query.fts.is_some() {
execute_fts(
docs,
query,
field.index_params,
candidate_ids,
fts_scores,
filter,
topk,
)?
} else {
execute_vector(
docs,
indexes,
query,
metric,
candidate_ids,
filter,
topk,
schema.vectors.iter().any(|field| {
field.name == query.field_name && field.data_type == DataType::VectorFp32
}),
)?
};
sort_scored_docs(&mut scored);
scored.truncate(topk);
let output_fields = query.output_fields.as_deref();
scored
.into_iter()
.map(|mut scored| {
scored.doc.set_score(score_to_f32(scored.exact_score)?)?;
if query.include_doc_id {
let pk = scored
.doc
.get_pk()
.ok_or_else(|| Error::internal("query result has no primary key"))?;
let doc_id = indexes.document_ordinal(pk).ok_or_else(|| {
Error::internal(format!(
"query result primary key '{pk}' has no document ID"
))
})?;
scored.doc.set_internal_id(Some(doc_id));
} else {
scored.doc.set_internal_id(None);
}
Ok(scored.doc.project(output_fields, query.include_vector))
})
.collect()
}
#[allow(clippy::too_many_arguments)]
fn execute_vector(
docs: &DocumentMap,
indexes: &IndexRegistry,
query: &SearchQuery,
metric: MetricType,
candidate_ids: Option<&CandidateSelection>,
filter: Option<&FilterExpr>,
topk: usize,
field_is_fp32: bool,
) -> Result<Vec<ScoredDoc>> {
let query_vector = resolve_query_vector(docs, query)?;
let radius = query.params.get("radius").and_then(Value::as_f64);
if let (Some(candidate_ids), Some(query_f32), true, None) = (
candidate_ids,
query.vector.as_deref(),
field_is_fp32,
filter,
) {
if let ResolvedQueryVector::Dense { norm, .. } = &query_vector {
if let Some(scored) = rerank_unquantized_candidates(
docs,
indexes,
candidate_ids,
&query.field_name,
query_f32,
*norm,
metric,
radius,
topk,
)? {
return Ok(scored);
}
}
}
let mut result = TopKCollector::new(topk);
if let Some(candidate_ids) = candidate_ids {
for id in candidate_ids.ids() {
if let Some(doc) = docs.get(id) {
score_vector_document(
&mut result,
doc,
query,
metric,
&query_vector,
radius,
filter,
)?;
}
}
} else {
for doc in docs.values() {
score_vector_document(
&mut result,
doc,
query,
metric,
&query_vector,
radius,
filter,
)?;
}
}
result.into_scored_docs()
}
#[allow(clippy::too_many_arguments)]
fn rerank_unquantized_candidates(
docs: &DocumentMap,
indexes: &IndexRegistry,
candidate_ids: &CandidateSelection,
field: &str,
query: &[f32],
query_norm: f64,
metric: MetricType,
radius: Option<f64>,
topk: usize,
) -> Result<Option<Vec<ScoredDoc>>> {
if topk == 0 {
return Ok(Some(Vec::new()));
}
let mut ranked = BinaryHeap::new();
for ordinal in candidate_ids.iter_ordinals() {
let Some(exact_score) =
indexes.exact_unquantized_f32_score_at(field, ordinal, query, query_norm, metric)
else {
return Ok(None);
};
if radius_excludes(exact_score, radius, metric)
|| dominated_by_score(&ranked, exact_score, topk)
{
continue;
}
let Some(id) = candidate_ids.id(ordinal) else {
return Ok(None);
};
retain_ranked(
&mut ranked,
RankedId {
exact_score,
id,
ordinal,
},
topk,
);
}
let mut scored = Vec::with_capacity(ranked.len());
for candidate in ranked {
let Some(doc) = docs.get(candidate.id) else {
return Ok(None);
};
scored.push(ScoredDoc::new(candidate.exact_score, doc.as_ref().clone())?);
}
Ok(Some(scored))
}
fn radius_excludes(score: f64, radius: Option<f64>, metric: MetricType) -> bool {
radius.is_some_and(|radius| {
if metric == MetricType::L2 {
score < -radius * radius
} else {
score < radius
}
})
}
struct RankedId<'a> {
exact_score: f64,
id: &'a str,
ordinal: u64,
}
impl PartialEq for RankedId<'_> {
fn eq(&self, other: &Self) -> bool {
self.cmp(other) == Ordering::Equal
}
}
impl Eq for RankedId<'_> {}
impl PartialOrd for RankedId<'_> {
fn partial_cmp(&self, other: &Self) -> Option<Ordering> {
Some(self.cmp(other))
}
}
impl Ord for RankedId<'_> {
fn cmp(&self, other: &Self) -> Ordering {
other
.exact_score
.total_cmp(&self.exact_score)
.then_with(|| self.id.cmp(other.id))
.then_with(|| self.ordinal.cmp(&other.ordinal))
}
}
fn retain_ranked<'a>(ranked: &mut BinaryHeap<RankedId<'a>>, candidate: RankedId<'a>, limit: usize) {
if ranked.len() < limit {
ranked.push(candidate);
return;
}
if ranked
.peek()
.is_some_and(|worst| candidate.cmp(worst) == Ordering::Less)
{
ranked.pop();
ranked.push(candidate);
}
}
fn dominated_by_score(ranked: &BinaryHeap<RankedId<'_>>, exact_score: f64, limit: usize) -> bool {
ranked.len() >= limit
&& ranked
.peek()
.is_some_and(|worst| exact_score.total_cmp(&worst.exact_score) == Ordering::Less)
}
fn resolve_query_vector(docs: &DocumentMap, query: &SearchQuery) -> Result<ResolvedQueryVector> {
if let Some(vector) = &query.vector {
let values: Vec<f64> = vector.iter().map(|value| f64::from(*value)).collect();
return Ok(ResolvedQueryVector::Dense {
norm: dense_query_norm(&values),
values,
});
}
if let Some(vector) = &query.binary_vector {
return Ok(ResolvedQueryVector::Binary(vector.clone()));
}
if let Some(values) = &query.sparse_vector {
return Ok(ResolvedQueryVector::Sparse(
values
.iter()
.map(|(index, value)| (*index, f64::from(*value)))
.collect(),
));
}
let id = query.id.as_deref().ok_or_else(|| {
Error::invalid_argument(
"query requires a dense vector, binary vector, sparse vector, or source id",
)
})?;
let source = docs
.get(id)
.ok_or_else(|| Error::not_found(format!("source document '{id}' not found")))?;
let vector = source.vector(&query.field_name).ok_or_else(|| {
Error::failed_precondition(format!(
"source document '{id}' has no vector in field '{}'",
query.field_name
))
})?;
if let Some(vector) = vector.to_dense_f64() {
return Ok(ResolvedQueryVector::Dense {
norm: dense_query_norm(&vector),
values: vector,
});
}
if let Some(vector) = vector.to_sparse_f64() {
return Ok(ResolvedQueryVector::Sparse(vector));
}
if let VectorValue::Binary32(values) | VectorValue::Binary64(values) = vector {
return Ok(ResolvedQueryVector::Binary(values.clone()));
}
Err(Error::failed_precondition(format!(
"source document '{id}' has no searchable vector in field '{}'",
query.field_name
)))
}
#[allow(clippy::too_many_arguments)]
fn score_vector_document<'a>(
result: &mut TopKCollector<'a>,
doc: &'a Doc,
query: &SearchQuery,
metric: MetricType,
query_vector: &ResolvedQueryVector,
radius: Option<f64>,
filter: Option<&FilterExpr>,
) -> Result<()> {
if !matches_filter(doc, filter) {
return Ok(());
}
let Some(vector) = doc.vector(&query.field_name) else {
return Ok(());
};
let score = match query_vector {
ResolvedQueryVector::Dense { values, norm } => {
let Some(score) = vector.dense_score(values, *norm, metric) else {
return Ok(());
};
score
}
ResolvedQueryVector::Binary(query) => {
let Some(score) = binary_score(query, vector) else {
return Ok(());
};
score
}
ResolvedQueryVector::Sparse(query) => {
let Some(stored) = vector.to_sparse_f64() else {
return Ok(());
};
sparse_score(query, &stored, metric)
}
};
if radius.is_some_and(|radius| {
if metric == MetricType::L2 {
score < -radius * radius
} else {
score < radius
}
}) {
return Ok(());
}
result.push(score, doc)
}
fn dense_query_norm(values: &[f64]) -> f64 {
values.iter().map(|value| value * value).sum::<f64>().sqrt()
}
fn sparse_score(
query: &BTreeMap<u32, f64>,
stored: &BTreeMap<u32, f64>,
metric: MetricType,
) -> f64 {
let dot = query
.iter()
.filter_map(|(index, value)| stored.get(index).map(|other| value * other))
.sum::<f64>();
match metric {
MetricType::L2 => {
let query_distance = query
.iter()
.map(|(index, value)| {
let difference = value - stored.get(index).copied().unwrap_or_default();
difference * difference
})
.sum::<f64>();
let stored_distance = stored
.iter()
.filter(|(index, _)| !query.contains_key(index))
.map(|(_, value)| value * value)
.sum::<f64>();
-(query_distance + stored_distance)
}
MetricType::Cosine => {
let query_norm = query
.values()
.map(|value| value * value)
.sum::<f64>()
.sqrt();
let stored_norm = stored
.values()
.map(|value| value * value)
.sum::<f64>()
.sqrt();
if query_norm == 0.0 || stored_norm == 0.0 {
0.0
} else {
dot / (query_norm * stored_norm)
}
}
MetricType::MipsL2 | MetricType::Ip | MetricType::Undefined => dot,
}
}
fn execute_fts(
docs: &DocumentMap,
query: &SearchQuery,
index_params: Option<&IndexParams>,
candidate_ids: Option<&CandidateSelection>,
indexed_scores: Option<&OrdinalScores>,
filter: Option<&FilterExpr>,
topk: usize,
) -> Result<Vec<ScoredDoc>> {
if let Some(indexed_scores) = indexed_scores {
let mut result = TopKCollector::new(topk);
for (id, score) in indexed_scores.entries() {
let Some(doc) = docs.get(id) else {
continue;
};
if matches_filter(doc, filter) {
result.push(score, doc)?;
}
}
return result.into_scored_docs();
}
let tokenizer = Tokenizer::from_index_params(index_params)?;
let mut parsed = parse_fts_query(query, &tokenizer)?;
let corpus: Vec<(&Doc, Vec<String>)> = docs
.values()
.filter_map(|doc| {
text_value(doc, &query.field_name).map(|text| (doc.as_ref(), tokenizer.tokenize(text)))
})
.collect();
if corpus.is_empty() {
return Ok(Vec::new());
}
parsed.expand_terms(
corpus
.iter()
.flat_map(|(_, tokens)| tokens.iter().map(String::as_str)),
);
let document_count = count_to_f64(corpus.len());
let average_length = corpus
.iter()
.map(|(_, tokens)| count_to_f64(tokens.len()))
.sum::<f64>()
/ document_count;
let document_frequency = document_frequencies(&corpus, parsed.all_terms());
let mut result = TopKCollector::new(topk);
for (doc, tokens) in corpus {
if candidate_ids.is_some_and(|ids| doc.get_pk().map_or(true, |id| !ids.contains(id))) {
continue;
}
if !matches_filter(doc, filter) {
continue;
}
let score = if let Some((terms, operator)) = parsed.simple() {
if operator == FtsDefaultOperator::And
&& terms.iter().any(|term| !tokens.contains(term))
{
continue;
}
bm25(
&tokens,
terms,
&document_frequency,
document_count,
average_length,
)
} else {
let mut context = ScanFtsEvalContext {
tokens: &tokens,
document_frequency: &document_frequency,
document_count,
average_length,
};
let Some(score) = parsed.score(&mut context) else {
continue;
};
score
};
if score <= 0.0 {
continue;
}
result.push(score, doc)?;
}
result.into_scored_docs()
}
struct ScanFtsEvalContext<'a> {
tokens: &'a [String],
document_frequency: &'a BTreeMap<String, usize>,
document_count: f64,
average_length: f64,
}
impl FtsEvalContext for ScanFtsEvalContext<'_> {
fn contains_term(&mut self, term: &str) -> bool {
self.tokens.iter().any(|token| token == term)
}
fn contains_phrase(&mut self, terms: &[String], slop: u32) -> bool {
contains_ordered_phrase(self.tokens, terms, slop)
}
fn term_score(&mut self, term: &str) -> f64 {
let frequency = count_to_f64(self.tokens.iter().filter(|token| *token == term).count());
if frequency == 0.0 {
return 0.0;
}
let document_frequency = self
.document_frequency
.get(term)
.copied()
.map_or(0.0, count_to_f64);
bm25_term_score(
frequency,
document_frequency,
self.document_count,
count_to_f64(self.tokens.len()),
self.average_length,
)
}
}
fn document_frequencies(
corpus: &[(&Doc, Vec<String>)],
query_terms: &[String],
) -> BTreeMap<String, usize> {
let mut frequencies = BTreeMap::new();
for term in query_terms {
frequencies.entry(term.clone()).or_insert_with(|| {
corpus
.iter()
.filter(|(_, tokens)| tokens.contains(term))
.count()
});
}
frequencies
}
fn bm25(
document_tokens: &[String],
query_terms: &[String],
document_frequencies: &BTreeMap<String, usize>,
document_count: f64,
average_length: f64,
) -> f64 {
if document_tokens.is_empty() || query_terms.is_empty() {
return 0.0;
}
let mut score = 0.0;
let document_length = count_to_f64(document_tokens.len());
for term in query_terms {
let frequency = count_to_f64(
document_tokens
.iter()
.filter(|token| *token == term)
.count(),
);
if frequency == 0.0 {
continue;
}
let document_frequency =
count_to_f64(document_frequencies.get(term).copied().unwrap_or_default());
score += bm25_term_score(
frequency,
document_frequency,
document_count,
document_length,
average_length,
);
}
score.max(0.0)
}
pub(super) fn normalize_scores(docs: &mut [Doc], method: &str) -> Result<()> {
if docs.is_empty() || method == "none" {
return Ok(());
}
let values: Vec<f64> = docs.iter().map(|doc| f64::from(doc.get_score())).collect();
match method {
"minmax" => {
let min = values.iter().copied().fold(f64::INFINITY, f64::min);
let max = values.iter().copied().fold(f64::NEG_INFINITY, f64::max);
let denominator = (max - min).max(1e-12);
for doc in docs {
let normalized = (f64::from(doc.get_score()) - min) / denominator;
doc.set_score(score_to_f32(normalized)?)?;
}
}
"zscore" => {
let mean = values.iter().sum::<f64>() / count_to_f64(values.len());
let variance = values
.iter()
.map(|value| {
let delta = *value - mean;
delta * delta
})
.sum::<f64>()
/ count_to_f64(values.len());
let denominator = variance.sqrt().max(1e-12);
for doc in docs {
let normalized = (f64::from(doc.get_score()) - mean) / denominator;
doc.set_score(score_to_f32(normalized)?)?;
}
}
_ => {
return Err(Error::invalid_argument(
"unknown score normalization method",
));
}
}
Ok(())
}
#[cfg(test)]
#[allow(clippy::bool_assert_comparison, clippy::float_cmp)]
mod tests {
use super::{
parse_filter_expression, parse_optional_filter, sort_scored_docs, ScoredCandidate,
TopKCollector,
};
use crate::doc::Doc;
use std::cmp::Ordering;
#[test]
fn bounded_topk_keeps_exact_scores_and_primary_key_ties() {
let docs: Vec<Doc> = ["z-low", "b-tie", "a-tie", "best", "c-tie"]
.into_iter()
.map(|id| Doc::with_pk(id).expect("document must be valid"))
.collect();
let mut collector = TopKCollector::new(3);
for (doc, score) in docs.iter().zip([1.0, 2.0, 2.0, 3.0, 2.0]) {
collector
.push(score, doc)
.expect("finite score must be accepted");
}
assert!(collector.push(f64::NAN, &docs[0]).is_err());
let mut ranked = collector
.into_scored_docs()
.expect("retained scores must be valid");
sort_scored_docs(&mut ranked);
assert_eq!(
ranked
.iter()
.map(|scored| scored.doc.get_pk().unwrap_or_default())
.collect::<Vec<_>>(),
vec!["best", "a-tie", "b-tie"]
);
}
#[test]
fn filter_parsers_reject_empty_and_accept_optional_none() {
assert!(parse_filter_expression("").is_err());
assert!(parse_filter_expression(" ").is_err());
assert!(parse_optional_filter(None).expect("none").is_none());
assert!(parse_optional_filter(Some("bucket == 1"))
.expect("parse")
.is_some());
}
#[test]
fn topk_zero_limit_and_scored_candidate_ordering() {
let a = Doc::with_pk("a").expect("a");
let b = Doc::with_pk("b").expect("b");
let mut empty = TopKCollector::new(0);
empty.push(1.0, &a).expect("limit zero accepts");
assert!(empty.into_scored_docs().expect("ok").is_empty());
let left = ScoredCandidate::new(1.0, &a).expect("left");
let right = ScoredCandidate::new(1.0, &b).expect("right");
assert_eq!(left.cmp(&right), Ordering::Less);
assert_eq!(left.partial_cmp(&right), Some(Ordering::Less));
assert!(!left.eq(&right));
}
#[test]
fn score_conversion_binary_scoring_and_doc_sort_cover_edge_paths() {
use super::{binary_score, count_to_f64, score_to_f32, sort_docs};
use crate::doc::VectorValue;
assert_eq!(score_to_f32(1.5).expect("ok"), 1.5);
assert!(score_to_f32(f64::NAN).is_err());
assert!(score_to_f32(f64::INFINITY).is_err());
assert!(score_to_f32(f64::from(f32::MAX) * 2.0).is_err());
assert_eq!(count_to_f64(7), 7.0);
assert!(binary_score(&[0xff], &VectorValue::Fp32(vec![1.0])).is_none());
assert!(binary_score(&[0xff, 0x00], &VectorValue::Binary32(vec![0xff])).is_none());
assert_eq!(
binary_score(&[0xff, 0x00], &VectorValue::Binary32(vec![0x0f, 0xff])),
Some(-(4.0 + 8.0))
);
let mut docs = vec![Doc::with_pk("b").expect("b"), Doc::with_pk("a").expect("a")];
docs[0].set_score(1.0).expect("score");
docs[1].set_score(1.0).expect("score");
sort_docs(&mut docs);
assert_eq!(docs[0].get_pk(), Some("a"));
assert_eq!(docs[1].get_pk(), Some("b"));
}
}