1use 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 LexicalIndex,
88 SparseIndex,
90 RequestedExact,
92 AnnReady,
94 AnnUnavailable,
96 FilterRequiresExact,
98 OperatorRequiresExact,
100 LowerEstimatedCost,
102 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 pub selectivity: f32,
118 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 #[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 #[serde(skip_serializing_if = "Option::is_none")]
154 pub filter_stats: Option<FilterStats>,
155}
156
157#[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 pub estimated_cost_units: f64,
233 #[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#[derive(Clone, Debug, PartialEq, Serialize)]
278pub struct PlanEstimate {
279 pub critical_path_cost: f64,
280 pub total_cost: f64,
281 #[serde(skip_serializing_if = "Option::is_none")]
283 pub calibrated_parallel_p90_ms: Option<f64>,
284 #[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 #[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 #[serde(skip_serializing_if = "Option::is_none")]
309 pub policy: Option<PolicyPlan>,
310}
311
312impl RetrievalPlan {
313 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 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 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 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 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 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 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; 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, })
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
1033fn 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 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 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}