Skip to main content

multivector/
planner.rs

1//! Deterministic physical planning for the retrieval API.
2use super::*;
3
4#[derive(Clone, Copy, Debug, Default, Deserialize, Eq, PartialEq, Serialize)]
5#[serde(rename_all = "snake_case")]
6pub enum QualityPreference {
7    Fast,
8    #[default]
9    Balanced,
10    High,
11}
12
13#[derive(Clone, Debug, Default, Deserialize, PartialEq, Serialize)]
14#[serde(deny_unknown_fields)]
15pub struct RetrievalObjective {
16    pub latency_budget_ms: Option<f64>,
17    pub context_budget_tokens: Option<usize>,
18    #[serde(default)]
19    pub quality: QualityPreference,
20}
21
22#[derive(Clone, Copy, Debug, Eq, PartialEq, Serialize)]
23#[serde(rename_all = "snake_case")]
24pub enum LogicalChannelKind {
25    Bm25,
26    Sparse,
27    Dense,
28    Multivector,
29}
30
31#[derive(Clone, Debug, Eq, PartialEq, Serialize)]
32pub struct LogicalChannel {
33    pub index: usize,
34    pub kind: LogicalChannelKind,
35    #[serde(skip_serializing_if = "Option::is_none")]
36    pub field: Option<String>,
37    pub candidate_limit: usize,
38}
39
40#[derive(Clone, Copy, Debug, Eq, PartialEq, Serialize)]
41#[serde(rename_all = "snake_case")]
42pub enum LogicalFusion {
43    ReciprocalRank,
44    Weighted,
45}
46
47#[derive(Clone, Debug, PartialEq, Serialize)]
48pub struct LogicalPlan {
49    pub objective: RetrievalObjective,
50    pub channels: Vec<LogicalChannel>,
51    pub fusion: LogicalFusion,
52    pub filtered: bool,
53    pub rerank: bool,
54    pub context_selection: bool,
55}
56
57#[derive(Clone, Copy, Debug, Deserialize, Eq, Hash, PartialEq, Serialize)]
58#[serde(rename_all = "snake_case")]
59pub enum PhysicalOperator {
60    Bm25,
61    SparseDot,
62    ExactDense,
63    HnswDense,
64    ExactMaxsim,
65    ExactFde,
66    HnswFde,
67}
68
69impl PhysicalOperator {
70    pub fn as_str(self) -> &'static str {
71        match self {
72            Self::Bm25 => "bm25",
73            Self::SparseDot => "sparse_dot",
74            Self::ExactDense => "exact_dense",
75            Self::HnswDense => "hnsw_dense",
76            Self::ExactMaxsim => "exact_maxsim",
77            Self::ExactFde => "exact_fde",
78            Self::HnswFde => "hnsw_fde",
79        }
80    }
81}
82
83#[derive(Clone, Copy, Debug, Eq, PartialEq, Serialize)]
84#[serde(rename_all = "snake_case")]
85pub enum PlanReason {
86    /// Index over a lexical/BM25 posting list.
87    LexicalIndex,
88    /// Index over a sparse inverted index.
89    SparseIndex,
90    /// Caller explicitly requested exact execution.
91    RequestedExact,
92    /// ANN graph available and selected by availability heuristic.
93    AnnReady,
94    /// ANN graph not built; falling back to exact scan.
95    AnnUnavailable,
96    /// Filter present; selectivity too low or post-filter disabled.
97    FilterRequiresExact,
98    /// Operator does not support ANN (e.g. named multivector exact-only).
99    OperatorRequiresExact,
100    /// Cost model estimated exact cheaper than HNSW for this corpus size.
101    LowerEstimatedCost,
102    /// Runtime calibration estimated lower p90 latency for this hardware.
103    LowerCalibratedLatency,
104}
105
106#[derive(Clone, Copy, Debug, Eq, PartialEq, Serialize)]
107#[serde(rename_all = "snake_case")]
108pub enum FilterStrategy {
109    None,
110    MetadataIndex,
111    MetadataScan,
112}
113
114#[derive(Clone, Debug, PartialEq, Serialize)]
115pub struct FilterStats {
116    /// Fraction of the current generation that passed the predicate.
117    pub selectivity: f32,
118    /// Physical strategy used to evaluate the filter.
119    pub filter_operator: FilterStrategy,
120}
121
122#[derive(Clone, Copy, Debug, Eq, PartialEq, Serialize)]
123#[serde(rename_all = "snake_case")]
124pub enum RepresentationKind {
125    Dense,
126    Multivector,
127    Sparse,
128}
129
130#[derive(Clone, Debug, Eq, PartialEq, Serialize)]
131pub struct FieldStats {
132    pub kind: RepresentationKind,
133    #[serde(skip_serializing_if = "Option::is_none")]
134    pub dimension: Option<usize>,
135    pub documents: usize,
136    /// Documents with this field inside the active filter, when filtered.
137    #[serde(skip_serializing_if = "Option::is_none")]
138    pub eligible_documents: Option<usize>,
139    pub graph_ready: bool,
140}
141
142#[derive(Clone, Debug, PartialEq, Serialize)]
143pub struct PlannerStats {
144    pub generation: u64,
145    pub documents: usize,
146    pub text_documents: usize,
147    pub token_documents: usize,
148    pub token_vectors: usize,
149    pub fde_dimension: usize,
150    pub fde_graph_ready: bool,
151    pub fields: BTreeMap<String, FieldStats>,
152    /// Populated when a filter is present; None for unfiltered plans.
153    #[serde(skip_serializing_if = "Option::is_none")]
154    pub filter_stats: Option<FilterStats>,
155}
156
157/// Collection-wide statistics maintained with each immutable state generation.
158/// Query planning must not scan the document map to rediscover these values.
159#[derive(Clone, Debug, Default)]
160pub(super) struct CachedPlannerStats {
161    text_documents: usize,
162    token_documents: usize,
163    token_vectors: usize,
164    field_documents: BTreeMap<String, usize>,
165}
166
167impl CachedPlannerStats {
168    pub(super) fn from_documents(
169        documents: &HashMap<String, DocumentRecord>,
170        schema: &BTreeMap<String, FieldSchema>,
171    ) -> Self {
172        let mut stats = Self {
173            field_documents: schema.keys().map(|name| (name.clone(), 0)).collect(),
174            ..Self::default()
175        };
176        for document in documents.values() {
177            stats.add(document, schema);
178        }
179        stats
180    }
181
182    pub(super) fn add(
183        &mut self,
184        document: &DocumentRecord,
185        schema: &BTreeMap<String, FieldSchema>,
186    ) {
187        self.text_documents += usize::from(document.fields.has_text());
188        self.token_documents += usize::from(document.tokens > 0);
189        self.token_vectors += document.tokens;
190        for name in schema.keys() {
191            let count = self.field_documents.entry(name.clone()).or_default();
192            *count += usize::from(document.fields.has_representation(name));
193        }
194    }
195
196    pub(super) fn remove(
197        &mut self,
198        document: &DocumentRecord,
199        schema: &BTreeMap<String, FieldSchema>,
200    ) {
201        self.text_documents = self
202            .text_documents
203            .checked_sub(usize::from(document.fields.has_text()))
204            .expect("cached text-document count is consistent");
205        self.token_documents = self
206            .token_documents
207            .checked_sub(usize::from(document.tokens > 0))
208            .expect("cached token-document count is consistent");
209        self.token_vectors = self
210            .token_vectors
211            .checked_sub(document.tokens)
212            .expect("cached token-vector count is consistent");
213        for name in schema.keys() {
214            let decrement = usize::from(document.fields.has_representation(name));
215            let count = self.field_documents.entry(name.clone()).or_default();
216            *count = count
217                .checked_sub(decrement)
218                .expect("cached field-document count is consistent");
219        }
220    }
221}
222
223#[derive(Clone, Debug, PartialEq, Serialize)]
224pub struct PlannedChannel {
225    pub index: usize,
226    pub operator: PhysicalOperator,
227    pub reason: PlanReason,
228    pub limit: usize,
229    #[serde(skip_serializing_if = "Option::is_none")]
230    pub ef_search: Option<usize>,
231    /// Hardware-independent work estimate.
232    pub estimated_cost_units: f64,
233    /// Calibrated p90 estimate after enough observations in this workload bucket.
234    #[serde(skip_serializing_if = "Option::is_none")]
235    pub estimated_latency_ms: Option<f64>,
236}
237
238#[derive(Clone, Copy, Debug, Deserialize, Eq, Hash, PartialEq, Serialize)]
239#[serde(rename_all = "snake_case")]
240pub enum FusionOperator {
241    NativeScore,
242    ReciprocalRank,
243    WeightedScore,
244}
245
246#[derive(Clone, Debug, PartialEq, Serialize)]
247pub struct RerankPlan {
248    pub operator: PhysicalOperator,
249    pub candidate_limit: usize,
250    pub adaptive: bool,
251}
252
253#[derive(Clone, Copy, Debug, Deserialize, Eq, Hash, PartialEq, Serialize)]
254#[serde(rename_all = "snake_case")]
255pub enum ContextOperator {
256    Ranked,
257    Mmr,
258}
259
260#[derive(Clone, Debug, PartialEq, Serialize)]
261pub struct ContextPlan {
262    pub operator: ContextOperator,
263    pub result_limit: usize,
264    pub context_budget_tokens: Option<usize>,
265}
266
267#[derive(Clone, Debug, PartialEq, Serialize)]
268#[serde(tag = "kind", content = "detail", rename_all = "snake_case")]
269pub enum PlanStage {
270    Parallel(Vec<PlannedChannel>),
271    Fusion(FusionOperator),
272    Rerank(RerankPlan),
273    Context(ContextPlan),
274}
275
276/// Relative work estimate. These units compare plans; they are not latency.
277#[derive(Clone, Debug, PartialEq, Serialize)]
278pub struct PlanEstimate {
279    pub critical_path_cost: f64,
280    pub total_cost: f64,
281    /// Calibrated p90 for the parallel candidate stage, excluding later stages.
282    #[serde(skip_serializing_if = "Option::is_none")]
283    pub calibrated_parallel_p90_ms: Option<f64>,
284    /// Sum of calibrated channel p90 values; useful as a CPU-work proxy.
285    #[serde(skip_serializing_if = "Option::is_none")]
286    pub calibrated_channel_total_p90_ms: Option<f64>,
287    #[serde(skip_serializing_if = "Option::is_none")]
288    pub calibrated_fusion_p90_ms: Option<f64>,
289    #[serde(skip_serializing_if = "Option::is_none")]
290    pub calibrated_rerank_p90_ms: Option<f64>,
291    #[serde(skip_serializing_if = "Option::is_none")]
292    pub calibrated_context_p90_ms: Option<f64>,
293    /// End-to-end estimate for the stages in this plan. This is available only
294    /// after every selected stage has enough observations in its workload bucket.
295    #[serde(skip_serializing_if = "Option::is_none")]
296    pub calibrated_total_p90_ms: Option<f64>,
297}
298
299#[derive(Clone, Debug, PartialEq, Serialize)]
300pub struct RetrievalPlan {
301    pub logical: LogicalPlan,
302    pub stats: PlannerStats,
303    pub eligible_documents: usize,
304    pub filter: FilterStrategy,
305    pub stages: Vec<PlanStage>,
306    pub estimate: PlanEstimate,
307    /// Set when planning_mode != Manual; contains the auto-generated channel list.
308    #[serde(skip_serializing_if = "Option::is_none")]
309    pub policy: Option<PolicyPlan>,
310}
311
312impl RetrievalPlan {
313    /// Returns the `PlannedChannel` slice from the first `Parallel` stage.
314    pub fn parallel_channels(&self) -> &[PlannedChannel] {
315        self.stages
316            .iter()
317            .find_map(|s| match s {
318                PlanStage::Parallel(ch) => Some(ch.as_slice()),
319                _ => None,
320            })
321            .unwrap_or(&[])
322    }
323
324    /// Returns the fusion operator from the plan, if present.
325    pub fn fusion_operator(&self) -> Option<FusionOperator> {
326        self.stages.iter().find_map(|s| match s {
327            PlanStage::Fusion(op) => Some(*op),
328            _ => None,
329        })
330    }
331
332    /// Returns the rerank plan, if present.
333    pub fn rerank_plan(&self) -> Option<&RerankPlan> {
334        self.stages.iter().find_map(|s| match s {
335            PlanStage::Rerank(r) => Some(r),
336            _ => None,
337        })
338    }
339
340    /// Returns the context plan.
341    pub fn context_plan(&self) -> Option<&ContextPlan> {
342        self.stages.iter().find_map(|s| match s {
343            PlanStage::Context(c) => Some(c),
344            _ => None,
345        })
346    }
347}
348
349fn invalid(message: impl Into<String>) -> IndexError {
350    IndexError::Invalid(message.into())
351}
352
353pub(super) fn validate_request(request: &RetrieveRequest) -> Result<(), IndexError> {
354    // In auto modes, prefetch is generated by the policy planner — allow empty.
355    let prefetch_empty = request.prefetch.is_empty();
356    let is_auto = request.planning_mode != PlanningMode::Manual;
357    if is_auto && request.query.is_none() {
358        return Err(invalid(
359            "planning_mode auto requires a query field with representations",
360        ));
361    }
362    if prefetch_empty && !is_auto
363        || request.prefetch.len() > 8
364        || request.limit == 0
365        || request.limit > 10_000
366        || request.context.neighbors > 8
367        || request.context.per_parent == Some(0)
368        || request
369            .context
370            .mmr
371            .is_some_and(|value| !value.is_finite() || !(0.0..=1.0).contains(&value))
372    {
373        return Err(invalid("invalid retrieval or context budgets"));
374    }
375    if request
376        .objective
377        .latency_budget_ms
378        .is_some_and(|budget| !budget.is_finite() || budget <= 0.0)
379    {
380        return Err(invalid("latency_budget_ms must be finite and positive"));
381    }
382    if request.objective.context_budget_tokens == Some(0) {
383        return Err(invalid("invalid retrieval objective"));
384    }
385    if let Some(filter) = &request.filter {
386        filter.validate(0)?;
387    }
388    Ok(())
389}
390
391impl MultiVectorIndex {
392    pub(super) fn prepare_request<'a>(
393        &self,
394        state: &State,
395        request: &'a RetrieveRequest,
396        eligible_documents: Option<(usize, FilterStrategy)>,
397        eligible_set: Option<&DocSet>,
398    ) -> (std::borrow::Cow<'a, RetrieveRequest>, Option<PolicyPlan>) {
399        if request.planning_mode == PlanningMode::Manual {
400            return (std::borrow::Cow::Borrowed(request), None);
401        }
402        let stats = planner_stats(
403            state,
404            self.fde.output_dimension(),
405            eligible_documents,
406            eligible_set,
407        );
408        let mut policy = policy::generate_policy_prefetch(
409            request.query.as_ref().expect("auto query validated"),
410            &stats,
411            request.limit,
412            request.objective.quality,
413            state.retrieval.schema(),
414        );
415        if request.planning_mode == PlanningMode::AutoWithOverrides {
416            policy::apply_overrides(&mut policy, &request.prefetch);
417        }
418        let mut effective = request.clone();
419        effective.prefetch = policy.generated_prefetch.clone();
420        (std::borrow::Cow::Owned(effective), Some(policy))
421    }
422
423    /// Compile a request without scoring documents.
424    pub fn plan(&self, request: &RetrieveRequest) -> Result<RetrievalPlan, IndexError> {
425        validate_request(request)?;
426        let state = self.snapshot();
427        let (eligible_set, filter_strategy) =
428            request
429                .filter
430                .as_ref()
431                .map_or((None, FilterStrategy::None), |filter| {
432                    state.retrieval.indexed_filter_set(filter).map_or_else(
433                        || {
434                            (
435                                Some(DocSet::from_numbers(
436                                    state.retrieval.next_id(),
437                                    state
438                                        .documents
439                                        .iter()
440                                        .filter(|(_, document)| filter.matches(&document.metadata))
441                                        .filter_map(|(id, _)| state.retrieval.number(id)),
442                                )),
443                                FilterStrategy::MetadataScan,
444                            )
445                        },
446                        |set| (Some(set), FilterStrategy::MetadataIndex),
447                    )
448                });
449        let eligible = eligible_set
450            .as_ref()
451            .map_or(state.documents.len(), DocSet::len);
452        let planner_filter = request.filter.as_ref().map(|_| (eligible, filter_strategy));
453        let (effective_request, policy_plan) =
454            self.prepare_request(&state, request, planner_filter, eligible_set.as_ref());
455        let request = effective_request.as_ref();
456        let mut plan = self.compile_plan(
457            &state,
458            request,
459            eligible,
460            filter_strategy,
461            eligible_set.as_ref(),
462        )?;
463        plan.policy = policy_plan;
464        Ok(plan)
465    }
466
467    pub(super) fn compile_plan(
468        &self,
469        state: &State,
470        request: &RetrieveRequest,
471        eligible_documents: usize,
472        filter_strategy: FilterStrategy,
473        eligible_set: Option<&DocSet>,
474    ) -> Result<RetrievalPlan, IndexError> {
475        validate_request(request)?;
476        if request.prefetch.is_empty() {
477            return Err(invalid(
478                "auto planning found no usable query representation",
479            ));
480        }
481
482        let stats = planner_stats(
483            state,
484            self.fde.output_dimension(),
485            request
486                .filter
487                .as_ref()
488                .map(|_| (eligible_documents, filter_strategy)),
489            eligible_set,
490        );
491        let calibration = self.calibration_snapshot();
492
493        let logical_channels = request
494            .prefetch
495            .iter()
496            .enumerate()
497            .map(|(index, channel)| {
498                let (kind, field, candidate_limit) = match channel {
499                    Channel::Bm25 { limit, .. } => (LogicalChannelKind::Bm25, None, *limit),
500                    Channel::Sparse { field, limit, .. } => {
501                        (LogicalChannelKind::Sparse, Some(field.clone()), *limit)
502                    }
503                    Channel::Dense { field, limit, .. } => {
504                        (LogicalChannelKind::Dense, Some(field.clone()), *limit)
505                    }
506                    Channel::Multivector { field, limit, .. } => {
507                        (LogicalChannelKind::Multivector, field.clone(), *limit)
508                    }
509                };
510                LogicalChannel {
511                    index,
512                    kind,
513                    field,
514                    candidate_limit,
515                }
516            })
517            .collect();
518
519        let logical = LogicalPlan {
520            objective: request.objective.clone(),
521            channels: logical_channels,
522            fusion: match request.fusion {
523                Fusion::Rrf { .. } => LogicalFusion::ReciprocalRank,
524                Fusion::Weighted { .. } => LogicalFusion::Weighted,
525            },
526            filtered: request.filter.is_some(),
527            rerank: request.rerank.is_some(),
528            context_selection: request.context != ContextOptions::default(),
529        };
530
531        let mut planned_channels = Vec::with_capacity(request.prefetch.len());
532        for (index, channel) in request.prefetch.iter().enumerate() {
533            let (operator, reason, limit, ef_search, cost) = match channel {
534                Channel::Bm25 { text, limit, k1, b } => {
535                    if text.len() > 65_536
536                        || !k1.is_finite()
537                        || *k1 < 0.0
538                        || !b.is_finite()
539                        || !(0.0..=1.0).contains(b)
540                    {
541                        return Err(invalid("invalid BM25 query or parameters"));
542                    }
543                    let cost = stats.text_documents.max(1) as f64 * 10.0;
544                    (
545                        PhysicalOperator::Bm25,
546                        PlanReason::LexicalIndex,
547                        *limit,
548                        None,
549                        cost,
550                    )
551                }
552                Channel::Sparse {
553                    field,
554                    vector,
555                    limit,
556                } => {
557                    if !state.retrieval.has_sparse_field(field) {
558                        return Err(invalid(format!("unknown sparse field {field:?}")));
559                    }
560                    vector
561                        .canonicalized()
562                        .map_err(|error| invalid(error.to_string()))?;
563                    let cost = stats
564                        .fields
565                        .get(field)
566                        .map_or(1, |field| field.documents.max(1))
567                        as f64
568                        * 5.0;
569                    (
570                        PhysicalOperator::SparseDot,
571                        PlanReason::SparseIndex,
572                        *limit,
573                        None,
574                        cost,
575                    )
576                }
577                Channel::Dense {
578                    field,
579                    vector,
580                    limit,
581                    backend,
582                    ef_search,
583                } => {
584                    retrieval::validate_matrix(std::slice::from_ref(vector))?;
585                    if state.retrieval.dense_dimension(field)? != vector.len()
586                        || !["auto", "exact", "hnsw"].contains(&backend.as_str())
587                        || *ef_search == 0
588                        || *ef_search > 65_536
589                    {
590                        return Err(invalid("invalid dense query dimension or backend"));
591                    }
592                    let ann_ready = state.named_ann.contains_key(field);
593                    if backend == "hnsw" && !ann_ready {
594                        return Err(invalid("dense ANN not built"));
595                    }
596                    let dim = state.retrieval.dense_dimension(field).unwrap_or(1) as f64;
597                    let ef = *ef_search as f64;
598                    let n = stats.documents.max(1) as f64;
599                    let f = eligible_documents.max(1) as f64;
600                    let filtered = request.filter.is_some();
601                    let filter_supported = request
602                        .filter
603                        .as_ref()
604                        .is_none_or(|filter| filter.ann_filter().is_some());
605                    if backend == "hnsw" && !filter_supported {
606                        return Err(invalid("filter is not supported by dense HNSW"));
607                    }
608                    // EXACT_SCALE: benchmarks on a 10k×128 fixture show exact
609                    // is ~26× slower than HNSW unfiltered, but the raw op counts
610                    // imply only ~2×. The 13× multiplier closes that gap so the
611                    // model's crossover (~4% selectivity) matches empirical data.
612                    // In-place filtered HNSW (annex-core) does not degrade with
613                    // selectivity the way post-filter HNSW does, so no selectivity
614                    // penalty is applied to cost_hnsw.
615                    let cost_exact = f * dim * 2.0 * 13.0;
616                    let cost_hnsw = ef * n.log2().ceil().max(1.0) * dim * 3.0;
617                    let _ = filtered; // crossover via calibrated cost_exact
618                    let exact_ms = calibrated_latency(
619                        &calibration,
620                        PhysicalOperator::ExactDense,
621                        channel,
622                        &stats,
623                        cost_exact,
624                    );
625                    let hnsw_ms = calibrated_latency(
626                        &calibration,
627                        PhysicalOperator::HnswDense,
628                        channel,
629                        &stats,
630                        cost_hnsw,
631                    );
632                    let (operator, reason) = choose_ann_with_cost(
633                        backend,
634                        ann_ready,
635                        filter_supported,
636                        PhysicalOperator::ExactDense,
637                        PhysicalOperator::HnswDense,
638                        cost_exact,
639                        cost_hnsw,
640                        exact_ms,
641                        hnsw_ms,
642                    );
643                    let cost = if operator == PhysicalOperator::HnswDense {
644                        cost_hnsw
645                    } else {
646                        cost_exact
647                    };
648                    (operator, reason, *limit, Some(*ef_search), cost)
649                }
650                Channel::Multivector {
651                    field: Some(field),
652                    vectors,
653                    limit,
654                    backend,
655                    ..
656                } => {
657                    retrieval::validate_matrix(vectors)?;
658                    let expected = FieldSchema::Multivector {
659                        dimension: vectors[0].len(),
660                    };
661                    if state.retrieval.schema().get(field) != Some(&expected) {
662                        return Err(invalid(format!(
663                            "unknown field or query dimension/kind mismatch: {field:?}"
664                        )));
665                    }
666                    if backend != "auto" && backend != "exact" {
667                        return Err(invalid("named multivectors support exact or auto"));
668                    }
669                    let reason = if backend == "exact" {
670                        PlanReason::RequestedExact
671                    } else {
672                        PlanReason::OperatorRequiresExact
673                    };
674                    let cost = estimate_channel_cost(
675                        PhysicalOperator::ExactMaxsim,
676                        eligible_documents,
677                        stats.documents,
678                        None,
679                        &stats,
680                        channel,
681                    );
682                    (PhysicalOperator::ExactMaxsim, reason, *limit, None, cost)
683                }
684                Channel::Multivector {
685                    field: None,
686                    vectors,
687                    limit,
688                    backend,
689                    ef_search,
690                } => {
691                    self.validate(vectors)?;
692                    if vectors.len() > 1024
693                        || *ef_search == 0
694                        || *ef_search > 65_536
695                        || !["auto", "exact", "hnsw"].contains(&backend.as_str())
696                    {
697                        return Err(invalid("invalid multivector backend or budget"));
698                    }
699                    let ann_ready = state.fde_ann.is_some();
700                    if backend == "hnsw" && !ann_ready {
701                        return Err(invalid("FDE ANN not built"));
702                    }
703                    let fde_dim = stats.fde_dimension.max(1) as f64;
704                    let ef = *ef_search as f64;
705                    let n = stats.documents.max(1) as f64;
706                    let f = eligible_documents.max(1) as f64;
707                    let filtered = request.filter.is_some();
708                    let filter_supported = request
709                        .filter
710                        .as_ref()
711                        .is_none_or(|filter| filter.ann_filter().is_some());
712                    if backend == "hnsw" && !filter_supported {
713                        return Err(invalid("filter is not supported by FDE HNSW"));
714                    }
715                    let cost_exact = f * fde_dim * 2.0 * 13.0;
716                    let cost_hnsw = ef * n.log2().ceil().max(1.0) * fde_dim * 3.0;
717                    let _ = filtered;
718                    let exact_ms = calibrated_latency(
719                        &calibration,
720                        PhysicalOperator::ExactFde,
721                        channel,
722                        &stats,
723                        cost_exact,
724                    );
725                    let hnsw_ms = calibrated_latency(
726                        &calibration,
727                        PhysicalOperator::HnswFde,
728                        channel,
729                        &stats,
730                        cost_hnsw,
731                    );
732                    let (operator, reason) = choose_ann_with_cost(
733                        backend,
734                        ann_ready,
735                        filter_supported,
736                        PhysicalOperator::ExactFde,
737                        PhysicalOperator::HnswFde,
738                        cost_exact,
739                        cost_hnsw,
740                        exact_ms,
741                        hnsw_ms,
742                    );
743                    let cost = if operator == PhysicalOperator::HnswFde {
744                        cost_hnsw
745                    } else {
746                        cost_exact
747                    };
748                    (operator, reason, *limit, Some(*ef_search), cost)
749                }
750            };
751            if limit == 0 || limit > 100_000 {
752                return Err(invalid("channel limit must be in 1..=100000"));
753            }
754            let estimated_latency_ms =
755                calibrated_latency(&calibration, operator, channel, &stats, cost);
756            planned_channels.push(PlannedChannel {
757                index,
758                operator,
759                reason,
760                limit,
761                ef_search,
762                estimated_cost_units: cost,
763                estimated_latency_ms,
764            });
765        }
766
767        match &request.fusion {
768            Fusion::Rrf { k } if !k.is_finite() || *k < 0.0 => {
769                return Err(invalid("RRF k must be finite and nonnegative"));
770            }
771            Fusion::Weighted { weights }
772                if weights.len() != planned_channels.len()
773                    || weights
774                        .iter()
775                        .any(|weight| !weight.is_finite() || *weight < 0.0)
776                    || weights.iter().all(|weight| *weight == 0.0) =>
777            {
778                return Err(invalid(
779                    "weighted fusion requires one nonnegative finite weight per channel",
780                ));
781            }
782            _ => {}
783        }
784
785        if let Some(rerank) = &request.rerank {
786            if rerank.limit < request.limit || rerank.limit > 100_000 {
787                return Err(invalid(
788                    "rerank limit must cover result limit and be <=100000",
789                ));
790            }
791            if let Some(policy) = &rerank.adaptive
792                && (policy.min_candidates < request.limit
793                    || policy.min_candidates > rerank.limit
794                    || !policy.agreement_threshold.is_finite()
795                    || !(0.0..=1.0).contains(&policy.agreement_threshold)
796                    || planned_channels.len() < 2)
797            {
798                return Err(invalid(
799                    "invalid adaptive rerank policy; needs multiple channels",
800                ));
801            }
802            retrieval::validate_matrix(&rerank.vectors)?;
803            if let Some(field) = &rerank.field {
804                let expected = FieldSchema::Multivector {
805                    dimension: rerank.vectors[0].len(),
806                };
807                if state.retrieval.schema().get(field) != Some(&expected) {
808                    return Err(invalid(format!(
809                        "unknown field or query dimension/kind mismatch: {field:?}"
810                    )));
811                }
812            } else {
813                self.validate(&rerank.vectors)?;
814            }
815        }
816
817        match (
818            request.context.mmr,
819            request.context.diversity_field.as_deref(),
820        ) {
821            (Some(_), Some(field)) => {
822                state.retrieval.dense_dimension(field)?;
823            }
824            (Some(_), None) => return Err(invalid("MMR requires a dense diversity_field")),
825            (None, Some(_)) => return Err(invalid("diversity_field requires mmr")),
826            (None, None) => {}
827        }
828
829        let fusion_op = match request.fusion {
830            Fusion::Rrf { .. } if planned_channels.len() == 1 => FusionOperator::NativeScore,
831            Fusion::Rrf { .. } => FusionOperator::ReciprocalRank,
832            Fusion::Weighted { .. } => FusionOperator::WeightedScore,
833        };
834
835        let rerank_plan = request.rerank.as_ref().map(|rerank| RerankPlan {
836            operator: PhysicalOperator::ExactMaxsim,
837            candidate_limit: rerank.limit,
838            adaptive: rerank.adaptive.is_some(),
839        });
840
841        let context_plan = ContextPlan {
842            operator: if request.context.mmr.is_some() {
843                ContextOperator::Mmr
844            } else {
845                ContextOperator::Ranked
846            },
847            result_limit: request.limit,
848            context_budget_tokens: request.objective.context_budget_tokens,
849        };
850
851        let channel_costs = planned_channels
852            .iter()
853            .map(|channel| channel.estimated_cost_units)
854            .collect::<Vec<_>>();
855        let parallel_critical_path = channel_costs
856            .iter()
857            .copied()
858            .fold(0.0_f64, f64::max)
859            .max(1.0);
860        let parallel_total = channel_costs.iter().sum::<f64>().max(1.0);
861        let calibrated_channel_times = planned_channels
862            .iter()
863            .map(|channel| channel.estimated_latency_ms)
864            .collect::<Option<Vec<_>>>();
865        let calibrated_parallel_p90_ms = calibrated_channel_times
866            .as_ref()
867            .map(|times| times.iter().copied().fold(0.0_f64, f64::max));
868        let calibrated_channel_total_p90_ms = calibrated_channel_times
869            .as_ref()
870            .map(|times| times.iter().sum());
871        let fusion_cost =
872            fusion_cost_units(planned_channels.iter().map(|channel| channel.limit).sum());
873        let rerank_cost = rerank_plan.as_ref().map_or(0.0, |r| {
874            rerank_cost_units(r.candidate_limit, &stats, self.config.dimension)
875        });
876        let context_pool = rerank_plan.as_ref().map_or_else(
877            || planned_channels.iter().map(|channel| channel.limit).sum(),
878            |rerank| rerank.candidate_limit,
879        );
880        let context_cost = context_cost_units(
881            context_plan.operator,
882            context_pool,
883            request.limit,
884            self.config.dimension,
885        );
886        let selectivity = stats.filter_stats.as_ref().map(|filter| filter.selectivity);
887        let fusion_key = calibration::key(
888            CalibrationTarget::Fusion(fusion_op),
889            0,
890            stats.documents,
891            selectivity,
892        );
893        let rerank_key = rerank_plan.as_ref().map(|_| {
894            calibration::key(
895                CalibrationTarget::RerankMaxsim,
896                self.config.dimension,
897                stats.documents,
898                selectivity,
899            )
900        });
901        let context_key = calibration::key(
902            CalibrationTarget::Context(context_plan.operator),
903            if context_plan.operator == ContextOperator::Mmr {
904                self.config.dimension
905            } else {
906                0
907            },
908            stats.documents,
909            selectivity,
910        );
911        let calibrated_fusion_p90_ms = calibration.estimate_ms(fusion_key, fusion_cost);
912        let calibrated_rerank_p90_ms =
913            rerank_key.and_then(|key| calibration.estimate_ms(key, rerank_cost));
914        let calibrated_rerank_stage_p90_ms = if rerank_plan.is_some() {
915            calibrated_rerank_p90_ms
916        } else {
917            Some(0.0)
918        };
919        let calibrated_context_p90_ms = calibration.estimate_ms(context_key, context_cost);
920        let calibrated_total_p90_ms = calibrated_parallel_p90_ms
921            .zip(calibrated_fusion_p90_ms)
922            .zip(calibrated_rerank_stage_p90_ms)
923            .zip(calibrated_context_p90_ms)
924            .map(|(((parallel, fusion), rerank), context)| parallel + fusion + rerank + context);
925
926        if let Some(budget) = request.objective.latency_budget_ms {
927            let estimate = calibrated_total_p90_ms.ok_or_else(|| {
928                invalid(
929                    "latency_budget_ms requires at least five observations for every selected stage",
930                )
931            })?;
932            if estimate > budget {
933                return Err(invalid(format!(
934                    "estimated p90 latency {estimate:.3} ms exceeds latency budget {budget:.3} ms"
935                )));
936            }
937        }
938
939        let mut stages = Vec::with_capacity(4);
940        stages.push(PlanStage::Parallel(planned_channels));
941        stages.push(PlanStage::Fusion(fusion_op));
942        if let Some(rp) = rerank_plan.clone() {
943            stages.push(PlanStage::Rerank(rp));
944        }
945        stages.push(PlanStage::Context(context_plan));
946
947        Ok(RetrievalPlan {
948            logical,
949            stats,
950            eligible_documents,
951            filter: filter_strategy,
952            stages,
953            estimate: PlanEstimate {
954                critical_path_cost: parallel_critical_path
955                    + fusion_cost
956                    + rerank_cost
957                    + context_cost,
958                total_cost: parallel_total + fusion_cost + rerank_cost + context_cost,
959                calibrated_parallel_p90_ms,
960                calibrated_channel_total_p90_ms,
961                calibrated_fusion_p90_ms,
962                calibrated_rerank_p90_ms,
963                calibrated_context_p90_ms,
964                calibrated_total_p90_ms,
965            },
966            policy: None, // populated by plan()/retrieve() for non-Manual modes
967        })
968    }
969}
970
971pub(super) fn planner_stats(
972    state: &State,
973    fde_dimension: usize,
974    eligible_documents: Option<(usize, FilterStrategy)>,
975    eligible_set: Option<&DocSet>,
976) -> PlannerStats {
977    let fields = state
978        .retrieval
979        .schema()
980        .iter()
981        .map(|(name, schema)| {
982            let (kind, dimension) = match schema {
983                FieldSchema::Dense { dimension } => (RepresentationKind::Dense, Some(*dimension)),
984                FieldSchema::Multivector { dimension } => {
985                    (RepresentationKind::Multivector, Some(*dimension))
986                }
987                FieldSchema::Sparse => (RepresentationKind::Sparse, None),
988            };
989            let documents = state
990                .planner_stats
991                .field_documents
992                .get(name)
993                .copied()
994                .unwrap_or(0);
995            (
996                name.clone(),
997                FieldStats {
998                    kind,
999                    dimension,
1000                    documents,
1001                    eligible_documents: eligible_set
1002                        .map(|eligible| state.retrieval.field_count_in(name, eligible)),
1003                    graph_ready: state.named_ann.contains_key(name),
1004                },
1005            )
1006        })
1007        .collect();
1008    let n = state.documents.len();
1009    let filter_stats = eligible_documents.map(|(eligible, filter_operator)| {
1010        let selectivity = if n == 0 {
1011            1.0_f32
1012        } else {
1013            (eligible as f32) / (n as f32)
1014        };
1015        FilterStats {
1016            selectivity: selectivity.clamp(0.0, 1.0),
1017            filter_operator,
1018        }
1019    });
1020    PlannerStats {
1021        generation: state.generation,
1022        documents: n,
1023        text_documents: state.planner_stats.text_documents,
1024        token_documents: state.planner_stats.token_documents,
1025        token_vectors: state.planner_stats.token_vectors,
1026        fde_dimension,
1027        fde_graph_ready: state.fde_ann.is_some(),
1028        fields,
1029        filter_stats,
1030    }
1031}
1032
1033/// Choose between exact and ANN operator using the cost model.
1034/// `cost_exact` and `cost_hnsw` are in cost units (use `estimate_channel_cost`
1035/// or the inline formulas in `compile_plan`).
1036fn choose_ann_with_cost(
1037    backend: &str,
1038    ann_ready: bool,
1039    ann_supported: bool,
1040    exact: PhysicalOperator,
1041    ann: PhysicalOperator,
1042    cost_exact: f64,
1043    cost_hnsw: f64,
1044    exact_ms: Option<f64>,
1045    hnsw_ms: Option<f64>,
1046) -> (PhysicalOperator, PlanReason) {
1047    if backend == "exact" {
1048        return (exact, PlanReason::RequestedExact);
1049    }
1050    if !ann_ready {
1051        return (exact, PlanReason::AnnUnavailable);
1052    }
1053    if !ann_supported {
1054        return (exact, PlanReason::FilterRequiresExact);
1055    }
1056    if backend == "hnsw" {
1057        return (ann, PlanReason::AnnReady);
1058    }
1059    if let (Some(exact_ms), Some(hnsw_ms)) = (exact_ms, hnsw_ms) {
1060        return if exact_ms <= hnsw_ms {
1061            (exact, PlanReason::LowerCalibratedLatency)
1062        } else {
1063            (ann, PlanReason::LowerCalibratedLatency)
1064        };
1065    }
1066    // ANN is available — choose by cost model.
1067    if cost_exact <= cost_hnsw {
1068        (exact, PlanReason::LowerEstimatedCost)
1069    } else {
1070        (ann, PlanReason::LowerEstimatedCost)
1071    }
1072}
1073
1074pub(super) fn calibrated_latency(
1075    calibration: &CalibrationSnapshot,
1076    operator: PhysicalOperator,
1077    channel: &Channel,
1078    stats: &PlannerStats,
1079    cost_units: f64,
1080) -> Option<f64> {
1081    calibration.estimate_ms(calibration_key(operator, channel, stats), cost_units)
1082}
1083
1084pub(super) fn calibration_key(
1085    operator: PhysicalOperator,
1086    channel: &Channel,
1087    stats: &PlannerStats,
1088) -> CalibrationKey {
1089    let dimension = match operator {
1090        PhysicalOperator::ExactDense | PhysicalOperator::HnswDense => channel_dim(channel, stats),
1091        PhysicalOperator::ExactFde | PhysicalOperator::HnswFde => stats.fde_dimension,
1092        PhysicalOperator::ExactMaxsim => match channel {
1093            Channel::Multivector { vectors, .. } => {
1094                vectors.first().map_or(stats.fde_dimension, Vec::len)
1095            }
1096            _ => stats.fde_dimension,
1097        },
1098        PhysicalOperator::Bm25 | PhysicalOperator::SparseDot => 0,
1099    };
1100    calibration::key(
1101        CalibrationTarget::Channel(operator),
1102        dimension,
1103        stats.documents,
1104        stats.filter_stats.as_ref().map(|filter| filter.selectivity),
1105    )
1106}
1107
1108pub(super) fn fusion_cost_units(candidates: usize) -> f64 {
1109    candidates.max(1) as f64
1110}
1111
1112pub(super) fn rerank_cost_units(candidates: usize, stats: &PlannerStats, dimension: usize) -> f64 {
1113    let avg_tokens = (stats.token_vectors as f64 / stats.token_documents.max(1) as f64).max(1.0);
1114    candidates.max(1) as f64 * avg_tokens * dimension.max(1) as f64 * 2.0
1115}
1116
1117pub(super) fn context_cost_units(
1118    operator: ContextOperator,
1119    candidates: usize,
1120    result_limit: usize,
1121    dimension: usize,
1122) -> f64 {
1123    match operator {
1124        ContextOperator::Ranked => candidates.max(1) as f64,
1125        ContextOperator::Mmr => {
1126            candidates.max(1) as f64 * result_limit.max(1) as f64 * dimension.max(1) as f64
1127        }
1128    }
1129}
1130
1131pub(super) fn stage_calibration_key(
1132    target: CalibrationTarget,
1133    dimension: usize,
1134    stats: &PlannerStats,
1135) -> CalibrationKey {
1136    calibration::key(
1137        target,
1138        dimension,
1139        stats.documents,
1140        stats.filter_stats.as_ref().map(|filter| filter.selectivity),
1141    )
1142}
1143
1144fn estimate_channel_cost(
1145    operator: PhysicalOperator,
1146    eligible: usize,
1147    corpus: usize,
1148    ef_search: Option<usize>,
1149    stats: &PlannerStats,
1150    channel: &Channel,
1151) -> f64 {
1152    let ef = ef_search.unwrap_or(256) as f64;
1153    let n = corpus.max(1) as f64;
1154    let f = eligible.max(1) as f64;
1155    let log2_n = n.log2().ceil().max(1.0);
1156
1157    match operator {
1158        PhysicalOperator::Bm25 => stats.text_documents.max(1) as f64 * 10.0,
1159        PhysicalOperator::SparseDot => match channel {
1160            Channel::Sparse { field, .. } => {
1161                stats
1162                    .fields
1163                    .get(field)
1164                    .map_or(1, |field| field.documents.max(1)) as f64
1165                    * 5.0
1166            }
1167            _ => unreachable!("sparse operator requires sparse channel"),
1168        },
1169        PhysicalOperator::ExactDense => {
1170            let dim = channel_dim(channel, stats).max(1) as f64;
1171            f * dim * 2.0 * 13.0
1172        }
1173        PhysicalOperator::HnswDense => {
1174            let dim = channel_dim(channel, stats).max(1) as f64;
1175            ef * log2_n * dim * 3.0
1176        }
1177        PhysicalOperator::ExactFde => {
1178            let fde_dim = stats.fde_dimension.max(1) as f64;
1179            f * fde_dim * 2.0 * 13.0
1180        }
1181        PhysicalOperator::HnswFde => {
1182            let fde_dim = stats.fde_dimension.max(1) as f64;
1183            ef * log2_n * fde_dim * 3.0
1184        }
1185        PhysicalOperator::ExactMaxsim => {
1186            // Approximate: scale FDE cost by token-to-vector ratio
1187            let fde_dim = stats.fde_dimension.max(1) as f64;
1188            let avg_tokens = (stats.token_vectors as f64 / stats.token_documents.max(1) as f64)
1189                .clamp(1.0, 512.0);
1190            f * fde_dim * avg_tokens * 2.0
1191        }
1192    }
1193}
1194
1195fn channel_dim(channel: &Channel, stats: &PlannerStats) -> usize {
1196    match channel {
1197        Channel::Dense { field, .. } => stats
1198            .fields
1199            .get(field)
1200            .and_then(|f| f.dimension)
1201            .unwrap_or(1),
1202        Channel::Multivector {
1203            field: Some(field), ..
1204        } => stats
1205            .fields
1206            .get(field)
1207            .and_then(|f| f.dimension)
1208            .unwrap_or(1),
1209        _ => 1,
1210    }
1211}