Skip to main content

summa_core/query/
reranker.rs

1//! L2 reranker: rerank L1 candidates by exact dense vector distance on stored vectors
2//!
3//! Optimized for throughput:
4//! - Candidates grouped by segment for batched I/O
5//! - Flat indexes sorted for sequential mmap access (OS readahead)
6//! - Bounded SIMD batches (not per-candidate scoring or unbounded raw buffers)
7//! - Reusable per-segment scratch buffers
8//! - unit_norm fast path: skip per-vector norm when vectors are pre-normalized
9
10use futures::{StreamExt, TryStreamExt};
11use rustc_hash::{FxHashMap, FxHashSet};
12use std::sync::Arc;
13use std::sync::atomic::{AtomicUsize, Ordering as AtomicOrdering};
14
15use crate::dsl::Field;
16
17use super::{MultiValueCombiner, ScoredPosition, SearchResult, compare_search_results_desc};
18
19/// Maximum stored vectors expanded from the document candidate set by one L2
20/// rerank request. Candidate documents are bounded separately by the server;
21/// this closes the multi-value/ordinal multiplier. 2026-08-22: raised
22/// 500k → 2M (and bytes 512 MiB → 1 GiB): a few hundred book-length
23/// documents in one RAG rerank pool legitimately expand past 500k chunk
24/// vectors. Expansion is streamed in MAX_RERANK_RAW_BATCH_BYTES batches, so
25/// this bounds total scoring work, not resident memory.
26const MAX_L2_RERANK_VECTORS: usize = 2_000_000;
27const MAX_L2_RERANK_VECTOR_BYTES: usize = 1024 * 1024 * 1024;
28const RERANK_SCORE_BATCH: usize = 4_096;
29const MAX_RERANK_RAW_BATCH_BYTES: usize = 8 * 1024 * 1024;
30const MAX_CONCURRENT_RERANK_SEGMENTS: usize = 8;
31
32#[derive(Clone, Copy, PartialEq, Eq)]
33enum RerankerKind {
34    Dense,
35    Binary,
36}
37
38fn validate_reranker_config<D: crate::directories::Directory + 'static>(
39    searcher: &crate::index::Searcher<D>,
40    config: &RerankerConfig,
41) -> crate::error::Result<RerankerKind> {
42    if !config.rrf_k.is_finite() || config.rrf_k < 0.0 {
43        return Err(crate::Error::Query(format!(
44            "reranker rrf_k must be finite and non-negative, got {}",
45            config.rrf_k
46        )));
47    }
48    config.combiner.validate().map_err(crate::Error::Query)?;
49    if config.vector.is_empty() == config.binary_vector.is_empty() {
50        return Err(crate::Error::Query(
51            "reranker must provide exactly one of vector or binary_vector".to_string(),
52        ));
53    }
54
55    let entry = searcher
56        .schema()
57        .get_field_entry(config.field)
58        .ok_or_else(|| crate::Error::FieldNotFound(config.field.0.to_string()))?;
59    if !config.binary_vector.is_empty() {
60        if entry.field_type != crate::dsl::FieldType::BinaryDenseVector {
61            return Err(crate::Error::InvalidFieldType {
62                expected: "binary_dense_vector".to_string(),
63                got: format!("{:?}", entry.field_type),
64            });
65        }
66        let field_config = entry.binary_dense_vector_config.as_ref().ok_or_else(|| {
67            crate::Error::Schema(format!(
68                "binary dense vector field '{}' has no configuration",
69                entry.name
70            ))
71        })?;
72        if field_config.dim == 0 || !field_config.dim.is_multiple_of(8) {
73            return Err(crate::Error::Schema(format!(
74                "binary dense vector field '{}' has invalid dimension {}",
75                entry.name, field_config.dim
76            )));
77        }
78        if config.binary_vector.len() != field_config.byte_len() {
79            return Err(crate::Error::Query(format!(
80                "reranker binary vector byte length {} does not match field '{}' byte length {}",
81                config.binary_vector.len(),
82                entry.name,
83                field_config.byte_len()
84            )));
85        }
86        if config.matryoshka_dims.is_some() {
87            return Err(crate::Error::Query(
88                "reranker matryoshka_dims is not supported for binary vectors".to_string(),
89            ));
90        }
91        return Ok(RerankerKind::Binary);
92    }
93
94    if entry.field_type != crate::dsl::FieldType::DenseVector {
95        return Err(crate::Error::InvalidFieldType {
96            expected: "dense_vector".to_string(),
97            got: format!("{:?}", entry.field_type),
98        });
99    }
100    let field_config = entry.dense_vector_config.as_ref().ok_or_else(|| {
101        crate::Error::Schema(format!(
102            "dense vector field '{}' has no configuration",
103            entry.name
104        ))
105    })?;
106    if config.vector.len() != field_config.dim {
107        return Err(crate::Error::Query(format!(
108            "reranker vector dimension {} does not match field '{}' dimension {}",
109            config.vector.len(),
110            entry.name,
111            field_config.dim
112        )));
113    }
114    if let Some((index, value)) = config
115        .vector
116        .iter()
117        .enumerate()
118        .find(|(_, value)| !value.is_finite())
119    {
120        return Err(crate::Error::Query(format!(
121            "reranker vector contains non-finite value {value} at index {index}"
122        )));
123    }
124    if config.unit_norm != field_config.unit_norm {
125        return Err(crate::Error::Query(format!(
126            "reranker unit_norm={} does not match field '{}' unit_norm={}",
127            config.unit_norm, entry.name, field_config.unit_norm
128        )));
129    }
130    if let Some(dims) = config.matryoshka_dims
131        && (dims == 0 || dims > field_config.dim)
132    {
133        return Err(crate::Error::Query(format!(
134            "reranker matryoshka_dims must be in 1..={}, got {dims}",
135            field_config.dim
136        )));
137    }
138    Ok(RerankerKind::Dense)
139}
140
141fn reserve_rerank_vectors(
142    vector_budget: &AtomicUsize,
143    byte_budget: &AtomicUsize,
144    count: usize,
145    vector_byte_size: usize,
146) -> crate::error::Result<()> {
147    let bytes = count.checked_mul(vector_byte_size).ok_or_else(|| {
148        crate::Error::Query("reranker stored-vector byte budget overflow".to_string())
149    })?;
150    byte_budget
151        .fetch_update(AtomicOrdering::Relaxed, AtomicOrdering::Relaxed, |used| {
152            used.checked_add(bytes)
153                .filter(|&next| next <= MAX_L2_RERANK_VECTOR_BYTES)
154        })
155        .map_err(|used| {
156            crate::Error::Query(format!(
157                "reranker reads more than {MAX_L2_RERANK_VECTOR_BYTES} stored vector bytes \
158                 (already reserved {used}, next document needs {bytes})"
159            ))
160        })?;
161
162    vector_budget
163        .fetch_update(AtomicOrdering::Relaxed, AtomicOrdering::Relaxed, |used| {
164            used.checked_add(count)
165                .filter(|&next| next <= MAX_L2_RERANK_VECTORS)
166        })
167        .map(|_| ())
168        .map_err(|used| {
169            crate::Error::Query(format!(
170                "reranker expands to more than {MAX_L2_RERANK_VECTORS} stored vectors \
171                 (already reserved {used}, next document has {count})"
172            ))
173        })
174}
175
176/// Coalesce a chunk's flat indexes (ascending) into `(buffer_start,
177/// flat_start, count)` runs of adjacent vectors. Multi-valued documents store
178/// their values consecutively, so a chunk usually needs far fewer range
179/// reads than vectors. A repeated index (duplicate candidate) starts a new
180/// run rather than being rejected, so results are unchanged.
181fn plan_flat_read_runs(
182    flat_indexes: impl Iterator<Item = usize>,
183    runs: &mut Vec<(usize, usize, usize)>,
184) {
185    runs.clear();
186    for (buffer_index, flat_index) in flat_indexes.enumerate() {
187        if let Some(run) = runs.last_mut()
188            && run.1.checked_add(run.2) == Some(flat_index)
189        {
190            run.2 += 1;
191            continue;
192        }
193        runs.push((buffer_index, flat_index, 1));
194    }
195}
196
197/// Read the planned runs into `raw`, one range read per run.
198async fn read_flat_vector_runs(
199    lazy_flat: &crate::segment::LazyFlatVectorData,
200    runs: &[(usize, usize, usize)],
201    raw: &mut [u8],
202) -> crate::error::Result<()> {
203    let vbs = lazy_flat.vector_byte_size();
204    for &(buffer_start, flat_start, count) in runs {
205        let bytes = lazy_flat
206            .read_vectors_batch(flat_start, count)
207            .await
208            .map_err(crate::error::Error::Io)?;
209        let start = buffer_start
210            .checked_mul(vbs)
211            .ok_or_else(|| crate::Error::Query("rerank buffer offset overflow".into()))?;
212        let end = start
213            .checked_add(bytes.len())
214            .ok_or_else(|| crate::Error::Query("rerank buffer range overflow".into()))?;
215        raw.get_mut(start..end)
216            .ok_or_else(|| crate::Error::Corruption("rerank buffer is too short".into()))?
217            .copy_from_slice(bytes.as_slice());
218    }
219    Ok(())
220}
221
222/// Candidates that could not be reranked (segment gone or no stored vectors
223/// for the field) are dropped from the result; say so loudly.
224fn report_skipped_candidates<D: crate::directories::Directory + 'static>(
225    searcher: &crate::index::Searcher<D>,
226    kind: &'static str,
227    field_id: u32,
228    skipped: u32,
229    total: usize,
230) {
231    if skipped == 0 {
232        return;
233    }
234    let index_label = searcher.schema().index_label();
235    crate::observe::rerank_candidates_skipped(index_label, kind, u64::from(skipped));
236    log::warn!(
237        "[{kind}_vector_rerank] index={index_label} field {field_id}: {skipped} of {total} \
238         candidates skipped (segment missing or no stored vectors for the field) and dropped \
239         from the reranked result"
240    );
241}
242
243#[inline]
244fn rerank_batch_len(vector_byte_size: usize) -> usize {
245    RERANK_SCORE_BATCH.min((MAX_RERANK_RAW_BATCH_BYTES / vector_byte_size.max(1)).max(1))
246}
247
248/// Precomputed query data for dense reranking (computed once, reused across segments).
249struct PrecompQuery<'a> {
250    query: &'a [f32],
251    inv_norm_q: f32,
252    query_f16: &'a [u16],
253}
254
255/// Batch SIMD scoring with precomputed query norm + f16 query.
256#[inline]
257#[allow(clippy::too_many_arguments)]
258fn score_batch_precomp(
259    pq: &PrecompQuery<'_>,
260    raw: &[u8],
261    quant: crate::dsl::DenseVectorQuantization,
262    dim: usize,
263    scores: &mut [f32],
264    unit_norm: bool,
265) -> crate::error::Result<()> {
266    let query = pq.query;
267    let inv_norm_q = pq.inv_norm_q;
268    let query_f16 = pq.query_f16;
269    use crate::dsl::DenseVectorQuantization;
270    use crate::structures::simd;
271    let element_size = quant.element_size();
272    let required_bytes = scores
273        .len()
274        .checked_mul(dim)
275        .and_then(|elements| elements.checked_mul(element_size))
276        .ok_or_else(|| {
277            crate::Error::Corruption("dense reranker batch size overflow".to_string())
278        })?;
279    if raw.len() < required_bytes {
280        return Err(crate::Error::Corruption(format!(
281            "dense reranker batch is truncated: need {required_bytes} bytes, got {}",
282            raw.len()
283        )));
284    }
285    if matches!(
286        quant,
287        DenseVectorQuantization::F32 | DenseVectorQuantization::F16
288    ) && required_bytes > 0
289        && !(raw.as_ptr() as usize).is_multiple_of(element_size)
290    {
291        return Err(crate::Error::Corruption(format!(
292            "dense reranker {:?} data is not {}-byte aligned",
293            quant, element_size
294        )));
295    }
296    match (quant, unit_norm) {
297        (DenseVectorQuantization::F32, false) => {
298            let num_floats = scores.len() * dim;
299            // Safety: Vec<u8> from the global allocator is guaranteed to be at least
300            // 8-byte aligned on 64-bit platforms (aligned to max_align_t). Assert at
301            // runtime to guard against custom allocators with weaker guarantees.
302            let vectors: &[f32] =
303                unsafe { std::slice::from_raw_parts(raw.as_ptr() as *const f32, num_floats) };
304            simd::batch_cosine_scores_precomp(query, vectors, dim, scores, inv_norm_q);
305        }
306        (DenseVectorQuantization::F32, true) => {
307            let num_floats = scores.len() * dim;
308            let vectors: &[f32] =
309                unsafe { std::slice::from_raw_parts(raw.as_ptr() as *const f32, num_floats) };
310            simd::batch_dot_scores_precomp(query, vectors, dim, scores, inv_norm_q);
311        }
312        (DenseVectorQuantization::F16, false) => {
313            simd::batch_cosine_scores_f16_precomp(query_f16, raw, dim, scores, inv_norm_q);
314        }
315        (DenseVectorQuantization::F16, true) => {
316            simd::batch_dot_scores_f16_precomp(query_f16, raw, dim, scores, inv_norm_q);
317        }
318        (DenseVectorQuantization::UInt8, false) => {
319            simd::batch_cosine_scores_u8_precomp(query, raw, dim, scores, inv_norm_q);
320        }
321        (DenseVectorQuantization::UInt8, true) => {
322            simd::batch_dot_scores_u8_precomp(query, raw, dim, scores, inv_norm_q);
323        }
324        (DenseVectorQuantization::Binary, _) => {
325            return Err(crate::Error::InvalidFieldType {
326                expected: "non-binary dense vector".to_string(),
327                got: "binary dense vector".to_string(),
328            });
329        }
330    }
331    Ok(())
332}
333
334/// Configuration for L2 dense/binary vector reranking
335#[derive(Debug, Clone)]
336pub struct RerankerConfig {
337    /// Vector field (dense or binary dense)
338    pub field: Field,
339    /// Query vector (f32, for dense fields)
340    pub vector: Vec<f32>,
341    /// Query vector (packed bits, for binary dense fields).
342    /// When non-empty, Hamming distance scoring is used instead of cosine.
343    pub binary_vector: Vec<u8>,
344    /// How to combine scores for multi-valued documents
345    pub combiner: MultiValueCombiner,
346    /// Whether stored vectors are pre-normalized to unit L2 norm.
347    /// When true, scoring uses dot-product only (skips per-vector norm — ~40% faster).
348    /// Ignored for binary fields.
349    pub unit_norm: bool,
350    /// Matryoshka pre-filter: number of leading dimensions to use for cheap
351    /// approximate scoring before full-dimension exact reranking.
352    /// Ignored for binary fields.
353    pub matryoshka_dims: Option<usize>,
354    /// Reciprocal Rank Fusion k parameter. When > 0, fuses L1 (first-stage) and
355    /// L2 (reranker) rankings: `score(d) = 1/(k + rank_L1) + 1/(k + rank_L2)`.
356    /// Typical value: 60. When 0, RRF is disabled and only L2 scores are used.
357    pub rrf_k: f32,
358}
359
360/// Score a single document against the query vector (used by tests).
361#[cfg(test)]
362use crate::structures::simd::cosine_similarity;
363#[cfg(test)]
364fn score_document(
365    doc: &crate::dsl::Document,
366    config: &RerankerConfig,
367) -> Option<(f32, Vec<ScoredPosition>)> {
368    let query_dim = config.vector.len();
369    let mut values: Vec<(u32, f32)> = doc
370        .get_all(config.field)
371        .filter_map(|fv| fv.as_dense_vector())
372        .enumerate()
373        .filter_map(|(ordinal, vec)| {
374            if vec.len() != query_dim {
375                return None;
376            }
377            let score = cosine_similarity(&config.vector, vec);
378            Some((ordinal as u32, score))
379        })
380        .collect();
381
382    if values.is_empty() {
383        return None;
384    }
385
386    let combined = config.combiner.combine(&values);
387
388    // Sort ordinals by score descending (best chunk first)
389    values.sort_unstable_by(|a, b| b.1.total_cmp(&a.1));
390    let positions: Vec<ScoredPosition> = values
391        .into_iter()
392        .map(|(ordinal, score)| ScoredPosition::new(ordinal, score))
393        .collect();
394
395    Some((combined, positions))
396}
397
398/// Apply Reciprocal Rank Fusion to combine L1 and L2 rankings.
399///
400/// `candidates` is sorted by L1 score descending (first-stage query output).
401/// `scored` is sorted by L2 score descending (reranker output).
402/// Replaces each result's score with the RRF fused score, re-sorts, and truncates.
403///
404/// Formula (Cormack, Clarke, Buettcher 2009):
405///   `RRF(d) = 1/(k + rank_L1(d)) + 1/(k + rank_L2(d))`
406/// where ranks are 1-based.
407fn apply_rrf(
408    candidates: &[SearchResult],
409    scored: &mut Vec<SearchResult>,
410    k: f32,
411    final_limit: usize,
412) {
413    // Build L1 rank map: (segment_id, doc_id) → 1-based rank
414    let l1_ranks: FxHashMap<(u128, u32), usize> = candidates
415        .iter()
416        .enumerate()
417        .map(|(idx, c)| ((c.segment_id, c.doc_id), idx + 1))
418        .collect();
419
420    // scored is sorted by L2 score desc → enumerate index + 1 = L2 rank
421    for (l2_idx, result) in scored.iter_mut().enumerate() {
422        let l1_rank = l1_ranks
423            .get(&(result.segment_id, result.doc_id))
424            .copied()
425            .unwrap_or(candidates.len() + 1);
426        result.score = super::fusion::rrf_contribution(k, l1_rank)
427            + super::fusion::rrf_contribution(k, l2_idx + 1);
428    }
429
430    scored.sort_unstable_by(compare_search_results_desc);
431    scored.truncate(final_limit);
432}
433
434/// Rerank L1 candidates by exact dense vector distance.
435///
436/// Groups candidates by segment for batched I/O, sorts flat indexes for
437/// sequential mmap access, and scores vectors in bounded SIMD batches.
438/// Scratch memory remains independent of the candidate count.
439///
440/// When `unit_norm` is set in the config, scoring uses dot-product only
441/// (skips per-vector norm computation — ~40% less work).
442pub async fn rerank<D: crate::directories::Directory + 'static>(
443    searcher: &crate::index::Searcher<D>,
444    candidates: &[SearchResult],
445    config: &RerankerConfig,
446    final_limit: usize,
447) -> crate::error::Result<Vec<SearchResult>> {
448    // Validate before empty-result early returns so malformed requests do not
449    // succeed or fail depending on whether the first-stage query found hits.
450    let kind = validate_reranker_config(searcher, config)?;
451    if final_limit == 0 || candidates.is_empty() {
452        return Ok(Vec::new());
453    }
454
455    // Dispatch: binary vector → Hamming, f32 vector → cosine/dot.
456    if kind == RerankerKind::Binary {
457        return rerank_binary(searcher, candidates, config, final_limit).await;
458    }
459
460    let t0 = std::time::Instant::now();
461    let field_id = config.field.0;
462    let query = &config.vector;
463    let query_dim = query.len();
464    let segments = searcher.segment_readers();
465    let seg_by_id = searcher.segment_map();
466
467    // Precompute query inverse-norm and f16 query once (reused across all segments)
468    use crate::structures::simd;
469    let norm_q_sq = simd::dot_product_f32(query, query, query_dim);
470    let inv_norm_q = if norm_q_sq < f32::EPSILON {
471        0.0
472    } else {
473        simd::fast_inv_sqrt(norm_q_sq)
474    };
475    let query_f16: Vec<u16> = query.iter().map(|&v| simd::f32_to_f16(v)).collect();
476    let pq = PrecompQuery {
477        query,
478        inv_norm_q,
479        query_f16: &query_f16,
480    };
481
482    // ── Phase 1: Group candidates by segment ──────────────────────────────
483    let mut segment_groups: FxHashMap<usize, Vec<usize>> = FxHashMap::default();
484    let mut skipped = 0u32;
485
486    for (ci, candidate) in candidates.iter().enumerate() {
487        if let Some(&si) = seg_by_id.get(&candidate.segment_id) {
488            segment_groups.entry(si).or_default().push(ci);
489        } else {
490            skipped += 1;
491        }
492    }
493
494    // ── Phase 2: Per-segment batched resolve + read + score (concurrent) ──
495    // Bound the fan-out so many immutable segments cannot each retain a raw
496    // scoring buffer and candidate scratch at the same time.
497    let query_ref = pq.query;
498    let inv_norm_q_val = pq.inv_norm_q;
499    let query_f16_ref = pq.query_f16;
500    let vector_budget = Arc::new(AtomicUsize::new(0));
501    let byte_budget = Arc::new(AtomicUsize::new(0));
502
503    let segment_futs = futures::stream::iter(segment_groups.into_iter().map(
504        |(si, candidate_indices)| {
505            #[allow(clippy::redundant_locals)]
506            let segments = &segments;
507            #[allow(clippy::redundant_locals)]
508            let candidates = candidates;
509            #[allow(clippy::redundant_locals)]
510            let query_ref = query_ref;
511            #[allow(clippy::redundant_locals)]
512            let query_f16_ref = query_f16_ref;
513            #[allow(clippy::redundant_locals)]
514            let config = config;
515            let vector_budget = Arc::clone(&vector_budget);
516            let byte_budget = Arc::clone(&byte_budget);
517            async move {
518                let mut scores: Vec<(usize, u32, f32)> = Vec::new();
519                let mut vectors = 0usize;
520                let mut seg_skipped = 0u32;
521
522                let Some(lazy_flat) = segments[si].flat_vectors().get(&field_id) else {
523                    return Ok::<_, crate::error::Error>((
524                        scores,
525                        vectors,
526                        candidate_indices.len() as u32,
527                    ));
528                };
529                if lazy_flat.dim != query_dim {
530                    return Err(crate::Error::Corruption(format!(
531                        "dense reranker field {field_id} stores dimension {}, expected {query_dim}",
532                        lazy_flat.dim
533                    )));
534                }
535                if lazy_flat.quantization == crate::dsl::DenseVectorQuantization::Binary {
536                    return Err(crate::Error::Corruption(format!(
537                        "dense reranker field {field_id} unexpectedly uses binary storage"
538                    )));
539                }
540
541                let vbs = lazy_flat.vector_byte_size();
542                let quant = lazy_flat.quantization;
543
544                // Resolve flat indexes for all candidates in this segment
545                let mut resolved: Vec<(usize, usize, u32)> = Vec::new();
546                for &ci in &candidate_indices {
547                    let local_doc_id = candidates[ci].doc_id;
548                    let (start, count) = lazy_flat.flat_indexes_for_doc_range(local_doc_id);
549                    if count == 0 {
550                        seg_skipped += 1;
551                        continue;
552                    }
553                    reserve_rerank_vectors(&vector_budget, &byte_budget, count, vbs)?;
554                    for j in 0..count {
555                        let (_, ordinal) = lazy_flat.get_doc_id(start + j);
556                        resolved.push((ci, start + j, ordinal as u32));
557                    }
558                }
559
560                if resolved.is_empty() {
561                    return Ok((scores, vectors, seg_skipped));
562                }
563
564                let n = resolved.len();
565                vectors = n;
566
567                // Sort by flat_idx for sequential mmap access. Prefetch only
568                // each bounded scoring chunk below; advising the entire
569                // candidate set at once can flood the page cache.
570                resolved.sort_unstable_by_key(|&(_, flat_idx, _)| flat_idx);
571
572                let batch_len = rerank_batch_len(vbs);
573                let max_batch = batch_len.min(n);
574                let max_raw_len = max_batch.checked_mul(vbs).ok_or_else(|| {
575                    crate::Error::Query("dense reranker buffer size overflow".into())
576                })?;
577                let mut raw_buf = vec![0u8; max_raw_len];
578                let mut runs: Vec<(usize, usize, usize)> = Vec::new();
579
580                // Reconstruct PrecompQuery from captured components
581                let pq = PrecompQuery {
582                    query: query_ref,
583                    inv_norm_q: inv_norm_q_val,
584                    query_f16: query_f16_ref,
585                };
586
587                // Matryoshka pre-filter
588                if let Some(mdims) = config.matryoshka_dims
589                    && mdims < query_dim
590                    && n > crate::query::max_candidate_limit(final_limit)
591                {
592                    let trunc_dim = mdims;
593                    let trunc_pq = PrecompQuery {
594                        query: &query_ref[..trunc_dim],
595                        inv_norm_q: {
596                            let nq = simd::dot_product_f32(
597                                &query_ref[..trunc_dim],
598                                &query_ref[..trunc_dim],
599                                trunc_dim,
600                            );
601                            if nq < f32::EPSILON {
602                                0.0
603                            } else {
604                                simd::fast_inv_sqrt(nq)
605                            }
606                        },
607                        query_f16: &query_f16_ref[..trunc_dim],
608                    };
609                    let trunc_vbs = trunc_dim * quant.element_size();
610                    let mut scores_buf = vec![0.0f32; n];
611                    for (chunk_idx, chunk) in resolved.chunks(batch_len).enumerate() {
612                        // Pack only the Matryoshka prefix for this batch. The
613                        // previous full-vector reads pulled every unused tail
614                        // into the page cache and then reread surviving vectors.
615                        let raw_len = chunk.len().checked_mul(trunc_vbs).ok_or_else(|| {
616                            crate::Error::Query("dense reranker buffer size overflow".into())
617                        })?;
618                        let raw = &mut raw_buf[..raw_len];
619                        for (buf_idx, &(_, flat_idx, _)) in chunk.iter().enumerate() {
620                            lazy_flat
621                                .read_vector_prefix_raw_into(
622                                    flat_idx,
623                                    trunc_vbs,
624                                    &mut raw[buf_idx * trunc_vbs..(buf_idx + 1) * trunc_vbs],
625                                )
626                                .await
627                                .map_err(crate::error::Error::Io)?;
628                        }
629                        let score_base = chunk_idx * batch_len;
630                        searcher.install_search_cpu(|| {
631                            score_batch_precomp(
632                                &trunc_pq,
633                                raw,
634                                quant,
635                                trunc_dim,
636                                &mut scores_buf[score_base..score_base + chunk.len()],
637                                config.unit_norm,
638                            )
639                        })?;
640                    }
641
642                    // Rank approximate *documents*, using every stored value
643                    // and the configured combiner. Selecting individual
644                    // vectors here can discard the other values needed by
645                    // Avg/Sum/LSE and can choose the wrong Max document.
646                    let mut approximate_ordinals: FxHashMap<usize, Vec<(u32, f32)>> =
647                        FxHashMap::default();
648                    for (resolved_index, &(ci, _, ordinal)) in resolved.iter().enumerate() {
649                        approximate_ordinals
650                            .entry(ci)
651                            .or_default()
652                            .push((ordinal, scores_buf[resolved_index]));
653                    }
654                    let mut ranked: Vec<(usize, f32)> = approximate_ordinals
655                        .into_iter()
656                        .map(|(ci, ordinals)| (ci, config.combiner.combine(&ordinals)))
657                        .collect();
658                    searcher.install_search_cpu(|| {
659                        ranked.sort_unstable_by(|a, b| {
660                            b.1.total_cmp(&a.1)
661                                .then_with(|| candidates[a.0].doc_id.cmp(&candidates[b.0].doc_id))
662                        });
663                    });
664                    let approximate_docs = ranked.len();
665                    let survivor_doc_limit =
666                        crate::query::max_candidate_limit(final_limit).min(approximate_docs);
667                    let survivor_docs: FxHashSet<usize> = ranked
668                        .into_iter()
669                        .take(survivor_doc_limit)
670                        .map(|(ci, _)| ci)
671                        .collect();
672                    // Full-score every value belonging to each surviving doc;
673                    // the final combiner must never see a truncated value set.
674                    let mut survivor_entries: Vec<_> = resolved
675                        .iter()
676                        .copied()
677                        .filter(|(ci, _, _)| survivor_docs.contains(ci))
678                        .collect();
679                    survivor_entries.sort_unstable_by_key(|&(_, flat_idx, _)| flat_idx);
680                    let mut full_scores = vec![0.0f32; max_batch.min(survivor_entries.len())];
681                    scores.reserve(survivor_entries.len());
682                    for chunk in survivor_entries.chunks(batch_len) {
683                        #[cfg(feature = "native")]
684                        lazy_flat.prefetch_vectors(
685                            chunk.iter().map(|&(_, flat_idx, _)| flat_idx),
686                        );
687                        let raw_len = chunk.len().checked_mul(vbs).ok_or_else(|| {
688                            crate::Error::Query("dense reranker buffer size overflow".into())
689                        })?;
690                        let raw = &mut raw_buf[..raw_len];
691                        plan_flat_read_runs(
692                            chunk.iter().map(|&(_, flat_idx, _)| flat_idx),
693                            &mut runs,
694                        );
695                        read_flat_vector_runs(lazy_flat, &runs, raw).await?;
696                        searcher.install_search_cpu(|| {
697                            score_batch_precomp(
698                                &pq,
699                                raw,
700                                quant,
701                                query_dim,
702                                &mut full_scores[..chunk.len()],
703                                config.unit_norm,
704                            )
705                        })?;
706                        for (buf_idx, &(ci, _, ordinal)) in chunk.iter().enumerate() {
707                            scores.push((ci, ordinal, full_scores[buf_idx]));
708                        }
709                    }
710
711                    let survivor_vectors = survivor_entries.len();
712                    log::debug!(
713            "[dense_vector_rerank] matryoshka pre-filter: {}/{} dims, {}/{} docs and {}/{} vectors survived",
714                        trunc_dim,
715                        query_dim,
716                        survivor_docs.len(),
717                        approximate_docs,
718                        survivor_vectors,
719                        n,
720                    );
721                } else {
722                    let mut scores_buf = vec![0.0f32; max_batch];
723                    scores.reserve(n);
724                    for chunk in resolved.chunks(batch_len) {
725                        #[cfg(feature = "native")]
726                        lazy_flat.prefetch_vectors(
727                            chunk.iter().map(|&(_, flat_idx, _)| flat_idx),
728                        );
729                        let raw_len = chunk.len().checked_mul(vbs).ok_or_else(|| {
730                            crate::Error::Query("dense reranker buffer size overflow".into())
731                        })?;
732                        let raw = &mut raw_buf[..raw_len];
733                        plan_flat_read_runs(
734                            chunk.iter().map(|&(_, flat_idx, _)| flat_idx),
735                            &mut runs,
736                        );
737                        read_flat_vector_runs(lazy_flat, &runs, raw).await?;
738                        searcher.install_search_cpu(|| {
739                            score_batch_precomp(
740                                &pq,
741                                raw,
742                                quant,
743                                query_dim,
744                                &mut scores_buf[..chunk.len()],
745                                config.unit_norm,
746                            )
747                        })?;
748                        for (buf_idx, &(ci, _, ordinal)) in chunk.iter().enumerate() {
749                            scores.push((ci, ordinal, scores_buf[buf_idx]));
750                        }
751                    }
752                }
753
754                Ok((scores, vectors, seg_skipped))
755            }
756        },
757    ))
758    .buffer_unordered(MAX_CONCURRENT_RERANK_SEGMENTS);
759    futures::pin_mut!(segment_futs);
760
761    let mut all_scores: Vec<(usize, u32, f32)> = Vec::new();
762    let mut total_vectors = 0usize;
763    while let Some((scores, vectors, seg_skipped)) = segment_futs.try_next().await? {
764        all_scores.extend(scores);
765        total_vectors = total_vectors.saturating_add(vectors);
766        skipped = skipped.saturating_add(seg_skipped);
767    }
768
769    let read_score_elapsed = t0.elapsed();
770    report_skipped_candidates(searcher, "dense", field_id, skipped, candidates.len());
771
772    if total_vectors == 0 {
773        log::debug!(
774            "[dense_vector_rerank] field {}: {} candidates, all skipped (no flat vectors)",
775            field_id,
776            candidates.len()
777        );
778        return Ok(Vec::new());
779    }
780
781    // ── Phase 3: Combine scores and build results ─────────────────────────
782    // Sort flat buffer by candidate_idx so contiguous runs belong to the same doc
783    all_scores.sort_unstable_by_key(|&(ci, _, _)| ci);
784
785    let mut scored: Vec<SearchResult> = Vec::with_capacity(
786        candidates
787            .len()
788            .min(crate::query::max_candidate_limit(final_limit)),
789    );
790    let mut ordinal_pairs: Vec<(u32, f32)> = Vec::new();
791    let mut i = 0;
792    while i < all_scores.len() {
793        let ci = all_scores[i].0;
794        let run_start = i;
795        while i < all_scores.len() && all_scores[i].0 == ci {
796            i += 1;
797        }
798        let run = &mut all_scores[run_start..i];
799
800        // Build (ordinal, score) slice for combiner (reuses hoisted buffer)
801        ordinal_pairs.clear();
802        ordinal_pairs.extend(run.iter().map(|&(_, ord, s)| (ord, s)));
803        let combined = config.combiner.combine(&ordinal_pairs);
804
805        // Sort positions by score descending (best chunk first)
806        run.sort_unstable_by(|a, b| b.2.total_cmp(&a.2));
807        let positions: Vec<ScoredPosition> = run
808            .iter()
809            .map(|&(_, ord, score)| ScoredPosition::new(ord, score))
810            .collect();
811
812        scored.push(SearchResult {
813            doc_id: candidates[ci].doc_id,
814            score: combined,
815            segment_id: candidates[ci].segment_id,
816            positions: vec![(field_id, positions)],
817        });
818    }
819
820    scored.sort_unstable_by(compare_search_results_desc);
821
822    if config.rrf_k > 0.0 {
823        apply_rrf(candidates, &mut scored, config.rrf_k, final_limit);
824    } else {
825        scored.truncate(final_limit);
826    }
827
828    log::debug!(
829        "[dense_vector_rerank] field {}: {} candidates -> {} results (skipped {}, {} vectors, unit_norm={}, rrf_k={}): read+score={:.1}ms total={:.1}ms",
830        field_id,
831        candidates.len(),
832        scored.len(),
833        skipped,
834        total_vectors,
835        config.unit_norm,
836        config.rrf_k,
837        read_score_elapsed.as_secs_f64() * 1000.0,
838        t0.elapsed().as_secs_f64() * 1000.0,
839    );
840
841    Ok(scored)
842}
843
844/// Rerank L1 candidates by exact Hamming distance on stored binary vectors.
845async fn rerank_binary<D: crate::directories::Directory + 'static>(
846    searcher: &crate::index::Searcher<D>,
847    candidates: &[SearchResult],
848    config: &RerankerConfig,
849    final_limit: usize,
850) -> crate::error::Result<Vec<SearchResult>> {
851    if config.binary_vector.is_empty() || candidates.is_empty() {
852        return Ok(Vec::new());
853    }
854
855    let t0 = std::time::Instant::now();
856    let field_id = config.field.0;
857    let query = &config.binary_vector;
858    let byte_len = query.len();
859    let segments = searcher.segment_readers();
860    let seg_by_id = searcher.segment_map();
861
862    // Group candidates by segment
863    let mut segment_groups: FxHashMap<usize, Vec<usize>> = FxHashMap::default();
864    let mut skipped = 0u32;
865    for (ci, cand) in candidates.iter().enumerate() {
866        if let Some(&seg_idx) = seg_by_id.get(&cand.segment_id) {
867            let reader = &segments[seg_idx];
868            if reader.flat_vectors().contains_key(&field_id) {
869                segment_groups.entry(seg_idx).or_default().push(ci);
870                continue;
871            }
872        }
873        skipped += 1;
874    }
875
876    // Bounded concurrent per-segment scoring (same pattern as dense reranker).
877    let vector_budget = Arc::new(AtomicUsize::new(0));
878    let byte_budget = Arc::new(AtomicUsize::new(0));
879    let segment_futs = futures::stream::iter(segment_groups.into_iter().map(
880        |(seg_idx, cand_indices)| {
881            #[allow(clippy::redundant_locals)]
882            let segments = &segments;
883            #[allow(clippy::redundant_locals)]
884            let candidates = candidates;
885            let vector_budget = Arc::clone(&vector_budget);
886            let byte_budget = Arc::clone(&byte_budget);
887            async move {
888                let mut scores: Vec<(usize, u32, f32)> = Vec::new();
889                let mut seg_skipped = 0u32;
890
891                let Some(lazy_flat) = segments[seg_idx].flat_vectors().get(&field_id) else {
892                    return Ok::<_, crate::error::Error>((scores, cand_indices.len() as u32));
893                };
894                if lazy_flat.quantization != crate::dsl::DenseVectorQuantization::Binary
895                    || !lazy_flat.dim.is_multiple_of(8)
896                {
897                    return Err(crate::Error::Corruption(format!(
898                        "binary reranker field {field_id} has invalid flat-vector metadata"
899                    )));
900                }
901                let vbs = lazy_flat.vector_byte_size();
902                if vbs != byte_len {
903                    return Err(crate::Error::Corruption(format!(
904                        "binary reranker field {field_id} stores {vbs} bytes/vector, expected {byte_len}"
905                    )));
906                }
907
908                // Resolve flat indexes
909                let mut resolved: Vec<(usize, usize)> = Vec::new();
910                for &ci in &cand_indices {
911                    let doc_id = candidates[ci].doc_id;
912                    let (start, count) = lazy_flat.flat_indexes_for_doc_range(doc_id);
913                    if count == 0 {
914                        seg_skipped += 1;
915                        continue;
916                    }
917                    reserve_rerank_vectors(&vector_budget, &byte_budget, count, vbs)?;
918                    for j in 0..count {
919                        resolved.push((ci, start + j));
920                    }
921                }
922                if resolved.is_empty() {
923                    return Ok((scores, seg_skipped));
924                }
925
926                resolved.sort_unstable_by_key(|&(_, flat_idx)| flat_idx);
927
928                let n = resolved.len();
929                let batch_len = rerank_batch_len(vbs);
930                let max_batch = batch_len.min(n);
931                let max_raw_len = max_batch.checked_mul(vbs).ok_or_else(|| {
932                    crate::Error::Query("binary reranker buffer size overflow".into())
933                })?;
934                let mut raw_buf = vec![0u8; max_raw_len];
935                let mut runs: Vec<(usize, usize, usize)> = Vec::new();
936                let mut scores_buf = vec![0f32; max_batch];
937                scores.reserve(n);
938
939                for chunk in resolved.chunks(batch_len) {
940                    let raw_len = chunk.len().checked_mul(vbs).ok_or_else(|| {
941                        crate::Error::Query("binary reranker buffer size overflow".into())
942                    })?;
943                    let raw = &mut raw_buf[..raw_len];
944                    plan_flat_read_runs(chunk.iter().map(|&(_, flat_idx)| flat_idx), &mut runs);
945                    read_flat_vector_runs(lazy_flat, &runs, raw).await?;
946                    searcher.install_search_cpu(|| {
947                        crate::structures::simd::batch_hamming_scores(
948                            query,
949                            raw,
950                            byte_len,
951                            lazy_flat.dim,
952                            &mut scores_buf[..chunk.len()],
953                        );
954                    });
955
956                    for (buf_idx, &(ci, flat_idx)) in chunk.iter().enumerate() {
957                        let (_, ordinal) = lazy_flat.get_doc_id(flat_idx);
958                        scores.push((ci, ordinal as u32, scores_buf[buf_idx]));
959                    }
960                }
961
962                Ok((scores, seg_skipped))
963            }
964        },
965    ))
966    .buffer_unordered(MAX_CONCURRENT_RERANK_SEGMENTS);
967    futures::pin_mut!(segment_futs);
968
969    // Combine ordinal scores per candidate and apply combiner
970    let mut cand_ordinal_scores: FxHashMap<usize, Vec<(u32, f32)>> = FxHashMap::default();
971    while let Some((scores, seg_skipped)) = segment_futs.try_next().await? {
972        skipped = skipped.saturating_add(seg_skipped);
973        for (ci, ordinal, score) in scores {
974            cand_ordinal_scores
975                .entry(ci)
976                .or_default()
977                .push((ordinal, score));
978        }
979    }
980    report_skipped_candidates(searcher, "binary", field_id, skipped, candidates.len());
981
982    let total_vectors = cand_ordinal_scores.len();
983    let mut scored: Vec<SearchResult> = Vec::with_capacity(total_vectors);
984    for (ci, ordinal_scores) in cand_ordinal_scores {
985        let combined = config.combiner.combine(&ordinal_scores);
986        let positions: Vec<ScoredPosition> = ordinal_scores
987            .iter()
988            .map(|&(ord, s)| ScoredPosition::new(ord, s))
989            .collect();
990        scored.push(SearchResult {
991            doc_id: candidates[ci].doc_id,
992            score: combined,
993            segment_id: candidates[ci].segment_id,
994            positions: vec![(field_id, positions)],
995        });
996    }
997
998    scored.sort_unstable_by(compare_search_results_desc);
999
1000    if config.rrf_k > 0.0 {
1001        apply_rrf(candidates, &mut scored, config.rrf_k, final_limit);
1002    } else {
1003        scored.truncate(final_limit);
1004    }
1005
1006    log::debug!(
1007        "[dense_vector_binary_rerank] field {}: {} candidates -> {} results ({} docs scored, bytes_per_vector={}, rrf_k={}): {:.1}ms",
1008        field_id,
1009        candidates.len(),
1010        scored.len(),
1011        total_vectors,
1012        byte_len,
1013        config.rrf_k,
1014        t0.elapsed().as_secs_f64() * 1000.0,
1015    );
1016
1017    Ok(scored)
1018}
1019
1020/// Request-owned preparation for one immutable dense scoring component. The
1021/// F16 representation is materialized only if a segment needs that codec.
1022pub(super) struct CandidateVectorPreparation {
1023    inv_norm_q: f32,
1024    query_f16: Vec<u16>,
1025}
1026
1027/// Backfill sorted, unique flat vector targets with the existing rerank range
1028/// planner and SIMD kernels, without ANN nomination.
1029pub(super) async fn score_vector_candidates<D: crate::directories::Directory + 'static>(
1030    searcher: &crate::index::Searcher<D>,
1031    flat: &crate::segment::LazyFlatVectorData,
1032    vector: &[f32],
1033    binary_vector: &[u8],
1034    unit_norm: bool,
1035    targets: &[u32],
1036    preparation: &mut Option<CandidateVectorPreparation>,
1037) -> crate::Result<Vec<f32>> {
1038    use crate::structures::simd;
1039    let pq = if binary_vector.is_empty() {
1040        let preparation = preparation.get_or_insert_with(|| {
1041            let norm = simd::dot_product_f32(vector, vector, vector.len());
1042            CandidateVectorPreparation {
1043                inv_norm_q: if norm < f32::EPSILON {
1044                    0.0
1045                } else {
1046                    simd::fast_inv_sqrt(norm)
1047                },
1048                query_f16: Vec::new(),
1049            }
1050        });
1051        if flat.quantization == crate::dsl::DenseVectorQuantization::F16
1052            && preparation.query_f16.is_empty()
1053        {
1054            preparation
1055                .query_f16
1056                .extend(vector.iter().map(|&v| simd::f32_to_f16(v)));
1057        }
1058        Some(PrecompQuery {
1059            query: vector,
1060            inv_norm_q: preparation.inv_norm_q,
1061            query_f16: &preparation.query_f16,
1062        })
1063    } else {
1064        None
1065    };
1066    let vbs = flat.vector_byte_size();
1067    let batch_len = rerank_batch_len(vbs);
1068    let mut raw = vec![0u8; batch_len.min(targets.len()) * vbs];
1069    let mut runs = Vec::new();
1070    let mut scores = vec![0.0; targets.len()];
1071    for (batch, out) in targets.chunks(batch_len).zip(scores.chunks_mut(batch_len)) {
1072        plan_flat_read_runs(batch.iter().map(|&target| target as usize), &mut runs);
1073        let raw = &mut raw[..batch.len() * vbs];
1074        read_flat_vector_runs(flat, &runs, raw).await?;
1075        if !binary_vector.is_empty() {
1076            searcher.install_search_cpu(|| {
1077                simd::batch_hamming_scores(binary_vector, raw, binary_vector.len(), flat.dim, out)
1078            });
1079        } else {
1080            searcher.install_search_cpu(|| {
1081                score_batch_precomp(
1082                    pq.as_ref().expect("dense query"),
1083                    raw,
1084                    flat.quantization,
1085                    flat.dim,
1086                    out,
1087                    unit_norm,
1088                )
1089            })?;
1090        }
1091    }
1092    Ok(scores)
1093}
1094
1095#[cfg(test)]
1096mod tests {
1097    use super::*;
1098    use crate::dsl::{Document, Field};
1099
1100    /// Adjacent flat indexes (multi-valued documents, consecutive candidates)
1101    /// coalesce into one range read; gaps and repeated indexes start new
1102    /// runs so every buffer slot is still filled exactly once.
1103    #[test]
1104    fn flat_read_runs_coalesce_adjacent_indexes_and_tolerate_duplicates() {
1105        let mut runs = Vec::new();
1106        plan_flat_read_runs([3usize, 4, 5, 9, 10, 20, 20, 21].into_iter(), &mut runs);
1107        assert_eq!(runs, vec![(0, 3, 3), (3, 9, 2), (5, 20, 1), (6, 20, 2)]);
1108
1109        plan_flat_read_runs(std::iter::empty(), &mut runs);
1110        assert!(runs.is_empty());
1111
1112        plan_flat_read_runs([usize::MAX].into_iter(), &mut runs);
1113        assert_eq!(runs, vec![(0, usize::MAX, 1)]);
1114    }
1115
1116    fn make_config(vector: Vec<f32>, combiner: MultiValueCombiner) -> RerankerConfig {
1117        RerankerConfig {
1118            field: Field(0),
1119            vector,
1120            binary_vector: Vec::new(),
1121            combiner,
1122            unit_norm: false,
1123            matryoshka_dims: None,
1124            rrf_k: 0.0,
1125        }
1126    }
1127
1128    #[test]
1129    fn rerank_batches_are_bounded_by_bytes() {
1130        assert_eq!(rerank_batch_len(1), RERANK_SCORE_BATCH);
1131        assert_eq!(
1132            rerank_batch_len(MAX_RERANK_RAW_BATCH_BYTES),
1133            1,
1134            "one very wide vector must still make progress"
1135        );
1136        assert!(
1137            rerank_batch_len(4_096) * 4_096 <= MAX_RERANK_RAW_BATCH_BYTES,
1138            "normal batches must stay within the raw scratch budget"
1139        );
1140    }
1141
1142    #[test]
1143    fn rerank_budget_bounds_count_and_bytes() {
1144        let vectors = AtomicUsize::new(0);
1145        let bytes = AtomicUsize::new(0);
1146        reserve_rerank_vectors(&vectors, &bytes, 2, 32).unwrap();
1147        assert_eq!(vectors.load(AtomicOrdering::Relaxed), 2);
1148        assert_eq!(bytes.load(AtomicOrdering::Relaxed), 64);
1149
1150        let vectors = AtomicUsize::new(0);
1151        let bytes = AtomicUsize::new(0);
1152        assert!(reserve_rerank_vectors(&vectors, &bytes, 2, MAX_L2_RERANK_VECTOR_BYTES).is_err());
1153    }
1154
1155    #[test]
1156    fn test_score_document_single_value() {
1157        let mut doc = Document::new();
1158        doc.add_dense_vector(Field(0), vec![1.0, 0.0, 0.0]);
1159
1160        let config = make_config(vec![1.0, 0.0, 0.0], MultiValueCombiner::Max);
1161        let (score, positions) = score_document(&doc, &config).unwrap();
1162        // cosine([1,0,0], [1,0,0]) = 1.0
1163        assert!((score - 1.0).abs() < 1e-6);
1164        assert_eq!(positions.len(), 1);
1165        assert_eq!(positions[0].position, 0); // ordinal 0
1166    }
1167
1168    #[test]
1169    fn test_score_document_orthogonal() {
1170        let mut doc = Document::new();
1171        doc.add_dense_vector(Field(0), vec![0.0, 1.0, 0.0]);
1172
1173        let config = make_config(vec![1.0, 0.0, 0.0], MultiValueCombiner::Max);
1174        let (score, _) = score_document(&doc, &config).unwrap();
1175        // cosine([1,0,0], [0,1,0]) = 0.0
1176        assert!(score.abs() < 1e-6);
1177    }
1178
1179    #[test]
1180    fn test_score_document_multi_value_max() {
1181        let mut doc = Document::new();
1182        doc.add_dense_vector(Field(0), vec![1.0, 0.0, 0.0]); // cos=1.0 (same direction)
1183        doc.add_dense_vector(Field(0), vec![0.0, 1.0, 0.0]); // cos=0.0 (orthogonal)
1184
1185        let config = make_config(vec![1.0, 0.0, 0.0], MultiValueCombiner::Max);
1186        let (score, positions) = score_document(&doc, &config).unwrap();
1187        assert!((score - 1.0).abs() < 1e-6);
1188        // Best chunk first
1189        assert_eq!(positions.len(), 2);
1190        assert_eq!(positions[0].position, 0); // ordinal 0 scored highest
1191        assert!((positions[0].score - 1.0).abs() < 1e-6);
1192    }
1193
1194    #[test]
1195    fn test_score_document_multi_value_avg() {
1196        let mut doc = Document::new();
1197        doc.add_dense_vector(Field(0), vec![1.0, 0.0, 0.0]); // cos=1.0
1198        doc.add_dense_vector(Field(0), vec![0.0, 1.0, 0.0]); // cos=0.0
1199
1200        let config = make_config(vec![1.0, 0.0, 0.0], MultiValueCombiner::Avg);
1201        let (score, _) = score_document(&doc, &config).unwrap();
1202        // avg(1.0, 0.0) = 0.5
1203        assert!((score - 0.5).abs() < 1e-6);
1204    }
1205
1206    #[test]
1207    fn test_score_document_missing_field() {
1208        let mut doc = Document::new();
1209        // Add to field 1, not field 0
1210        doc.add_dense_vector(Field(1), vec![1.0, 0.0, 0.0]);
1211
1212        let config = make_config(vec![1.0, 0.0, 0.0], MultiValueCombiner::Max);
1213        assert!(score_document(&doc, &config).is_none());
1214    }
1215
1216    #[test]
1217    fn test_score_document_wrong_field_type() {
1218        let mut doc = Document::new();
1219        doc.add_text(Field(0), "not a vector");
1220
1221        let config = make_config(vec![1.0, 0.0, 0.0], MultiValueCombiner::Max);
1222        assert!(score_document(&doc, &config).is_none());
1223    }
1224
1225    #[test]
1226    fn test_score_document_dimension_mismatch() {
1227        let mut doc = Document::new();
1228        doc.add_dense_vector(Field(0), vec![1.0, 0.0]); // 2D
1229
1230        let config = make_config(vec![1.0, 0.0, 0.0], MultiValueCombiner::Max); // 3D query
1231        assert!(score_document(&doc, &config).is_none());
1232    }
1233
1234    #[test]
1235    fn test_score_document_empty_query_vector() {
1236        let mut doc = Document::new();
1237        doc.add_dense_vector(Field(0), vec![1.0, 0.0, 0.0]);
1238
1239        let config = make_config(vec![], MultiValueCombiner::Max);
1240        // Empty query can't match any stored vector (dimension mismatch)
1241        assert!(score_document(&doc, &config).is_none());
1242    }
1243
1244    fn make_result(doc_id: u32, score: f32, segment_id: u128) -> SearchResult {
1245        SearchResult {
1246            doc_id,
1247            score,
1248            segment_id,
1249            positions: Vec::new(),
1250        }
1251    }
1252
1253    #[test]
1254    fn test_rrf_basic_fusion() {
1255        // L1 ranking: doc A(rank 1), B(rank 2), C(rank 3)
1256        let candidates = vec![
1257            make_result(1, 10.0, 1), // A: L1 rank 1
1258            make_result(2, 8.0, 1),  // B: L1 rank 2
1259            make_result(3, 5.0, 1),  // C: L1 rank 3
1260        ];
1261
1262        // L2 ranking (reversed): C(rank 1), B(rank 2), A(rank 3)
1263        let mut scored = vec![
1264            make_result(3, 0.9, 1), // C: L2 rank 1
1265            make_result(2, 0.7, 1), // B: L2 rank 2
1266            make_result(1, 0.3, 1), // A: L2 rank 3
1267        ];
1268
1269        let k = 60.0;
1270        apply_rrf(&candidates, &mut scored, k, 10);
1271
1272        // B should win: rank 2 in both → 2/(k+2) vs split ranks for A and C
1273        // A: 1/(61) + 1/(63) = 0.01639 + 0.01587 = 0.03226
1274        // B: 1/(62) + 1/(62) = 0.01613 + 0.01613 = 0.03226
1275        // C: 1/(63) + 1/(61) = 0.01587 + 0.01639 = 0.03226
1276        // All equal! (symmetric: rank sum = 4 for each)
1277        // Actually: A: 1/61 + 1/63, B: 1/62 + 1/62, C: 1/63 + 1/61
1278        // A = C by symmetry, B is slightly different
1279        // 1/61 + 1/63 = (63+61)/(61*63) = 124/3843 = 0.032267
1280        // 1/62 + 1/62 = 2/62 = 1/31 = 0.032258
1281        // So A = C > B (very slightly). Top result should be doc 3 (C) or doc 1 (A)
1282        // since they have the same RRF score but C appeared first in scored.
1283
1284        assert_eq!(scored.len(), 3);
1285        // All three should have very similar RRF scores
1286        let spread = scored[0].score - scored[2].score;
1287        assert!(
1288            spread < 0.001,
1289            "All docs have near-equal RRF scores, spread={spread}"
1290        );
1291    }
1292
1293    #[test]
1294    fn test_rrf_clear_winner() {
1295        // Doc X is rank 1 in both L1 and L2 → should clearly win
1296        let candidates = vec![
1297            make_result(1, 10.0, 1), // X: L1 rank 1
1298            make_result(2, 8.0, 1),  // Y: L1 rank 2
1299            make_result(3, 5.0, 1),  // Z: L1 rank 3
1300        ];
1301
1302        // L2 ranking: X still rank 1
1303        let mut scored = vec![
1304            make_result(1, 0.95, 1), // X: L2 rank 1
1305            make_result(3, 0.50, 1), // Z: L2 rank 2
1306            make_result(2, 0.30, 1), // Y: L2 rank 3
1307        ];
1308
1309        let k = 60.0;
1310        apply_rrf(&candidates, &mut scored, k, 10);
1311
1312        // X: 1/(61) + 1/(61) = 2/61 = 0.03279 (best)
1313        // Y: 1/(62) + 1/(63) = 0.03200 (worst)
1314        // Z: 1/(63) + 1/(62) = 0.03200 (same as Y by symmetry)
1315        assert_eq!(scored[0].doc_id, 1, "Doc 1 (rank 1 in both) should be top");
1316        assert!(scored[0].score > scored[1].score);
1317    }
1318
1319    #[test]
1320    fn test_rrf_truncation() {
1321        let candidates = vec![
1322            make_result(1, 10.0, 1),
1323            make_result(2, 8.0, 1),
1324            make_result(3, 5.0, 1),
1325            make_result(4, 3.0, 1),
1326            make_result(5, 1.0, 1),
1327        ];
1328
1329        let mut scored = vec![
1330            make_result(5, 0.9, 1),
1331            make_result(4, 0.8, 1),
1332            make_result(3, 0.7, 1),
1333            make_result(2, 0.6, 1),
1334            make_result(1, 0.5, 1),
1335        ];
1336
1337        apply_rrf(&candidates, &mut scored, 60.0, 3);
1338        assert_eq!(scored.len(), 3, "Should truncate to final_limit=3");
1339    }
1340
1341    #[test]
1342    fn test_rrf_missing_l1_candidate() {
1343        // L1 has docs 1, 2. L2 scored doc 3 which wasn't in L1 candidates.
1344        let candidates = vec![make_result(1, 10.0, 1), make_result(2, 8.0, 1)];
1345
1346        let mut scored = vec![
1347            make_result(3, 0.9, 1), // not in L1 → gets worst L1 rank
1348            make_result(1, 0.5, 1),
1349        ];
1350
1351        apply_rrf(&candidates, &mut scored, 60.0, 10);
1352
1353        // Doc 1: L1 rank 1 → 1/61, L2 rank 2 → 1/62  = 0.03252
1354        // Doc 3: L1 rank 3 (fallback) → 1/63, L2 rank 1 → 1/61 = 0.03226
1355        // Doc 1 should win because it has a better L1 rank
1356        assert_eq!(scored[0].doc_id, 1);
1357    }
1358
1359    #[test]
1360    fn test_rrf_small_k() {
1361        // With small k, rank differences matter more
1362        let candidates = vec![make_result(1, 10.0, 1), make_result(2, 8.0, 1)];
1363
1364        let mut scored = vec![
1365            make_result(2, 0.9, 1), // L2 rank 1
1366            make_result(1, 0.5, 1), // L2 rank 2
1367        ];
1368
1369        apply_rrf(&candidates, &mut scored, 1.0, 10);
1370
1371        // k=1: Doc 1: 1/(1+1) + 1/(1+2) = 0.5 + 0.333 = 0.833
1372        //       Doc 2: 1/(1+2) + 1/(1+1) = 0.333 + 0.5 = 0.833
1373        // With k=1 and symmetric ranks, scores are equal
1374        let diff = (scored[0].score - scored[1].score).abs();
1375        assert!(
1376            diff < 1e-6,
1377            "Symmetric ranks should produce equal RRF scores"
1378        );
1379    }
1380}