Skip to main content

multivector/
retrieval.rs

1//! Durable document representations and one retrieval plan over a generation.
2use super::*;
3use annex::vector::sparse::{SparseIndex, SparseVector};
4use std::collections::{BTreeMap, BTreeSet};
5use std::time::Instant;
6
7#[derive(Clone, Debug, Deserialize, Serialize)]
8#[serde(tag = "kind", rename_all = "snake_case", deny_unknown_fields)]
9pub enum Representation {
10    Dense { vector: Vector },
11    Multivector { vectors: Vec<Vector> },
12    Sparse { vector: SparseVector },
13}
14
15#[derive(Clone, Debug, Default, Deserialize, Serialize)]
16#[serde(deny_unknown_fields)]
17pub struct Chunk {
18    pub parent: String,
19    pub position: u32,
20}
21
22#[derive(Clone, Debug, Default, Deserialize, Serialize)]
23#[serde(deny_unknown_fields)]
24pub struct RetrievalDocument {
25    pub id: String,
26    #[serde(default)]
27    pub vectors: Vec<Vector>,
28    #[serde(default)]
29    pub metadata: Value,
30    pub text: Option<String>,
31    #[serde(default)]
32    pub representations: BTreeMap<String, Representation>,
33    pub chunk: Option<Chunk>,
34}
35
36#[derive(Clone, Debug, Deserialize, Serialize)]
37pub(super) enum StoredRepresentation {
38    Dense {
39        location: ObjectLocation,
40        dimension: usize,
41    },
42    Multivector {
43        location: ObjectLocation,
44        dimension: usize,
45        tokens: usize,
46    },
47    Sparse(SparseVector),
48}
49#[derive(Clone, Debug, Default, Deserialize, Serialize)]
50pub(super) struct Fields {
51    text: Option<String>,
52    representations: BTreeMap<String, StoredRepresentation>,
53    chunk: Option<Chunk>,
54}
55
56#[derive(Clone, Copy, Debug, Deserialize, Eq, PartialEq, Serialize)]
57#[serde(tag = "kind", rename_all = "snake_case", deny_unknown_fields)]
58pub(super) enum FieldSchema {
59    Dense { dimension: usize },
60    Multivector { dimension: usize },
61    Sparse,
62}
63
64#[derive(Clone, Default)]
65pub(super) struct RetrievalState {
66    next_id: u64,
67    by_id: HashMap<String, u64>,
68    ids: HashMap<u64, String>,
69    vocabulary: HashMap<String, u32>,
70    lexical: SparseIndex,
71    sparse: HashMap<String, SparseIndex>,
72    schema: BTreeMap<String, FieldSchema>,
73    document_chunks: HashMap<u64, Chunk>,
74    chunks: HashMap<String, BTreeMap<u32, BTreeSet<String>>>,
75    analyzer: Analyzer,
76}
77
78fn invalid(message: impl Into<String>) -> IndexError {
79    IndexError::Invalid(message.into())
80}
81fn sparse_error(error: annex::vector::sparse::SparseError) -> IndexError {
82    invalid(error.to_string())
83}
84
85impl RetrievalState {
86    pub(super) fn from_schema(
87        schema: BTreeMap<String, FieldSchema>,
88        analyzer: Analyzer,
89    ) -> Result<Self, IndexError> {
90        if schema.iter().any(|(name, shape)| {
91            name.is_empty()
92                || name.len() > 128
93                || matches!(
94                    shape,
95                    FieldSchema::Dense { dimension }
96                        | FieldSchema::Multivector { dimension }
97                        if !(1..=65_536).contains(dimension)
98                )
99        }) {
100            return Err(invalid("invalid persisted representation schema"));
101        }
102        Ok(Self {
103            schema,
104            analyzer,
105            ..Self::default()
106        })
107    }
108
109    pub(super) fn schema(&self) -> &BTreeMap<String, FieldSchema> {
110        &self.schema
111    }
112
113    pub(super) fn dense_dimension(&self, field: &str) -> Result<usize, IndexError> {
114        match self.schema.get(field) {
115            Some(FieldSchema::Dense { dimension }) => Ok(*dimension),
116            _ => Err(invalid("unknown dense field")),
117        }
118    }
119    pub(super) fn remove(&mut self, id: &str) {
120        if let Some(number) = self.by_id.remove(id) {
121            if let Some(chunk) = self.document_chunks.remove(&number) {
122                if let Some(positions) = self.chunks.get_mut(&chunk.parent) {
123                    if let Some(ids) = positions.get_mut(&chunk.position) {
124                        ids.remove(id);
125                        if ids.is_empty() {
126                            positions.remove(&chunk.position);
127                        }
128                    }
129                    if positions.is_empty() {
130                        self.chunks.remove(&chunk.parent);
131                    }
132                }
133            }
134            self.ids.remove(&number);
135            self.lexical.delete(number);
136            for index in self.sparse.values() {
137                index.delete(number);
138            }
139        }
140    }
141    pub(super) fn insert(&mut self, id: &str, fields: &Fields) -> Result<(), IndexError> {
142        self.remove(id);
143        let number = self.next_id;
144        self.next_id = number
145            .checked_add(1)
146            .ok_or_else(|| invalid("document IDs exhausted"))?;
147        self.by_id.insert(id.to_owned(), number);
148        self.ids.insert(number, id.to_owned());
149        if let Some(chunk) = &fields.chunk {
150            self.chunks
151                .entry(chunk.parent.clone())
152                .or_default()
153                .entry(chunk.position)
154                .or_default()
155                .insert(id.to_owned());
156            self.document_chunks.insert(number, chunk.clone());
157        }
158        if let Some(text) = &fields.text {
159            let mut pairs = Vec::new();
160            for (term, count) in self.analyzer.analyze(text) {
161                let next = u32::try_from(self.vocabulary.len())
162                    .map_err(|_| invalid("lexical vocabulary exhausted"))?;
163                pairs.push((*self.vocabulary.entry(term).or_insert(next), count));
164            }
165            self.lexical
166                .upsert(number, &SparseVector::from_pairs(pairs))
167                .map_err(sparse_error)?;
168        }
169        for (name, representation) in &fields.representations {
170            let shape = match representation {
171                StoredRepresentation::Dense { dimension, .. } => FieldSchema::Dense {
172                    dimension: *dimension,
173                },
174                StoredRepresentation::Multivector { dimension, .. } => FieldSchema::Multivector {
175                    dimension: *dimension,
176                },
177                StoredRepresentation::Sparse(vector) => {
178                    self.sparse
179                        .entry(name.clone())
180                        .or_default()
181                        .upsert(number, vector)
182                        .map_err(sparse_error)?;
183                    FieldSchema::Sparse
184                }
185            };
186            if self
187                .schema
188                .get(name)
189                .is_some_and(|expected| *expected != shape)
190            {
191                return Err(invalid(format!(
192                    "representation {name:?} has a different kind or dimension"
193                )));
194            }
195            self.schema.insert(name.clone(), shape);
196        }
197        Ok(())
198    }
199
200    fn neighbors<'a>(&'a self, chunk: &Chunk, radius: u32) -> impl Iterator<Item = &'a str> {
201        let start = chunk.position.saturating_sub(radius);
202        let end = chunk.position.saturating_add(radius);
203        self.chunks
204            .get(&chunk.parent)
205            .into_iter()
206            .flat_map(move |positions| positions.range(start..=end))
207            .flat_map(|(_, ids)| ids.iter().map(String::as_str))
208    }
209}
210
211#[derive(Clone, Debug, Deserialize, Serialize)]
212#[serde(tag = "op", rename_all = "snake_case", deny_unknown_fields)]
213pub enum Predicate {
214    Eq {
215        field: String,
216        value: Value,
217    },
218    In {
219        field: String,
220        values: Vec<Value>,
221    },
222    Range {
223        field: String,
224        gte: Option<f64>,
225        lte: Option<f64>,
226    },
227    And {
228        filters: Vec<Predicate>,
229    },
230    Or {
231        filters: Vec<Predicate>,
232    },
233    Not {
234        filter: Box<Predicate>,
235    },
236}
237impl Predicate {
238    fn validate(&self, depth: usize) -> Result<(), IndexError> {
239        if depth > 16 {
240            return Err(invalid("filter nesting exceeds 16"));
241        }
242        match self {
243            Self::Eq { field, value } => {
244                if field.is_empty() || value.is_object() || value.is_array() {
245                    return Err(invalid("eq requires a field and scalar value"));
246                }
247            }
248            Self::In { field, values } => {
249                if field.is_empty()
250                    || values.is_empty()
251                    || values.len() > 1024
252                    || values.iter().any(|v| v.is_object() || v.is_array())
253                {
254                    return Err(invalid("in requires 1..=1024 scalar values"));
255                }
256            }
257            Self::Range { field, gte, lte } => {
258                if field.is_empty()
259                    || (gte.is_none() && lte.is_none())
260                    || gte.iter().chain(lte).any(|v| !v.is_finite())
261                    || matches!((gte,lte), (Some(a),Some(b)) if a>b)
262                {
263                    return Err(invalid("invalid numeric range"));
264                }
265            }
266            Self::And { filters } | Self::Or { filters } => {
267                if filters.is_empty() || filters.len() > 64 {
268                    return Err(invalid("boolean filters require 1..=64 children"));
269                }
270                for f in filters {
271                    f.validate(depth + 1)?;
272                }
273            }
274            Self::Not { filter } => filter.validate(depth + 1)?,
275        }
276        Ok(())
277    }
278    fn matches(&self, metadata: &Value) -> bool {
279        let value = |field: &str| {
280            if field.starts_with('/') {
281                metadata.pointer(field)
282            } else {
283                metadata.get(field)
284            }
285        };
286        match self {
287            Self::Eq {
288                field,
289                value: expected,
290            } => value(field).is_some_and(|v| v == expected),
291            Self::In { field, values } => value(field).is_some_and(|v| match v {
292                Value::Array(a) => a.iter().any(|x| values.contains(x)),
293                _ => values.contains(v),
294            }),
295            Self::Range { field, gte, lte } => value(field)
296                .and_then(Value::as_f64)
297                .is_some_and(|v| gte.is_none_or(|lo| v >= lo) && lte.is_none_or(|hi| v <= hi)),
298            Self::And { filters } => filters.iter().all(|f| f.matches(metadata)),
299            Self::Or { filters } => filters.iter().any(|f| f.matches(metadata)),
300            Self::Not { filter } => !filter.matches(metadata),
301        }
302    }
303}
304fn default_limit() -> usize {
305    100
306}
307fn default_k1() -> f32 {
308    1.2
309}
310fn default_b() -> f32 {
311    0.75
312}
313fn default_ef() -> usize {
314    256
315}
316#[derive(Clone, Debug, Deserialize, Serialize)]
317#[serde(tag = "kind", rename_all = "snake_case", deny_unknown_fields)]
318pub enum Channel {
319    Bm25 {
320        text: String,
321        #[serde(default = "default_limit")]
322        limit: usize,
323        #[serde(default = "default_k1")]
324        k1: f32,
325        #[serde(default = "default_b")]
326        b: f32,
327    },
328    Sparse {
329        field: String,
330        vector: SparseVector,
331        #[serde(default = "default_limit")]
332        limit: usize,
333    },
334    Dense {
335        field: String,
336        vector: Vector,
337        #[serde(default = "default_limit")]
338        limit: usize,
339        #[serde(default = "default_backend")]
340        backend: String,
341        #[serde(default = "default_ef")]
342        ef_search: usize,
343    },
344    Multivector {
345        field: Option<String>,
346        vectors: Vec<Vector>,
347        #[serde(default = "default_limit")]
348        limit: usize,
349        #[serde(default = "default_backend")]
350        backend: String,
351        #[serde(default = "default_ef")]
352        ef_search: usize,
353    },
354}
355fn default_backend() -> String {
356    "auto".into()
357}
358#[derive(Clone, Debug, Deserialize, Serialize)]
359#[serde(tag = "kind", rename_all = "snake_case", deny_unknown_fields)]
360pub enum Fusion {
361    Rrf {
362        #[serde(default = "default_rrf")]
363        k: f32,
364    },
365    Weighted {
366        weights: Vec<f32>,
367    },
368}
369fn default_rrf() -> f32 {
370    10.
371}
372impl Default for Fusion {
373    fn default() -> Self {
374        Self::Rrf { k: 10. }
375    }
376}
377#[derive(Clone, Debug, Deserialize, Serialize)]
378#[serde(deny_unknown_fields)]
379pub struct Rerank {
380    pub vectors: Vec<Vector>,
381    pub field: Option<String>,
382    #[serde(default = "default_limit")]
383    pub limit: usize,
384    pub adaptive: Option<AdaptiveRerank>,
385}
386#[derive(Clone, Debug, Deserialize, Serialize)]
387#[serde(deny_unknown_fields)]
388pub struct AdaptiveRerank {
389    pub min_candidates: usize,
390    pub agreement_threshold: f32,
391}
392#[derive(Clone, Debug, Default, Deserialize, Serialize)]
393#[serde(deny_unknown_fields)]
394pub struct ContextOptions {
395    pub per_parent: Option<usize>,
396    #[serde(default)]
397    pub neighbors: u32,
398    #[serde(default)]
399    pub deduplicate: bool,
400    pub mmr: Option<f32>,
401    pub diversity_field: Option<String>,
402}
403#[derive(Clone, Debug, Deserialize, Serialize)]
404#[serde(deny_unknown_fields)]
405pub struct RetrieveRequest {
406    pub prefetch: Vec<Channel>,
407    #[serde(default)]
408    pub fusion: Fusion,
409    pub filter: Option<Predicate>,
410    pub rerank: Option<Rerank>,
411    #[serde(default = "ten")]
412    pub limit: usize,
413    #[serde(default)]
414    pub context: ContextOptions,
415}
416fn ten() -> usize {
417    10
418}
419#[derive(Clone, Debug, Serialize)]
420pub struct ContextHit {
421    pub id: String,
422    pub score: f32,
423    pub metadata: Value,
424    pub text: Option<String>,
425    pub chunk: Option<Chunk>,
426    pub sources: Vec<usize>,
427    #[serde(skip_serializing_if = "Option::is_none")]
428    pub expanded_from: Option<String>,
429}
430#[derive(Clone, Debug, Serialize)]
431pub struct RetrievalTrace {
432    pub generation: u64,
433    pub eligible_documents: usize,
434    pub channels: Vec<Value>,
435    pub fused_candidates: usize,
436    pub reranked_candidates: usize,
437    pub channel_agreement: Option<f32>,
438    pub elapsed_ms: f64,
439}
440#[derive(Clone, Debug, Serialize)]
441pub struct RetrievalResponse {
442    pub matches: Vec<ContextHit>,
443    pub trace: RetrievalTrace,
444}
445
446#[derive(Default)]
447struct ContextSelection {
448    matches: Vec<ContextHit>,
449    ids: HashSet<String>,
450    texts: HashSet<blake3::Hash>,
451    parents: HashMap<String, usize>,
452}
453
454impl ContextSelection {
455    fn accepts(
456        &self,
457        id: &str,
458        fields: &Fields,
459        text_key: Option<blake3::Hash>,
460        options: &ContextOptions,
461    ) -> bool {
462        !self.ids.contains(id)
463            && text_key.is_none_or(|key| !self.texts.contains(&key))
464            && fields.chunk.as_ref().is_none_or(|chunk| {
465                options.per_parent.is_none_or(|limit| {
466                    self.parents.get(&chunk.parent).copied().unwrap_or(0) < limit
467                })
468            })
469    }
470
471    fn push(&mut self, hit: ContextHit, text_key: Option<blake3::Hash>) {
472        self.ids.insert(hit.id.clone());
473        if let Some(key) = text_key {
474            self.texts.insert(key);
475        }
476        if let Some(chunk) = &hit.chunk {
477            *self.parents.entry(chunk.parent.clone()).or_default() += 1;
478        }
479        self.matches.push(hit);
480    }
481}
482
483impl Fields {
484    pub(super) fn has_dense(&self, field: &str) -> bool {
485        matches!(
486            self.representations.get(field),
487            Some(StoredRepresentation::Dense { .. })
488        )
489    }
490    pub(super) fn relocate(
491        &mut self,
492        source: &[u8],
493        destination: &FixedVectorStore,
494    ) -> Result<(), IndexError> {
495        for representation in self.representations.values_mut() {
496            match representation {
497                StoredRepresentation::Dense { location, .. }
498                | StoredRepresentation::Multivector { location, .. } => {
499                    *location = destination.copy_record(source, *location)?;
500                }
501                StoredRepresentation::Sparse(_) => (),
502            }
503        }
504        Ok(())
505    }
506    pub(super) fn prepare(
507        document: &RetrievalDocument,
508        stores: &SegmentStores,
509    ) -> Result<Self, IndexError> {
510        if document.id.is_empty()
511            || document.id.len() > 4096
512            || document.representations.len() > 32
513            || document.text.as_ref().is_some_and(|s| s.len() > 1_048_576)
514        {
515            return Err(invalid(
516                "document ID, text or representation count exceeds limits",
517            ));
518        }
519        if let Some(chunk) = &document.chunk {
520            if chunk.parent.is_empty() || chunk.parent.len() > 4096 {
521                return Err(invalid("invalid chunk parent"));
522            }
523        }
524        let mut fields = Self {
525            text: document.text.clone(),
526            chunk: document.chunk.clone(),
527            ..Self::default()
528        };
529        for (name, value) in &document.representations {
530            if name.is_empty() || name.len() > 128 {
531                return Err(invalid("representation names require 1..=128 bytes"));
532            }
533            let stored = match value {
534                Representation::Sparse { vector } => {
535                    if vector.len() > 65_536 {
536                        return Err(invalid("sparse representation exceeds 65536 features"));
537                    }
538                    StoredRepresentation::Sparse(vector.canonicalized().map_err(sparse_error)?)
539                }
540                Representation::Dense { vector } => {
541                    validate_matrix(std::slice::from_ref(vector))?;
542                    StoredRepresentation::Dense {
543                        location: stores.fde.put(&normalize(vector))?,
544                        dimension: vector.len(),
545                    }
546                }
547                Representation::Multivector { vectors } => {
548                    validate_matrix(vectors)?;
549                    let flat: Vec<_> = vectors.iter().flat_map(|v| normalize(v)).collect();
550                    StoredRepresentation::Multivector {
551                        location: stores.fde.put(&flat)?,
552                        dimension: vectors[0].len(),
553                        tokens: vectors.len(),
554                    }
555                }
556            };
557            fields.representations.insert(name.clone(), stored);
558        }
559        Ok(fields)
560    }
561    pub(super) fn verify(&self, mapped: &[u8]) -> Result<(), IndexError> {
562        for representation in self.representations.values() {
563            match representation {
564                StoredRepresentation::Dense {
565                    location,
566                    dimension,
567                } => {
568                    verify_record(mapped, *location, true)?;
569                    let v = FixedVectorStore::get(mapped, *location, *dimension)?;
570                    validate_matrix(&[v.to_vec()])?;
571                }
572                StoredRepresentation::Multivector {
573                    location,
574                    dimension,
575                    tokens,
576                } => {
577                    verify_record(mapped, *location, true)?;
578                    let size = dimension
579                        .checked_mul(*tokens)
580                        .ok_or_else(|| invalid("representation size overflow"))?;
581                    let v = FixedVectorStore::get(mapped, *location, size)?;
582                    if *dimension == 0 || *tokens == 0 || v.iter().any(|v| !v.is_finite()) {
583                        return Err(invalid("invalid multivector record"));
584                    }
585                }
586                StoredRepresentation::Sparse(v) => {
587                    v.canonicalized().map_err(sparse_error)?;
588                }
589            }
590        }
591        Ok(())
592    }
593}
594fn validate_matrix(vectors: &[Vector]) -> Result<(), IndexError> {
595    let dimension = vectors.first().map_or(0, Vec::len);
596    if vectors.is_empty()
597        || vectors.len() > 8192
598        || dimension == 0
599        || dimension > 65_536
600        || vectors
601            .len()
602            .checked_mul(dimension)
603            .is_none_or(|n| n > 16_777_216)
604        || vectors
605            .iter()
606            .any(|v| v.len() != dimension || v.iter().any(|x| !x.is_finite()))
607    {
608        return Err(invalid("invalid or excessive vector shape"));
609    }
610    Ok(())
611}
612fn top(mut scores: Vec<(String, f32)>, limit: usize) -> Vec<(String, f32)> {
613    let order =
614        |a: &(String, f32), b: &(String, f32)| b.1.total_cmp(&a.1).then_with(|| a.0.cmp(&b.0));
615    if scores.len() > limit {
616        scores.select_nth_unstable_by(limit, order);
617        scores.truncate(limit);
618    }
619    scores.sort_unstable_by(order);
620    scores
621}
622
623impl MultiVectorIndex {
624    pub(super) fn dense_vector<'a>(
625        &self,
626        s: &'a State,
627        id: &str,
628        field: &str,
629    ) -> Result<&'a [f32], IndexError> {
630        match s.documents[id].fields.representations.get(field) {
631            Some(StoredRepresentation::Dense {
632                location,
633                dimension,
634            }) => Ok(FixedVectorStore::get(
635                s.record_fde(&s.documents[id]),
636                *location,
637                *dimension,
638            )?),
639            _ => Err(invalid("missing dense vector")),
640        }
641    }
642    pub fn retrieve(&self, request: &RetrieveRequest) -> Result<RetrievalResponse, IndexError> {
643        let started = Instant::now();
644        if request.prefetch.is_empty()
645            || request.prefetch.len() > 8
646            || request.limit == 0
647            || request.limit > 10_000
648            || request.context.neighbors > 8
649            || request.context.per_parent == Some(0)
650            || request
651                .context
652                .mmr
653                .is_some_and(|v| !v.is_finite() || !(0.0..=1.0).contains(&v))
654        {
655            return Err(invalid("invalid retrieval or context budgets"));
656        }
657        if let Some(filter) = &request.filter {
658            filter.validate(0)?;
659        }
660        let s = self.snapshot();
661        let eligible: HashSet<_> = s
662            .documents
663            .iter()
664            .filter(|(_, d)| {
665                request
666                    .filter
667                    .as_ref()
668                    .is_none_or(|f| f.matches(&d.metadata))
669            })
670            .map(|(id, _)| id.as_str())
671            .collect();
672        let mut channels = Vec::new();
673        let mut lists = Vec::new();
674        for channel in &request.prefetch {
675            let at = Instant::now();
676            let limit = match channel {
677                Channel::Bm25 { limit, .. }
678                | Channel::Sparse { limit, .. }
679                | Channel::Dense { limit, .. }
680                | Channel::Multivector { limit, .. } => *limit,
681            };
682            if limit == 0 || limit > 100_000 {
683                return Err(invalid("channel limit must be in 1..=100000"));
684            }
685            let allowed = |number| {
686                s.retrieval
687                    .ids
688                    .get(&number)
689                    .is_some_and(|id| eligible.contains(id.as_str()))
690            };
691            let external = |hits: Vec<(u64, f32)>| {
692                hits.into_iter()
693                    .map(|(id, score)| (s.retrieval.ids[&id].clone(), score))
694                    .collect::<Vec<_>>()
695            };
696            let tie_break = |a, b| s.retrieval.ids[&a].cmp(&s.retrieval.ids[&b]);
697            let (scores, backend) = match channel {
698                Channel::Bm25 { text, k1, b, .. } => {
699                    if text.len() > 65_536 {
700                        return Err(invalid("query text exceeds 65536 bytes"));
701                    }
702                    // BM25 query weights are raw analyzed query term
703                    // frequencies: repeated query terms contribute repeatedly.
704                    let query = SparseVector::from_pairs(
705                        s.retrieval.analyzer.analyze(text).into_iter().filter_map(
706                            |(term, count)| {
707                                s.retrieval.vocabulary.get(&term).map(|&id| (id, count))
708                            },
709                        ),
710                    );
711                    (
712                        external(
713                            s.retrieval
714                                .lexical
715                                .search_bm25_filtered_by(&query, limit, *k1, *b, allowed, tie_break)
716                                .map_err(sparse_error)?,
717                        ),
718                        "bm25",
719                    )
720                }
721                Channel::Sparse { field, vector, .. } => {
722                    let index = s
723                        .retrieval
724                        .sparse
725                        .get(field)
726                        .ok_or_else(|| invalid(format!("unknown sparse field {field:?}")))?;
727                    (
728                        external(
729                            index
730                                .search_dot_filtered_by(vector, limit, allowed, tie_break)
731                                .map_err(sparse_error)?,
732                        ),
733                        "sparse_dot",
734                    )
735                }
736                Channel::Dense {
737                    field,
738                    vector,
739                    backend,
740                    ef_search,
741                    ..
742                } => {
743                    validate_matrix(std::slice::from_ref(vector))?;
744                    if s.retrieval.dense_dimension(field)? != vector.len()
745                        || !["auto", "exact", "hnsw"].contains(&backend.as_str())
746                        || *ef_search == 0
747                        || *ef_search > 65_536
748                    {
749                        return Err(invalid("invalid dense query dimension or backend"));
750                    }
751                    if request.filter.is_none()
752                        && backend != "exact"
753                        && s.named_ann.contains_key(field)
754                    {
755                        (
756                            self.ann_scores(
757                                &s,
758                                &s.named_ann[field],
759                                &normalize(vector),
760                                limit,
761                                *ef_search,
762                            )?,
763                            "hnsw_dense",
764                        )
765                    } else {
766                        if backend == "hnsw" && !s.named_ann.contains_key(field) {
767                            return Err(invalid("dense ANN not built"));
768                        }
769                        (
770                            self.named_scores(
771                                &s,
772                                field,
773                                std::slice::from_ref(vector),
774                                false,
775                                &eligible,
776                                limit,
777                            )?,
778                            "exact_dense",
779                        )
780                    }
781                }
782                Channel::Multivector {
783                    field: Some(field),
784                    vectors,
785                    backend,
786                    ..
787                } => {
788                    if backend != "auto" && backend != "exact" {
789                        return Err(invalid("named multivectors support exact or auto"));
790                    }
791                    (
792                        self.named_scores(&s, field, vectors, true, &eligible, limit)?,
793                        "exact_maxsim",
794                    )
795                }
796                Channel::Multivector {
797                    field: None,
798                    vectors,
799                    backend,
800                    ef_search,
801                    ..
802                } => {
803                    self.validate(vectors)?;
804                    if vectors.len() > 1024
805                        || *ef_search == 0
806                        || *ef_search > 65_536
807                        || !["auto", "exact", "hnsw"].contains(&backend.as_str())
808                    {
809                        return Err(invalid("invalid multivector backend or budget"));
810                    }
811                    let normalized: Vec<_> = vectors.iter().map(|v| normalize(v)).collect();
812                    if request.filter.is_none() && backend != "exact" && s.fde_ann.is_some() {
813                        (
814                            self.ann_fde_scores(
815                                &s,
816                                &self.fde.encode_query(&normalized),
817                                limit,
818                                *ef_search,
819                            )?,
820                            "hnsw_fde",
821                        )
822                    } else {
823                        if backend == "hnsw" && s.fde_ann.is_none() {
824                            return Err(invalid("FDE ANN not built"));
825                        }
826                        let scores = self.exact_fde_scores_filtered(
827                            &s,
828                            &normalized,
829                            Some(limit),
830                            Some(&eligible),
831                        )?;
832                        (scores, "exact_fde")
833                    }
834                }
835            };
836            channels.push(serde_json::json!({"backend": backend,"candidates": scores.len(),"elapsed_ms": at.elapsed().as_secs_f64()*1000.}));
837            lists.push(scores);
838        }
839        let agreement = if lists.len() > 1 {
840            let head: HashSet<_> = lists[0]
841                .iter()
842                .take(request.limit)
843                .map(|h| h.0.as_str())
844                .collect();
845            Some(
846                lists[1..]
847                    .iter()
848                    .map(|list| {
849                        let other: HashSet<_> = list
850                            .iter()
851                            .take(request.limit)
852                            .map(|h| h.0.as_str())
853                            .collect();
854                        let union = head.union(&other).count();
855                        if union == 0 {
856                            0.
857                        } else {
858                            head.intersection(&other).count() as f32 / union as f32
859                        }
860                    })
861                    .fold(1., f32::min),
862            )
863        } else {
864            None
865        };
866        let mut fused: HashMap<String, (f32, Vec<usize>)> = HashMap::new();
867        match &request.fusion {
868            Fusion::Rrf { k } if !k.is_finite() || *k < 0. => {
869                return Err(invalid("RRF k must be finite and nonnegative"));
870            }
871            Fusion::Weighted { weights }
872                if weights.len() != lists.len()
873                    || weights.iter().any(|w| !w.is_finite() || *w < 0.)
874                    || weights.iter().all(|w| *w == 0.) =>
875            {
876                return Err(invalid(
877                    "weighted fusion requires one nonnegative finite weight per channel",
878                ));
879            }
880            _ => (),
881        }
882        for (channel, list) in lists.iter().enumerate() {
883            for (rank, (id, score)) in list.iter().enumerate() {
884                let contribution = match &request.fusion {
885                    Fusion::Rrf { k } => 1. / (k + rank as f32 + 1.),
886                    Fusion::Weighted { weights } => weights[channel] * score,
887                };
888                let entry = fused.entry(id.clone()).or_default();
889                entry.0 += contribution;
890                if !entry.0.is_finite() {
891                    return Err(invalid("fusion score overflow"));
892                }
893                entry.1.push(channel);
894            }
895        }
896        // One channel retains its score; fusion only combines independent lists.
897        let mut ranked = if lists.len() == 1 && matches!(request.fusion, Fusion::Rrf { .. }) {
898            lists.pop().unwrap()
899        } else {
900            top(
901                fused
902                    .iter()
903                    .map(|(id, (score, _))| (id.clone(), *score))
904                    .collect(),
905                fused.len(),
906            )
907        };
908        let fused_candidates = ranked.len();
909        let mut reranked = 0;
910        if let Some(rerank) = &request.rerank {
911            if rerank.limit < request.limit || rerank.limit > 100_000 {
912                return Err(invalid(
913                    "rerank limit must cover result limit and be <=100000",
914                ));
915            }
916            let budget = if let Some(policy) = &rerank.adaptive {
917                if policy.min_candidates < request.limit
918                    || policy.min_candidates > rerank.limit
919                    || !policy.agreement_threshold.is_finite()
920                    || !(0.0..=1.0).contains(&policy.agreement_threshold)
921                    || agreement.is_none()
922                {
923                    return Err(invalid(
924                        "invalid adaptive rerank policy; needs multiple channels",
925                    ));
926                }
927                if agreement.unwrap() >= policy.agreement_threshold {
928                    policy.min_candidates
929                } else {
930                    rerank.limit
931                }
932            } else {
933                rerank.limit
934            };
935            ranked.truncate(budget);
936            reranked = ranked.len();
937            if let Some(field) = &rerank.field {
938                let pool: HashSet<_> = ranked.iter().map(|(id, _)| id.as_str()).collect();
939                ranked =
940                    self.named_scores(&s, field, &rerank.vectors, true, &pool, rerank.limit)?;
941                if ranked.len() != reranked {
942                    return Err(invalid("rerank field missing from a candidate"));
943                }
944            } else {
945                self.validate(&rerank.vectors)?;
946                if ranked.iter().any(|(id, _)| s.documents[id].tokens == 0) {
947                    return Err(invalid(
948                        "default multivector missing from a rerank candidate",
949                    ));
950                }
951                let normalized: Vec<_> = rerank.vectors.iter().map(|v| normalize(v)).collect();
952                ranked = self
953                    .rescore(&s, &normalized, ranked, rerank.limit, Some(rerank.limit))?
954                    .into_iter()
955                    .map(|h| (h.id, h.score))
956                    .collect();
957            }
958        }
959        let matches = self.context(&s, ranked, &fused, &eligible, request)?;
960        Ok(RetrievalResponse {
961            matches,
962            trace: RetrievalTrace {
963                generation: s.generation,
964                eligible_documents: eligible.len(),
965                channels,
966                fused_candidates,
967                reranked_candidates: reranked,
968                channel_agreement: agreement,
969                elapsed_ms: started.elapsed().as_secs_f64() * 1000.,
970            },
971        })
972    }
973
974    fn named_scores(
975        &self,
976        s: &State,
977        field: &str,
978        query: &[Vector],
979        multivector: bool,
980        eligible: &HashSet<&str>,
981        limit: usize,
982    ) -> Result<Vec<(String, f32)>, IndexError> {
983        validate_matrix(query)?;
984        let expected = if multivector {
985            FieldSchema::Multivector {
986                dimension: query[0].len(),
987            }
988        } else {
989            FieldSchema::Dense {
990                dimension: query[0].len(),
991            }
992        };
993        if s.retrieval.schema.get(field) != Some(&expected) {
994            return Err(invalid(format!(
995                "unknown field or query dimension/kind mismatch: {field:?}"
996            )));
997        }
998        let normalized: Vec<_> = query.iter().map(|v| normalize(v)).collect();
999        let scores = eligible
1000            .par_iter()
1001            .filter_map(|&id| {
1002                let d = s.documents.get(id)?;
1003                d.fields.representations.get(field).map(|r| (id, d, r))
1004            })
1005            .map(|(id, d, r)| {
1006                let (location, dimension, count) = match r {
1007                    StoredRepresentation::Dense {
1008                        location,
1009                        dimension,
1010                    } => (*location, *dimension, 1),
1011                    StoredRepresentation::Multivector {
1012                        location,
1013                        dimension,
1014                        tokens,
1015                    } => (*location, *dimension, *tokens),
1016                    _ => unreachable!(),
1017                };
1018                let vector = FixedVectorStore::get(s.record_fde(d), location, dimension * count)?;
1019                let score = if multivector {
1020                    maxsim_flat(&normalized, vector, dimension)
1021                } else {
1022                    dot(&normalized[0], vector)
1023                };
1024                Ok::<_, IndexError>((id.to_owned(), score))
1025            })
1026            .collect::<Result<Vec<_>, _>>()?;
1027        Ok(top(scores, limit))
1028    }
1029
1030    fn context(
1031        &self,
1032        s: &State,
1033        ranked: Vec<(String, f32)>,
1034        fused: &HashMap<String, (f32, Vec<usize>)>,
1035        eligible: &HashSet<&str>,
1036        request: &RetrieveRequest,
1037    ) -> Result<Vec<ContextHit>, IndexError> {
1038        let options = &request.context;
1039        let mmr_vectors = if options.mmr.is_some() {
1040            if ranked.len() > 4096 {
1041                return Err(invalid("MMR pool exceeds 4096 candidates"));
1042            }
1043            let field = options
1044                .diversity_field
1045                .as_ref()
1046                .ok_or_else(|| invalid("MMR requires a dense diversity_field"))?;
1047            s.retrieval.dense_dimension(field)?;
1048            Some(
1049                ranked
1050                    .iter()
1051                    .map(|(id, _)| self.dense_vector(s, id, field))
1052                    .collect::<Result<Vec<_>, _>>()?,
1053            )
1054        } else {
1055            if options.diversity_field.is_some() {
1056                return Err(invalid("diversity_field requires mmr"));
1057            }
1058            None
1059        };
1060        let text_key = |id: &str| {
1061            if options.deduplicate {
1062                s.documents[id]
1063                    .fields
1064                    .text
1065                    .as_ref()
1066                    .map(|text| blake3::hash(text.as_bytes()))
1067            } else {
1068                None
1069            }
1070        };
1071        // Cache hashes only for MMR's repeatedly inspected, bounded pool. Plain
1072        // ranking hashes documents lazily, stopping as soon as context is full.
1073        let candidate_keys: Vec<_> = if mmr_vectors.is_some() {
1074            ranked.iter().map(|(id, _)| text_key(id)).collect()
1075        } else {
1076            Vec::new()
1077        };
1078        let min = ranked
1079            .iter()
1080            .map(|(_, score)| f64::from(*score))
1081            .fold(f64::INFINITY, f64::min);
1082        let max = ranked
1083            .iter()
1084            .map(|(_, score)| f64::from(*score))
1085            .fold(f64::NEG_INFINITY, f64::max);
1086        let relevance = |score: f32| {
1087            if max > min {
1088                ((f64::from(score) - min) / (max - min)) as f32
1089            } else {
1090                1.
1091            }
1092        };
1093        let mut selected = ContextSelection::default();
1094        let add = |selected: &mut ContextSelection,
1095                   id: &str,
1096                   score: f32,
1097                   expanded_from: Option<String>,
1098                   key: Option<blake3::Hash>| {
1099            let d = &s.documents[id];
1100            if !eligible.contains(id) || !selected.accepts(id, &d.fields, key, options) {
1101                return false;
1102            }
1103            selected.push(
1104                ContextHit {
1105                    id: id.to_owned(),
1106                    score,
1107                    metadata: d.metadata.clone(),
1108                    text: d.fields.text.clone(),
1109                    chunk: d.fields.chunk.clone(),
1110                    sources: fused.get(id).map(|v| v.1.clone()).unwrap_or_default(),
1111                    expanded_from,
1112                },
1113                key,
1114            );
1115            true
1116        };
1117        let mut remaining = vec![true; ranked.len()];
1118        let mut redundancy = vec![f32::NEG_INFINITY; ranked.len()];
1119        let mut seeds = 0usize;
1120        let mut cursor = 0;
1121        while selected.matches.len() < request.limit {
1122            let next = if let Some(lambda) = options.mmr {
1123                let mut best: Option<(usize, f32)> = None;
1124                for (i, (id, score)) in ranked.iter().enumerate() {
1125                    if !remaining[i] {
1126                        continue;
1127                    }
1128                    if !eligible.contains(id.as_str())
1129                        || !selected.accepts(
1130                            id,
1131                            &s.documents[id].fields,
1132                            candidate_keys[i],
1133                            options,
1134                        )
1135                    {
1136                        // Group and duplicate exclusions cannot become eligible
1137                        // later; excluded candidates must not affect diversity.
1138                        remaining[i] = false;
1139                        continue;
1140                    }
1141                    let value = if seeds == 0 {
1142                        relevance(*score)
1143                    } else {
1144                        lambda * relevance(*score) - (1. - lambda) * redundancy[i]
1145                    };
1146                    if best.is_none_or(|(_, previous)| value > previous) {
1147                        best = Some((i, value));
1148                    }
1149                }
1150                best.map(|(i, _)| i)
1151            } else if cursor < ranked.len() {
1152                let next = cursor;
1153                cursor += 1;
1154                Some(next)
1155            } else {
1156                None
1157            };
1158            let Some(next) = next else { break };
1159            remaining[next] = false;
1160            let (id, score) = &ranked[next];
1161            let key = if mmr_vectors.is_some() {
1162                candidate_keys[next]
1163            } else {
1164                text_key(id)
1165            };
1166            if !add(&mut selected, id, *score, None, key) {
1167                continue;
1168            }
1169            seeds += 1;
1170            if selected.matches.len() >= request.limit {
1171                break;
1172            }
1173            // MMR diversifies accepted retrieval seeds. Neighbor expansion adds
1174            // context around those seeds without requiring a dense neighbor field.
1175            if let Some(vectors) = &mmr_vectors {
1176                for i in 0..ranked.len() {
1177                    if remaining[i] {
1178                        redundancy[i] = redundancy[i].max(dot(vectors[i], vectors[next]));
1179                    }
1180                }
1181            }
1182            if options.neighbors > 0 {
1183                if let Some(chunk) = &s.documents[id].fields.chunk {
1184                    for neighbor in s.retrieval.neighbors(chunk, options.neighbors) {
1185                        if !eligible.contains(neighbor) || selected.ids.contains(neighbor) {
1186                            continue;
1187                        }
1188                        add(
1189                            &mut selected,
1190                            neighbor,
1191                            *score,
1192                            Some(id.clone()),
1193                            text_key(neighbor),
1194                        );
1195                        if selected.matches.len() >= request.limit {
1196                            break;
1197                        }
1198                    }
1199                }
1200            }
1201        }
1202        Ok(selected.matches)
1203    }
1204}
1205
1206#[cfg(test)]
1207mod tests {
1208    use super::*;
1209    use crate::storage::FAIL_COMMIT;
1210    use serde_json::json;
1211
1212    fn document(
1213        id: &str,
1214        text: &str,
1215        vector: Vector,
1216        tenant: &str,
1217        position: u32,
1218    ) -> RetrievalDocument {
1219        serde_json::from_value(json!({"id":id,"text":text,"metadata":{"tenant":tenant,"year":2026},"chunk":{"parent":tenant,"position":position},"representations":{"semantic":{"kind":"dense","vector":vector},"tokens":{"kind":"multivector","vectors":[vector]},"sparse":{"kind":"sparse","vector":{"indices":[u32::MAX],"values":[position as f32+1.]}}}})).unwrap()
1220    }
1221
1222    fn schema_document(id: &str, representation: Representation) -> RetrievalDocument {
1223        RetrievalDocument {
1224            id: id.into(),
1225            representations: BTreeMap::from([("semantic".into(), representation)]),
1226            ..RetrievalDocument::default()
1227        }
1228    }
1229
1230    fn assert_schema_mismatch(error: IndexError) {
1231        assert!(
1232            error.to_string().contains("different kind or dimension"),
1233            "unexpected error: {error}"
1234        );
1235    }
1236    fn request() -> RetrieveRequest {
1237        serde_json::from_value(json!({"prefetch":[{"kind":"dense","field":"semantic","vector":[1.,0.],"limit":3},{"kind":"bm25","text":"E123 repair","limit":3}],"limit":3,"filter":{"op":"eq","field":"tenant","value":"a"}})).unwrap()
1238    }
1239    fn index(path: &Path) -> MultiVectorIndex {
1240        let index = MultiVectorIndex::open(path, IndexConfig::new(2)).unwrap();
1241        index
1242            .upsert_records(vec![
1243                document("a", "E123 repair", vec![0.8, 0.6], "a", 0),
1244                document("b", "hardware guide", vec![1., 0.], "a", 1),
1245                document("c", "E123 E123 repair repair", vec![1., 0.], "denied", 0),
1246            ])
1247            .unwrap();
1248        index
1249    }
1250    #[test]
1251    fn hybrid_filters_before_top_k_and_preserves_named_fields_on_reopen() {
1252        let dir = tempfile::tempdir().unwrap();
1253        let index = index(dir.path());
1254        let mut query = request();
1255        for channel in &mut query.prefetch {
1256            match channel {
1257                Channel::Dense { limit, .. } | Channel::Bm25 { limit, .. } => *limit = 1,
1258                _ => (),
1259            }
1260        }
1261        let response = index.retrieve(&query).unwrap();
1262        assert_eq!(response.trace.eligible_documents, 2);
1263        assert_eq!(
1264            response
1265                .matches
1266                .iter()
1267                .map(|h| h.id.as_str())
1268                .collect::<Vec<_>>(),
1269            vec!["a", "b"]
1270        );
1271        assert!(response.matches.iter().all(|h| h.metadata["tenant"] == "a"));
1272        let mut sparse = query.clone();
1273        sparse.prefetch = vec![Channel::Sparse {
1274            field: "sparse".into(),
1275            vector: SparseVector::from_pairs([(u32::MAX, 1.)]),
1276            limit: 10,
1277        }];
1278        assert_eq!(index.retrieve(&sparse).unwrap().matches[0].id, "b");
1279        query.rerank = Some(Rerank {
1280            field: Some("tokens".into()),
1281            vectors: vec![vec![1., 0.]],
1282            limit: 3,
1283            adaptive: None,
1284        });
1285        assert_eq!(index.retrieve(&query).unwrap().matches[0].id, "b");
1286        let before = serde_json::to_value(index.retrieve(&query).unwrap().matches).unwrap();
1287        drop(index);
1288        let reopened = MultiVectorIndex::open(dir.path(), IndexConfig::new(2)).unwrap();
1289        assert_eq!(
1290            serde_json::to_value(reopened.retrieve(&query).unwrap().matches).unwrap(),
1291            before
1292        );
1293        assert!(reopened.delete("b").unwrap());
1294        assert_eq!(reopened.retrieve(&sparse).unwrap().matches[0].id, "a");
1295        reopened
1296            .upsert_records(vec![document("a", "different", vec![0., 1.], "a", 0)])
1297            .unwrap();
1298        query.prefetch.retain(|c| matches!(c, Channel::Bm25 { .. }));
1299        query.rerank = None;
1300        assert!(reopened.retrieve(&query).unwrap().matches.is_empty());
1301    }
1302
1303    fn analyzer_config(analyzer: TextAnalyzer) -> IndexConfig {
1304        IndexConfig {
1305            analyzer,
1306            ..IndexConfig::new(2)
1307        }
1308    }
1309    fn bm25_only(text: &str) -> RetrieveRequest {
1310        serde_json::from_value(json!({
1311            "prefetch": [{"kind": "bm25", "text": text, "limit": 10}],
1312            "limit": 10
1313        }))
1314        .unwrap()
1315    }
1316
1317    #[test]
1318    fn bm25_query_term_frequency_weights_repeated_terms() {
1319        let dir = tempfile::tempdir().unwrap();
1320        let index = MultiVectorIndex::open(dir.path(), IndexConfig::new(2)).unwrap();
1321        index
1322            .upsert_records(vec![
1323                document("alpha", "alpha shared", vec![1., 0.], "t", 0),
1324                document("beta", "beta shared", vec![0., 1.], "t", 0),
1325            ])
1326            .unwrap();
1327        // Equal-length documents and equal df tie on a single occurrence of
1328        // each term; the stable tie-break puts "alpha" first.
1329        let tied = index.retrieve(&bm25_only("alpha beta")).unwrap();
1330        assert_eq!(tied.matches[0].id, "alpha");
1331        assert_eq!(tied.matches[0].score, tied.matches[1].score);
1332        // A repeated query term must double its contribution so "beta" wins.
1333        let repeated = index.retrieve(&bm25_only("alpha beta beta")).unwrap();
1334        assert_eq!(repeated.matches[0].id, "beta");
1335        let beta = repeated.matches.iter().find(|h| h.id == "beta").unwrap();
1336        let alpha = repeated.matches.iter().find(|h| h.id == "alpha").unwrap();
1337        assert_eq!(beta.score, tied.matches[0].score * 2.0);
1338        assert_eq!(alpha.score, tied.matches[0].score);
1339    }
1340
1341    #[test]
1342    fn english_analyzer_stems_queries_and_survives_reopen() {
1343        let dir = tempfile::tempdir().unwrap();
1344        let index =
1345            MultiVectorIndex::open(dir.path(), analyzer_config(TextAnalyzer::english())).unwrap();
1346        index
1347            .upsert_records(vec![
1348                document("runner", "The runner runs fast", vec![1., 0.], "t", 0),
1349                document("walker", "walking quickly", vec![0., 1.], "t", 0),
1350            ])
1351            .unwrap();
1352        // "running" stems to "run", matching the indexed "runs"; "the" is a
1353        // stop word on both sides, so it cannot retrieve anything.
1354        let ranked = |index: &MultiVectorIndex, text: &str| {
1355            index
1356                .retrieve(&bm25_only(text))
1357                .unwrap()
1358                .matches
1359                .iter()
1360                .map(|h| h.id.clone())
1361                .collect::<Vec<_>>()
1362        };
1363        assert_eq!(ranked(&index, "running"), ["runner"]);
1364        assert!(ranked(&index, "the").is_empty());
1365        drop(index);
1366        let reopened =
1367            MultiVectorIndex::open(dir.path(), analyzer_config(TextAnalyzer::english())).unwrap();
1368        assert_eq!(ranked(&reopened, "running"), ["runner"]);
1369        drop(reopened);
1370        // The analyzer is part of the persisted configuration contract.
1371        let Err(mismatch) = MultiVectorIndex::open(dir.path(), IndexConfig::new(2)) else {
1372            panic!("reopening with a different analyzer must fail");
1373        };
1374        assert!(matches!(mismatch, IndexError::Config { .. }));
1375    }
1376
1377    #[test]
1378    fn analyzer_config_is_validated_and_legacy_configs_default_to_plain() {
1379        let dir = tempfile::tempdir().unwrap();
1380        let Err(error) = MultiVectorIndex::open(
1381            dir.path(),
1382            analyzer_config(TextAnalyzer {
1383                max_token_length: Some(0),
1384                ..TextAnalyzer::plain()
1385            }),
1386        ) else {
1387            panic!("invalid analyzer config must be rejected");
1388        };
1389        assert!(error.to_string().contains("max_token_length"));
1390        let legacy: IndexConfig = serde_json::from_value(json!({
1391            "dimension": 2, "centroids": 2, "residual_bits": 2, "probes": 2,
1392            "fde_repetitions": 2, "fde_ksim": 2, "fde_projected": 2
1393        }))
1394        .unwrap();
1395        assert_eq!(legacy.analyzer, TextAnalyzer::plain());
1396        let custom: TextAnalyzer = serde_json::from_value(json!({
1397            "stem": true, "stopwords": "english",
1398            "ascii_folding": true, "max_token_length": 32
1399        }))
1400        .unwrap();
1401        assert!(custom.stem && custom.max_token_length == Some(32));
1402    }
1403
1404    #[test]
1405    fn named_field_schema_survives_deleting_every_document_and_reopen() {
1406        let dir = tempfile::tempdir().unwrap();
1407        let index = MultiVectorIndex::open(dir.path(), IndexConfig::new(2)).unwrap();
1408        index
1409            .upsert_records(vec![schema_document(
1410                "dense",
1411                Representation::Dense {
1412                    vector: vec![1., 0.],
1413                },
1414            )])
1415            .unwrap();
1416        assert!(index.delete("dense").unwrap());
1417
1418        assert_schema_mismatch(
1419            index
1420                .upsert_records(vec![schema_document(
1421                    "wrong-before-reopen",
1422                    Representation::Sparse {
1423                        vector: SparseVector::from_pairs([(1, 1.)]),
1424                    },
1425                )])
1426                .unwrap_err(),
1427        );
1428        drop(index);
1429
1430        let reopened = MultiVectorIndex::open(dir.path(), IndexConfig::new(2)).unwrap();
1431        assert_eq!(reopened.stats().documents, 0);
1432        assert_schema_mismatch(
1433            reopened
1434                .upsert_records(vec![schema_document(
1435                    "wrong-after-reopen",
1436                    Representation::Sparse {
1437                        vector: SparseVector::from_pairs([(1, 1.)]),
1438                    },
1439                )])
1440                .unwrap_err(),
1441        );
1442        reopened
1443            .upsert_records(vec![schema_document(
1444                "same-schema",
1445                Representation::Dense {
1446                    vector: vec![0., 1.],
1447                },
1448            )])
1449            .unwrap();
1450    }
1451
1452    #[test]
1453    fn legacy_manifest_infers_and_persists_named_field_schema() {
1454        let dir = tempfile::tempdir().unwrap();
1455        let index = MultiVectorIndex::open(dir.path(), IndexConfig::new(2)).unwrap();
1456        index
1457            .upsert_records(vec![schema_document(
1458                "legacy",
1459                Representation::Dense {
1460                    vector: vec![1., 0.],
1461                },
1462            )])
1463            .unwrap();
1464        drop(index);
1465
1466        let path = dir.path().join("manifest.json");
1467        let mut envelope: ManifestEnvelope =
1468            serde_json::from_slice(&fs::read(&path).unwrap()).unwrap();
1469        let mut payload: Value = serde_json::from_str(&envelope.manifest).unwrap();
1470        payload
1471            .as_object_mut()
1472            .unwrap()
1473            .remove("representation_schema");
1474        envelope.manifest = serde_json::to_string(&payload).unwrap();
1475        envelope.checksum_blake3 = blake3::hash(envelope.manifest.as_bytes())
1476            .to_hex()
1477            .to_string();
1478        fs::write(&path, serde_json::to_vec(&envelope).unwrap()).unwrap();
1479
1480        let restored = MultiVectorIndex::open(dir.path(), IndexConfig::new(2)).unwrap();
1481        assert!(restored.delete("legacy").unwrap());
1482        drop(restored);
1483
1484        let envelope: ManifestEnvelope = serde_json::from_slice(&fs::read(&path).unwrap()).unwrap();
1485        let manifest: Manifest = serde_json::from_str(&envelope.manifest).unwrap();
1486        assert_eq!(
1487            manifest.representation_schema.get("semantic"),
1488            Some(&FieldSchema::Dense { dimension: 2 })
1489        );
1490
1491        let reopened = MultiVectorIndex::open(dir.path(), IndexConfig::new(2)).unwrap();
1492        assert_schema_mismatch(
1493            reopened
1494                .upsert_records(vec![schema_document(
1495                    "wrong",
1496                    Representation::Multivector {
1497                        vectors: vec![vec![1., 0.]],
1498                    },
1499                )])
1500                .unwrap_err(),
1501        );
1502    }
1503    #[test]
1504    fn context_grouping_expansion_deduplication_and_mmr_obey_scope() {
1505        let dir = tempfile::tempdir().unwrap();
1506        let index = index(dir.path());
1507        index
1508            .upsert_records(vec![document(
1509                "duplicate",
1510                "E123 repair",
1511                vec![0.8, 0.6],
1512                "a",
1513                2,
1514            )])
1515            .unwrap();
1516        let mut query = request();
1517        query.prefetch.retain(|c| matches!(c, Channel::Bm25 { .. }));
1518        query.context = ContextOptions {
1519            neighbors: 1,
1520            deduplicate: true,
1521            ..Default::default()
1522        };
1523        let response = index.retrieve(&query).unwrap();
1524        assert_eq!(
1525            response
1526                .matches
1527                .iter()
1528                .map(|h| h.id.as_str())
1529                .collect::<Vec<_>>(),
1530            vec!["a", "b"]
1531        );
1532        assert_eq!(response.matches[1].expanded_from.as_deref(), Some("a"));
1533        query.context.per_parent = Some(1);
1534        assert_eq!(index.retrieve(&query).unwrap().matches.len(), 1);
1535        query = request();
1536        query.context.mmr = Some(0.2);
1537        query.context.diversity_field = Some("semantic".into());
1538        assert_eq!(index.retrieve(&query).unwrap().matches.len(), 3);
1539        query.filter = Some(Predicate::And {
1540            filters: vec![Predicate::Range {
1541                field: "year".into(),
1542                gte: Some(2027.),
1543                lte: None,
1544            }],
1545        });
1546        assert!(index.retrieve(&query).unwrap().matches.is_empty());
1547    }
1548    #[test]
1549    fn context_constraints_refill_mmr_and_skip_rejected_seed_neighbors() {
1550        let dir = tempfile::tempdir().unwrap();
1551        let index = MultiVectorIndex::open(dir.path(), IndexConfig::new(2)).unwrap();
1552        index
1553            .upsert_records(vec![
1554                document("a", "same", vec![1., 0.], "p", 0),
1555                document("b", "same", vec![0.99, 0.01], "p", 1),
1556                document("c", "different", vec![0.8, 0.2], "q", 0),
1557            ])
1558            .unwrap();
1559        let mut query: RetrieveRequest = serde_json::from_value(json!({
1560            "prefetch":[{"kind":"dense","field":"semantic","vector":[1.,0.],"limit":3}],
1561            "limit":2,"context":{"mmr":1.,"diversity_field":"semantic","per_parent":1,"deduplicate":true}
1562        })).unwrap();
1563        let result = index.retrieve(&query).unwrap();
1564        assert_eq!(
1565            result
1566                .matches
1567                .iter()
1568                .map(|h| h.id.as_str())
1569                .collect::<Vec<_>>(),
1570            vec!["a", "c"]
1571        );
1572        index
1573            .upsert_records(vec![
1574                document("b", "same", vec![1., 0.], "q", 0),
1575                document("c", "neighbor", vec![0., 1.], "q", 1),
1576            ])
1577            .unwrap();
1578        query = serde_json::from_value(json!({
1579            "prefetch":[{"kind":"bm25","text":"same","limit":3}],
1580            "limit":3,"context":{"neighbors":1,"deduplicate":true}
1581        }))
1582        .unwrap();
1583        // Rejected duplicate b must neither expand c nor retain its old p position.
1584        assert_eq!(
1585            index
1586                .retrieve(&query)
1587                .unwrap()
1588                .matches
1589                .iter()
1590                .map(|h| h.id.as_str())
1591                .collect::<Vec<_>>(),
1592            vec!["a"]
1593        );
1594        query.prefetch = vec![Channel::Bm25 {
1595            text: "same".into(),
1596            limit: 1,
1597            k1: 1.2,
1598            b: 0.75,
1599        }];
1600        query.context = ContextOptions::default();
1601        assert_eq!(index.retrieve(&query).unwrap().matches[0].id, "a");
1602        drop(index);
1603        let index = MultiVectorIndex::open(dir.path(), IndexConfig::new(2)).unwrap();
1604        assert_eq!(index.retrieve(&query).unwrap().matches[0].id, "a");
1605    }
1606
1607    #[test]
1608    fn failed_hybrid_batches_and_compaction_keep_whole_generations() {
1609        for stage in ["fde_partial_write", "manifest_written", "manifest_renamed"] {
1610            let dir = tempfile::tempdir().unwrap();
1611            let index = index(dir.path());
1612            let before = index.stats().generation;
1613            FAIL_COMMIT.with(|f| f.set(Some((stage, 1))));
1614            let result = index.upsert_records(vec![
1615                document("a", "replaced", vec![0., 1.], "a", 0),
1616                document("b", "replaced", vec![0., 1.], "a", 1),
1617            ]);
1618            assert!(result.is_err(), "{stage}");
1619            let committed = stage == "manifest_renamed";
1620            assert_eq!(index.stats().generation, before + u64::from(committed));
1621            let mut q = request();
1622            q.prefetch.retain(|c| matches!(c, Channel::Bm25 { .. }));
1623            assert_eq!(index.retrieve(&q).unwrap().matches.is_empty(), committed);
1624            drop(index);
1625            let index = MultiVectorIndex::open(dir.path(), IndexConfig::new(2)).unwrap();
1626            assert_eq!(index.retrieve(&q).unwrap().matches.is_empty(), committed);
1627        }
1628        for stage in [
1629            "compact_partial_write",
1630            "compaction_copied",
1631            "manifest_written",
1632            "manifest_renamed",
1633            "directory_synced",
1634        ] {
1635            let dir = tempfile::tempdir().unwrap();
1636            let index = index(dir.path());
1637            let before = index.stats().generation;
1638            let expected =
1639                serde_json::to_value(index.retrieve(&request()).unwrap().matches).unwrap();
1640            FAIL_COMMIT.with(|f| f.set(Some((stage, 1))));
1641            assert!(index.compact().is_err(), "{stage}");
1642            assert_eq!(
1643                index.stats().generation,
1644                before + u64::from(["manifest_renamed", "directory_synced"].contains(&stage))
1645            );
1646            assert_eq!(
1647                serde_json::to_value(index.retrieve(&request()).unwrap().matches).unwrap(),
1648                expected
1649            );
1650            drop(index);
1651            let index = MultiVectorIndex::open(dir.path(), IndexConfig::new(2)).unwrap();
1652            assert_eq!(
1653                serde_json::to_value(index.retrieve(&request()).unwrap().matches).unwrap(),
1654                expected
1655            );
1656        }
1657    }
1658    #[test]
1659    fn sealed_generations_and_named_ann_keep_overlay_and_adaptive_budgets() {
1660        let dir = tempfile::tempdir().unwrap();
1661        let index = index(dir.path());
1662        index.build_dense_ann("semantic", 4, 16).unwrap();
1663        index.seal().unwrap();
1664        index
1665            .upsert_records(vec![document("a", "E123 repair", vec![1., 0.], "a", 0)])
1666            .unwrap();
1667        let mut q = request();
1668        q.filter = None;
1669        q.prefetch.retain(|c| matches!(c, Channel::Dense { .. }));
1670        assert_eq!(
1671            index.retrieve(&q).unwrap().trace.channels[0]["backend"],
1672            "hnsw_dense"
1673        );
1674        assert!(index.delete("b").unwrap());
1675        assert!(
1676            index
1677                .retrieve(&q)
1678                .unwrap()
1679                .matches
1680                .iter()
1681                .all(|h| h.id != "b")
1682        );
1683        assert_eq!(index.stats().storage_segments, 2);
1684        index.compact().unwrap();
1685        assert_eq!(index.stats().storage_segments, 1);
1686        let before = serde_json::to_value(index.retrieve(&q).unwrap().matches).unwrap();
1687        drop(index);
1688        let index = MultiVectorIndex::open(dir.path(), IndexConfig::new(2)).unwrap();
1689        assert_eq!(
1690            serde_json::to_value(index.retrieve(&q).unwrap().matches).unwrap(),
1691            before
1692        );
1693        assert_eq!(
1694            index.retrieve(&q).unwrap().trace.channels[0]["backend"],
1695            "exact_dense"
1696        );
1697        let mut q = request();
1698        q.limit = 1;
1699        q.rerank = Some(Rerank {
1700            field: Some("tokens".into()),
1701            vectors: vec![vec![1., 0.]],
1702            limit: 3,
1703            adaptive: Some(AdaptiveRerank {
1704                min_candidates: 1,
1705                agreement_threshold: 0.5,
1706            }),
1707        });
1708        let response = index.retrieve(&q).unwrap();
1709        assert_eq!(response.trace.reranked_candidates, 1);
1710        assert_eq!(response.trace.channel_agreement, Some(1.));
1711    }
1712    #[test]
1713    fn compaction_reclaims_overwrites_but_pins_active_readers_and_raw_scores() {
1714        let dir = tempfile::tempdir().unwrap();
1715        let index = index(dir.path());
1716        for _ in 0..4 {
1717            index
1718                .upsert_records(vec![document("a", "E123 repair", vec![0.8, 0.6], "a", 0)])
1719                .unwrap();
1720        }
1721        let snapshot = index.snapshot();
1722        let expected = index
1723            .named_scores(
1724                &snapshot,
1725                "semantic",
1726                &[vec![1., 0.]],
1727                false,
1728                &HashSet::from(["a", "b"]),
1729                10,
1730            )
1731            .unwrap();
1732        let report = index.compact().unwrap();
1733        assert!(report["bytes_after"].as_u64().unwrap() < report["bytes_before"].as_u64().unwrap());
1734        assert!(dir.path().join("fde/fde.bin").exists());
1735        assert_eq!(
1736            index
1737                .named_scores(
1738                    &snapshot,
1739                    "semantic",
1740                    &[vec![1., 0.]],
1741                    false,
1742                    &HashSet::from(["a", "b"]),
1743                    10
1744                )
1745                .unwrap(),
1746            expected
1747        );
1748        drop(snapshot);
1749        assert!(!dir.path().join("fde").exists());
1750        drop(index);
1751        let reopened = MultiVectorIndex::open(dir.path(), IndexConfig::new(2)).unwrap();
1752        assert_eq!(
1753            reopened
1754                .named_scores(
1755                    &reopened.snapshot(),
1756                    "semantic",
1757                    &[vec![1., 0.]],
1758                    false,
1759                    &HashSet::from(["a", "b"]),
1760                    10
1761                )
1762                .unwrap(),
1763            expected
1764        );
1765        reopened
1766            .upsert_records(vec![document("new", "E123", vec![1., 0.], "a", 2)])
1767            .unwrap();
1768        assert_eq!(reopened.stats().documents, 4);
1769    }
1770
1771    #[test]
1772    fn global_development_choice_is_the_rrf_default() {
1773        assert!(matches!(Fusion::default(), Fusion::Rrf { k } if k == 10.));
1774    }
1775}