1use std::collections::{hash_map::Entry, HashMap};
4
5use uuid::Uuid;
6
7use khive_score::DeterministicScore;
8use khive_storage::types::{TextSearchHit, VectorSearchHit};
9
10use crate::error::RuntimeResult;
11use crate::retrieval::{RankScoreKind, SearchHit, SearchSignals, SearchSource};
12use crate::runtime::KhiveRuntime;
13
14pub use khive_fusion::FusionStrategy;
15
16pub type CandidateStream = Vec<(Uuid, DeterministicScore)>;
20
21pub type RankedHit = (Uuid, DeterministicScore);
23
24#[async_trait::async_trait]
32pub trait FusionExecutor: Send + Sync + 'static {
33 fn rank_score_kind(&self) -> RankScoreKind;
35
36 async fn fuse(
41 &self,
42 streams: Vec<CandidateStream>,
43 params: &serde_json::Value,
44 limit: usize,
45 ) -> RuntimeResult<Vec<RankedHit>>;
46}
47
48pub(crate) async fn rrf_fuse_k(
50 rt: &KhiveRuntime,
51 text_hits: Vec<TextSearchHit>,
52 vector_hits: Vec<VectorSearchHit>,
53 k: usize,
54 limit: usize,
55) -> RuntimeResult<Vec<SearchHit>> {
56 rt.fuse_with_strategy(text_hits, vector_hits, &FusionStrategy::Rrf { k }, limit)
57 .await
58}
59
60impl KhiveRuntime {
61 pub(crate) async fn fuse_with_strategy(
72 &self,
73 text_hits: Vec<TextSearchHit>,
74 vector_hits: Vec<VectorSearchHit>,
75 strategy: &FusionStrategy,
76 limit: usize,
77 ) -> RuntimeResult<Vec<SearchHit>> {
78 match strategy {
79 FusionStrategy::VectorOnly => {
80 self.fuse_sources(Vec::new(), vector_hits, strategy, limit)
81 .await
82 }
83 FusionStrategy::KeywordOnly => {
84 self.fuse_sources(text_hits, Vec::new(), strategy, limit)
85 .await
86 }
87 FusionStrategy::Rrf { .. }
88 | FusionStrategy::Weighted { .. }
89 | FusionStrategy::Union
90 | FusionStrategy::Custom { .. } => {
91 self.fuse_sources(text_hits, vector_hits, strategy, limit)
92 .await
93 }
94 }
95 }
96
97 async fn fuse_sources(
98 &self,
99 text_hits: Vec<TextSearchHit>,
100 vector_hits: Vec<VectorSearchHit>,
101 strategy: &FusionStrategy,
102 limit: usize,
103 ) -> RuntimeResult<Vec<SearchHit>> {
104 let mut metadata: HashMap<Uuid, SearchHit> =
105 HashMap::with_capacity(text_hits.len() + vector_hits.len());
106 let prefer_maximum_signal = matches!(
107 strategy,
108 FusionStrategy::Weighted { .. } | FusionStrategy::Union
109 );
110
111 let text_source: Vec<(Uuid, DeterministicScore)> = text_hits
112 .into_iter()
113 .map(|h| {
114 let hit = SearchHit {
115 entity_id: h.subject_id,
116 score: h.score,
117 rank_score_kind: RankScoreKind::Keyword,
118 signals: SearchSignals {
119 vector_similarity: None,
120 keyword_score: Some(h.score),
121 },
122 source: SearchSource::Text,
123 title: h.title,
124 snippet: h.snippet,
125 };
126 let id = hit.entity_id;
127 let score = hit.score;
128 merge_metadata(&mut metadata, hit, prefer_maximum_signal);
129 (id, score)
130 })
131 .collect();
132
133 let vector_source: Vec<(Uuid, DeterministicScore)> = vector_hits
134 .into_iter()
135 .map(|h| {
136 let hit = SearchHit {
137 entity_id: h.subject_id,
138 score: h.score,
139 rank_score_kind: RankScoreKind::Vector,
140 signals: SearchSignals {
141 vector_similarity: Some(h.score),
142 keyword_score: None,
143 },
144 source: SearchSource::Vector,
145 title: None,
146 snippet: None,
147 };
148 let id = hit.entity_id;
149 let score = hit.score;
150 merge_metadata(&mut metadata, hit, prefer_maximum_signal);
151 (id, score)
152 })
153 .collect();
154
155 let sources: Vec<Vec<(Uuid, DeterministicScore)>> = vec![vector_source, text_source];
158
159 let (rank_score_kind, fused) = self.dispatch_fusion(sources, strategy, limit).await?;
160
161 Ok(fused
162 .into_iter()
163 .filter_map(|(id, score)| {
164 let mut hit = metadata.remove(&id)?;
165 hit.score = score;
166 hit.rank_score_kind = rank_score_kind;
167 Some(hit)
168 })
169 .collect())
170 }
171
172 async fn dispatch_fusion(
181 &self,
182 sources: Vec<Vec<(Uuid, DeterministicScore)>>,
183 strategy: &FusionStrategy,
184 limit: usize,
185 ) -> RuntimeResult<(RankScoreKind, Vec<RankedHit>)> {
186 let rank_score_kind = match strategy {
187 FusionStrategy::Rrf { .. } => RankScoreKind::Rrf,
188 FusionStrategy::VectorOnly => RankScoreKind::Vector,
189 FusionStrategy::KeywordOnly => RankScoreKind::Keyword,
190 FusionStrategy::Weighted { .. } => RankScoreKind::Weighted,
191 FusionStrategy::Union => RankScoreKind::Union,
192 FusionStrategy::Custom { name, params } => {
193 let executor = self.fusion_executor(name)?;
194 let rank_score_kind = executor.rank_score_kind();
195 if limit == 0 || sources.iter().all(Vec::is_empty) {
196 return Ok((rank_score_kind, Vec::new()));
197 }
198 let mut hits = executor.fuse(sources, params, limit).await?;
199 hits.sort_by(khive_fusion::cmp_desc_then_id);
200 hits.truncate(limit);
201 return Ok((rank_score_kind, hits));
202 }
203 };
204 Ok((
205 rank_score_kind,
206 khive_fusion::fuse(sources, strategy, limit)?,
207 ))
208 }
209}
210
211fn merge_metadata(
212 metadata: &mut HashMap<Uuid, SearchHit>,
213 hit: SearchHit,
214 prefer_maximum_signal: bool,
215) {
216 match metadata.entry(hit.entity_id) {
217 Entry::Occupied(mut entry) => {
218 let existing = entry.get_mut();
219 existing.source = merge_sources(existing.source, hit.source);
220 existing.signals.vector_similarity = if prefer_maximum_signal {
223 existing
224 .signals
225 .vector_similarity
226 .max(hit.signals.vector_similarity)
227 } else {
228 existing
229 .signals
230 .vector_similarity
231 .or(hit.signals.vector_similarity)
232 };
233 existing.signals.keyword_score = if prefer_maximum_signal {
234 existing
235 .signals
236 .keyword_score
237 .max(hit.signals.keyword_score)
238 } else {
239 existing.signals.keyword_score.or(hit.signals.keyword_score)
240 };
241 if existing.title.is_none() {
242 existing.title = hit.title;
243 }
244 if existing.snippet.is_none() {
245 existing.snippet = hit.snippet;
246 }
247 }
248 Entry::Vacant(entry) => {
249 entry.insert(hit);
250 }
251 }
252}
253
254fn merge_sources(left: SearchSource, right: SearchSource) -> SearchSource {
255 match (left, right) {
256 (SearchSource::Both, _) | (_, SearchSource::Both) => SearchSource::Both,
257 (SearchSource::Text, SearchSource::Vector) | (SearchSource::Vector, SearchSource::Text) => {
258 SearchSource::Both
259 }
260 (SearchSource::Text, SearchSource::Text) => SearchSource::Text,
261 (SearchSource::Vector, SearchSource::Vector) => SearchSource::Vector,
262 }
263}
264
265#[cfg(test)]
266mod tests {
267 use super::*;
268 use chrono::Utc;
269 use khive_storage::types::{
270 TextDocument, TextFilter, TextQueryMode, TextSearchHit, TextSearchRequest, VectorSearchHit,
271 VectorSearchRequest,
272 };
273 use khive_storage::Entity;
274 use khive_types::SubstrateKind;
275 use lattice_embed::EmbeddingModel;
276 use std::collections::HashSet;
277 use std::sync::Arc;
278
279 use crate::error::RuntimeError;
280 use crate::retrieval::CANDIDATE_MULTIPLIER;
281 use crate::runtime::NamespaceToken;
282 use crate::RuntimeConfig;
283
284 fn text_hit(id: Uuid, score: f64, title: &str) -> TextSearchHit {
285 TextSearchHit {
286 subject_id: id,
287 score: DeterministicScore::from_f64(score),
288 rank: 1,
289 title: Some(title.to_string()),
290 snippet: Some("...".to_string()),
291 }
292 }
293
294 fn vector_hit(id: Uuid, score: f64) -> VectorSearchHit {
295 VectorSearchHit {
296 subject_id: id,
297 score: DeterministicScore::from_f64(score),
298 rank: 1,
299 }
300 }
301
302 fn evidence_runtime() -> KhiveRuntime {
303 let backend = Arc::new(crate::StorageBackend::memory().expect("in-memory backend"));
304 backend.prepare_core_schema().expect("core schema");
305 KhiveRuntime::from_backend(
306 backend,
307 RuntimeConfig {
308 db_path: None,
309 events_split: None,
310 actor_id: Some("test:fusion-evidence".into()),
311 ..RuntimeConfig::no_embeddings()
312 },
313 )
314 }
315
316 #[tokio::test]
317 async fn fusion_evidence_labels_builtin_strategies_and_preserves_components() {
318 let rt = evidence_runtime();
319 let id = Uuid::from_u128(1);
320 let keyword = DeterministicScore::from_raw(1_i64 << 30);
321 let vector = DeterministicScore::from_raw(3_i64 << 30);
322 for (strategy, kind, raw_score, signals) in [
323 (
324 FusionStrategy::Rrf { k: 60 },
325 RankScoreKind::Rrf,
326 140_818_600,
327 SearchSignals {
328 vector_similarity: Some(vector),
329 keyword_score: Some(keyword),
330 },
331 ),
332 (
333 FusionStrategy::VectorOnly,
334 RankScoreKind::Vector,
335 vector.to_raw(),
336 SearchSignals {
337 vector_similarity: Some(vector),
338 keyword_score: None,
339 },
340 ),
341 (
342 FusionStrategy::KeywordOnly,
343 RankScoreKind::Keyword,
344 keyword.to_raw(),
345 SearchSignals {
346 vector_similarity: None,
347 keyword_score: Some(keyword),
348 },
349 ),
350 (
351 FusionStrategy::weighted(vec![0.5, 0.5]),
352 RankScoreKind::Weighted,
353 1_i64 << 32,
354 SearchSignals {
355 vector_similarity: Some(vector),
356 keyword_score: Some(keyword),
357 },
358 ),
359 (
360 FusionStrategy::Union,
361 RankScoreKind::Union,
362 vector.to_raw(),
363 SearchSignals {
364 vector_similarity: Some(vector),
365 keyword_score: Some(keyword),
366 },
367 ),
368 ] {
369 let hits = rt
370 .fuse_with_strategy(
371 vec![text_hit(id, 0.25, "candidate")],
372 vec![vector_hit(id, 0.75)],
373 &strategy,
374 10,
375 )
376 .await
377 .unwrap();
378 assert_eq!(hits.len(), 1);
379 assert_eq!(hits[0].entity_id, id);
380 assert_eq!(hits[0].score.to_raw(), raw_score);
381 assert_eq!(hits[0].rank_score_kind, kind);
382 assert_eq!(hits[0].signals, signals);
383 }
384 assert_eq!(RankScoreKind::Rrf.as_str(), "rrf");
385 assert_eq!(RankScoreKind::Vector.as_str(), "vector");
386 assert_eq!(RankScoreKind::Keyword.as_str(), "keyword");
387 assert_eq!(RankScoreKind::Weighted.as_str(), "weighted");
388 assert_eq!(RankScoreKind::Union.as_str(), "union");
389 }
390
391 #[tokio::test]
392 async fn fusion_evidence_distinguishes_absence_from_zero() {
393 let rt = evidence_runtime();
394 let id = Uuid::from_u128(1);
395 for (text, vector, signals) in [
396 (
397 vec![text_hit(id, 0.0, "zero keyword")],
398 vec![],
399 SearchSignals {
400 vector_similarity: None,
401 keyword_score: Some(DeterministicScore::ZERO),
402 },
403 ),
404 (
405 vec![],
406 vec![vector_hit(id, 0.0)],
407 SearchSignals {
408 vector_similarity: Some(DeterministicScore::ZERO),
409 keyword_score: None,
410 },
411 ),
412 ] {
413 let hits = rt
414 .fuse_with_strategy(text, vector, &FusionStrategy::rrf(), 10)
415 .await
416 .unwrap();
417 assert_eq!(hits.len(), 1);
418 assert_eq!(hits[0].signals, signals);
419 }
420 assert_eq!(
421 SearchSignals::default(),
422 SearchSignals {
423 vector_similarity: None,
424 keyword_score: None,
425 }
426 );
427 }
428
429 #[tokio::test]
430 async fn fusion_evidence_golden_preserves_true_ties_across_permutations() {
431 let rt = evidence_runtime();
432 let a = Uuid::from_u128(1);
433 let b = Uuid::from_u128(2);
434 let expected = vec![
435 (
436 a,
437 139_682_966,
438 RankScoreKind::Rrf,
439 SearchSignals {
440 vector_similarity: Some(DeterministicScore::from_raw(1_i64 << 31)),
441 keyword_score: Some(DeterministicScore::from_raw(1_i64 << 30)),
442 },
443 ),
444 (
445 b,
446 139_682_966,
447 RankScoreKind::Rrf,
448 SearchSignals {
449 vector_similarity: Some(DeterministicScore::from_raw(1_i64 << 32)),
450 keyword_score: Some(DeterministicScore::from_raw(3_i64 << 30)),
451 },
452 ),
453 ];
454 for (text_ids, vector_ids) in [([a, b], [b, a]), ([b, a], [a, b])] {
455 for _ in 0..4 {
456 let text = text_ids
457 .into_iter()
458 .map(|id| text_hit(id, if id == a { 0.25 } else { 0.75 }, "candidate"))
459 .collect();
460 let vector = vector_ids
461 .into_iter()
462 .map(|id| vector_hit(id, if id == a { 0.5 } else { 1.0 }))
463 .collect();
464 let hits = rt
465 .fuse_with_strategy(text, vector, &FusionStrategy::Rrf { k: 60 }, 10)
466 .await
467 .unwrap();
468 assert_eq!(hits.len(), 2);
469 assert_eq!(hits[0].score, hits[1].score);
470 assert!(hits.iter().all(|hit| hit.source == SearchSource::Both));
471 let snapshot: Vec<_> = hits
472 .iter()
473 .map(|hit| {
474 (
475 hit.entity_id,
476 hit.score.to_raw(),
477 hit.rank_score_kind,
478 hit.signals,
479 )
480 })
481 .collect();
482 assert_eq!(snapshot, expected);
483 }
484 }
485 }
486
487 #[tokio::test]
488 async fn fusion_evidence_duplicate_selection_follows_strategy() {
489 let rt = evidence_runtime();
490 let id = Uuid::from_u128(1);
491 for (strategy, keyword_raw) in [
492 (FusionStrategy::rrf(), 1_i64 << 30),
493 (FusionStrategy::Union, 3_i64 << 30),
494 (FusionStrategy::weighted(vec![0.5, 0.5]), 3_i64 << 30),
495 ] {
496 let hits = rt
497 .fuse_with_strategy(
498 vec![text_hit(id, 0.25, "first"), text_hit(id, 0.75, "second")],
499 vec![vector_hit(id, 0.5)],
500 &strategy,
501 10,
502 )
503 .await
504 .unwrap();
505 assert_eq!(hits.len(), 1);
506 assert_eq!(
507 hits[0].signals,
508 SearchSignals {
509 vector_similarity: Some(DeterministicScore::from_raw(1_i64 << 31)),
510 keyword_score: Some(DeterministicScore::from_raw(keyword_raw)),
511 }
512 );
513 assert_eq!(hits[0].title.as_deref(), Some("first"));
514 }
515 }
516
517 #[tokio::test]
518 async fn custom_fusion_evidence_uses_declared_kind() {
519 let rt = evidence_runtime();
520 let id = Uuid::from_u128(1);
521 rt.register_fusion_strategy("invert", Arc::new(InvertScoreExecutor));
522 let strategy =
523 FusionStrategy::try_custom("invert".into(), serde_json::Value::Null).unwrap();
524 let hits = rt
525 .fuse_with_strategy(vec![text_hit(id, 0.75, "candidate")], vec![], &strategy, 10)
526 .await
527 .unwrap();
528 assert_eq!(hits.len(), 1);
529 assert_eq!(hits[0].score.to_raw(), 1_i64 << 30);
530 assert_eq!(hits[0].rank_score_kind, RankScoreKind::Weighted);
531 assert_eq!(
532 hits[0].signals,
533 SearchSignals {
534 vector_similarity: None,
535 keyword_score: Some(DeterministicScore::from_raw(3_i64 << 30)),
536 }
537 );
538 }
539
540 fn cosine_fixture_vector(dimensions: usize, x: f32, y: f32) -> Vec<f32> {
541 let mut vector = vec![0.0; dimensions];
542 vector[0] = x;
543 vector[1] = y;
544 vector
545 }
546
547 async fn stale_full_prefix_fixture() -> (
548 KhiveRuntime,
549 NamespaceToken,
550 &'static str,
551 Vec<f32>,
552 Vec<TextSearchHit>,
553 Vec<VectorSearchHit>,
554 HashSet<Uuid>,
555 ) {
556 let model = EmbeddingModel::AllMiniLmL6V2;
557 let dimensions = model.dimensions();
558 let rt = KhiveRuntime::new(RuntimeConfig {
559 db_path: None,
560 embedding_model: Some(model),
561 additional_embedding_models: vec![],
562 ..RuntimeConfig::default()
563 })
564 .unwrap();
565 let tok = NamespaceToken::local();
566 let query_text = "fusionrefillterm";
567 let query_vector = cosine_fixture_vector(dimensions, 1.0, 0.0);
568
569 let common_stale_a = Uuid::from_u128(1);
570 let common_stale_b = Uuid::from_u128(2);
571 let text_only_stale = Uuid::from_u128(3);
572 let vector_only_stale = Uuid::from_u128(4);
573
574 let live_text = Entity::new("local", "concept", "live text candidate");
575 let live_vector = Entity::new("local", "concept", "live vector candidate");
576 rt.entities(&tok)
577 .unwrap()
578 .upsert_entities(vec![live_text.clone(), live_vector.clone()])
579 .await
580 .unwrap();
581
582 let document = |subject_id, repetitions: usize| TextDocument {
583 subject_id,
584 kind: SubstrateKind::Entity,
585 record_kind: None,
586 namespace: "local".to_string(),
587 title: None,
588 body: std::iter::repeat_n(query_text, repetitions)
589 .collect::<Vec<_>>()
590 .join(" "),
591 tags: vec![],
592 metadata: None,
593 updated_at: Utc::now(),
594 };
595 rt.text(&tok)
596 .unwrap()
597 .upsert_documents(vec![
598 document(common_stale_a, 12),
599 document(common_stale_b, 8),
600 document(text_only_stale, 4),
601 document(live_text.id, 1),
602 ])
603 .await
604 .unwrap();
605
606 let vectors = rt.vectors(&tok).unwrap();
607 for (id, vector) in [
608 (common_stale_a, cosine_fixture_vector(dimensions, 1.0, 0.0)),
609 (common_stale_b, cosine_fixture_vector(dimensions, 0.8, 0.6)),
610 (
611 vector_only_stale,
612 cosine_fixture_vector(dimensions, 0.5, 0.866_025_4),
613 ),
614 (live_vector.id, cosine_fixture_vector(dimensions, -1.0, 0.0)),
615 ] {
616 vectors
617 .insert(
618 id,
619 SubstrateKind::Entity,
620 "local",
621 "entity.body",
622 vec![vector],
623 )
624 .await
625 .unwrap();
626 }
627
628 let text_hits = rt
629 .text(&tok)
630 .unwrap()
631 .search(TextSearchRequest {
632 query: query_text.to_string(),
633 mode: TextQueryMode::Plain,
634 filter: Some(TextFilter {
635 namespaces: vec!["local".to_string()],
636 ..TextFilter::default()
637 }),
638 top_k: CANDIDATE_MULTIPLIER,
639 snippet_chars: 0,
640 })
641 .await
642 .unwrap();
643 let vector_hits = vectors
644 .search(VectorSearchRequest {
645 query_vectors: vec![query_vector.clone()],
646 top_k: CANDIDATE_MULTIPLIER,
647 namespace: Some("local".to_string()),
648 kind: Some(SubstrateKind::Entity),
649 embedding_model: None,
650 filter: None,
651 backend_hints: None,
652 })
653 .await
654 .unwrap();
655
656 assert_eq!(text_hits.len(), CANDIDATE_MULTIPLIER as usize);
657 assert_eq!(vector_hits.len(), CANDIDATE_MULTIPLIER as usize);
658 let live = HashSet::from([live_text.id, live_vector.id]);
659 (
660 rt,
661 tok,
662 query_text,
663 query_vector,
664 text_hits,
665 vector_hits,
666 live,
667 )
668 }
669
670 #[tokio::test]
672 async fn rrf_custom_k_differs_from_k60() {
673 let rt = KhiveRuntime::memory().unwrap();
674 let a = Uuid::new_v4();
675 let b = Uuid::new_v4();
676 let text = vec![text_hit(a, 0.9, "a"), text_hit(b, 0.1, "b")];
680 let hits_k1 = rt
681 .fuse_with_strategy(text.clone(), vec![], &FusionStrategy::Rrf { k: 1 }, 10)
682 .await
683 .unwrap();
684 let hits_k60 = rt
685 .fuse_with_strategy(text, vec![], &FusionStrategy::Rrf { k: 60 }, 10)
686 .await
687 .unwrap();
688 assert_eq!(hits_k1[0].entity_id, a);
690 assert_eq!(hits_k60[0].entity_id, a);
691 assert!(hits_k1[0].score > hits_k60[0].score);
693 }
694
695 #[tokio::test]
697 async fn weighted_ordering_depends_on_weights() {
698 let rt = KhiveRuntime::memory().unwrap();
699 let a = Uuid::new_v4();
700 let b = Uuid::new_v4();
701 let text = vec![text_hit(a, 0.9, "a"), text_hit(b, 0.1, "b")];
703 let vec_hits = vec![vector_hit(b, 0.9), vector_hit(a, 0.1)];
704
705 let heavy_vector = rt
706 .fuse_with_strategy(
707 text.clone(),
708 vec_hits.clone(),
709 &FusionStrategy::Weighted {
710 weights: vec![0.7, 0.3],
711 },
712 10,
713 )
714 .await
715 .unwrap();
716 let heavy_keyword = rt
717 .fuse_with_strategy(
718 text,
719 vec_hits,
720 &FusionStrategy::Weighted {
721 weights: vec![0.3, 0.7],
722 },
723 10,
724 )
725 .await
726 .unwrap();
727
728 assert_eq!(heavy_vector[0].entity_id, b);
729 assert_eq!(heavy_keyword[0].entity_id, a);
730 }
731
732 #[tokio::test]
734 async fn weighted_scale_invariant() {
735 let rt = KhiveRuntime::memory().unwrap();
736 let a = Uuid::new_v4();
737 let b = Uuid::new_v4();
738 let text = vec![text_hit(a, 0.9, "a"), text_hit(b, 0.1, "b")];
739 let vec_hits = vec![vector_hit(b, 0.9), vector_hit(a, 0.1)];
740
741 let w1 = rt
742 .fuse_with_strategy(
743 text.clone(),
744 vec_hits.clone(),
745 &FusionStrategy::Weighted {
746 weights: vec![0.7, 0.3],
747 },
748 10,
749 )
750 .await
751 .unwrap();
752 let w2 = rt
753 .fuse_with_strategy(
754 text,
755 vec_hits,
756 &FusionStrategy::Weighted {
757 weights: vec![7.0, 3.0],
758 },
759 10,
760 )
761 .await
762 .unwrap();
763
764 assert_eq!(w1[0].entity_id, w2[0].entity_id);
765 assert_eq!(w1[1].entity_id, w2[1].entity_id);
766 let diff = (w1[0].score.to_f64() - w2[0].score.to_f64()).abs();
767 assert!(diff < 1e-9, "scores differ by {diff}");
768 }
769
770 #[tokio::test]
772 async fn weighted_zero_weights_equal_fallback() {
773 let rt = KhiveRuntime::memory().unwrap();
774 let a = Uuid::new_v4();
775 let b = Uuid::new_v4();
776 let text = vec![text_hit(a, 0.9, "a"), text_hit(b, 0.1, "b")];
778 let vec_hits = vec![vector_hit(a, 0.9), vector_hit(b, 0.1)];
779
780 let hits = rt
781 .fuse_with_strategy(
782 text,
783 vec_hits,
784 &FusionStrategy::Weighted {
785 weights: vec![0.0, 0.0],
786 },
787 10,
788 )
789 .await
790 .unwrap();
791 assert_eq!(hits[0].entity_id, a);
792 }
793
794 #[tokio::test]
796 async fn weighted_negative_weight_clamped() {
797 let rt = KhiveRuntime::memory().unwrap();
798 let a = Uuid::new_v4();
799 let text = vec![text_hit(a, 0.9, "a")];
800 let hits = rt
802 .fuse_with_strategy(
803 text,
804 vec![],
805 &FusionStrategy::Weighted {
806 weights: vec![-0.5, 1.0],
807 },
808 10,
809 )
810 .await
811 .unwrap();
812 assert_eq!(hits.len(), 1);
813 assert_eq!(hits[0].entity_id, a);
814 }
815
816 #[tokio::test]
817 async fn weighted_empty_arm_keeps_canonical_position() {
818 let rt = KhiveRuntime::memory().unwrap();
819 let text_only = Uuid::new_v4();
820 let hits = rt
821 .fuse_with_strategy(
822 vec![text_hit(text_only, 0.9, "text")],
823 vec![],
824 &FusionStrategy::Weighted {
825 weights: vec![1.0, 0.0],
827 },
828 10,
829 )
830 .await
831 .unwrap();
832 assert!(
833 hits.is_empty(),
834 "dropping the empty vector arm would incorrectly rebind text to its weight"
835 );
836 }
837
838 #[tokio::test]
840 async fn union_max_score_per_entity() {
841 let rt = KhiveRuntime::memory().unwrap();
842 let a = Uuid::new_v4();
843 let text = vec![text_hit(a, 0.3, "a")];
844 let vec_hits = vec![vector_hit(a, 0.9)];
845
846 let hits = rt
847 .fuse_with_strategy(text, vec_hits, &FusionStrategy::Union, 10)
848 .await
849 .unwrap();
850 assert_eq!(hits.len(), 1);
851 assert!((hits[0].score.to_f64() - 0.9).abs() < 1e-6);
852 assert_eq!(hits[0].source, SearchSource::Both);
853 }
854
855 #[tokio::test]
857 async fn vector_only_drops_text() {
858 let rt = KhiveRuntime::memory().unwrap();
859 let a = Uuid::new_v4();
860 let b = Uuid::new_v4();
861 let text = vec![text_hit(b, 0.9, "b")];
862 let vec_hits = vec![vector_hit(a, 0.8)];
863
864 let hits = rt
865 .fuse_with_strategy(text, vec_hits, &FusionStrategy::VectorOnly, 10)
866 .await
867 .unwrap();
868 assert_eq!(hits.len(), 1);
869 assert_eq!(hits[0].entity_id, a);
870 assert_eq!(hits[0].source, SearchSource::Vector);
871 assert!(hits[0].title.is_none());
872 }
873
874 #[tokio::test]
875 async fn keyword_only_drops_vector() {
876 let rt = KhiveRuntime::memory().unwrap();
877 let text_id = Uuid::new_v4();
878 let vector_id = Uuid::new_v4();
879 let hits = rt
880 .fuse_with_strategy(
881 vec![text_hit(text_id, 0.8, "text")],
882 vec![vector_hit(vector_id, 0.9)],
883 &FusionStrategy::KeywordOnly,
884 10,
885 )
886 .await
887 .unwrap();
888 assert_eq!(hits.len(), 1);
889 assert_eq!(hits[0].entity_id, text_id);
890 assert_eq!(hits[0].source, SearchSource::Text);
891 }
892
893 struct ReverseOrderExecutor;
896
897 #[async_trait::async_trait]
898 impl FusionExecutor for ReverseOrderExecutor {
899 fn rank_score_kind(&self) -> RankScoreKind {
900 RankScoreKind::Union
901 }
902
903 async fn fuse(
904 &self,
905 streams: Vec<CandidateStream>,
906 _params: &serde_json::Value,
907 _limit: usize,
908 ) -> RuntimeResult<Vec<RankedHit>> {
909 let mut flat: Vec<_> = streams.into_iter().flatten().collect();
910 flat.reverse();
911 Ok(flat)
912 }
913 }
914
915 struct InvertScoreExecutor;
921
922 #[async_trait::async_trait]
923 impl FusionExecutor for InvertScoreExecutor {
924 fn rank_score_kind(&self) -> RankScoreKind {
925 RankScoreKind::Weighted
926 }
927
928 async fn fuse(
929 &self,
930 streams: Vec<CandidateStream>,
931 _params: &serde_json::Value,
932 _limit: usize,
933 ) -> RuntimeResult<Vec<RankedHit>> {
934 Ok(streams
935 .into_iter()
936 .flatten()
937 .map(|(id, score)| (id, DeterministicScore::from_f64(1.0 - score.to_f64())))
938 .collect())
939 }
940 }
941
942 struct EqualScoreExecutor;
947
948 #[async_trait::async_trait]
949 impl FusionExecutor for EqualScoreExecutor {
950 fn rank_score_kind(&self) -> RankScoreKind {
951 RankScoreKind::Weighted
952 }
953
954 async fn fuse(
955 &self,
956 streams: Vec<CandidateStream>,
957 _params: &serde_json::Value,
958 _limit: usize,
959 ) -> RuntimeResult<Vec<RankedHit>> {
960 Ok(streams
961 .into_iter()
962 .flatten()
963 .map(|(id, _)| (id, DeterministicScore::from_f64(1.0)))
964 .collect())
965 }
966 }
967
968 #[tokio::test]
971 async fn custom_strategy_dispatches_through_executor_and_differs_from_rrf() {
972 let rt = KhiveRuntime::memory().unwrap();
973 let a = Uuid::new_v4();
974 let b = Uuid::new_v4();
975 let text = vec![text_hit(a, 0.9, "a"), text_hit(b, 0.5, "b")];
976
977 rt.register_fusion_strategy("invert", Arc::new(InvertScoreExecutor));
978 let strategy =
979 FusionStrategy::try_custom("invert".to_string(), serde_json::Value::Null).unwrap();
980
981 let custom = rt
982 .fuse_with_strategy(text.clone(), vec![], &strategy, 10)
983 .await
984 .unwrap();
985 let rrf = rt
986 .fuse_with_strategy(text, vec![], &FusionStrategy::rrf(), 10)
987 .await
988 .unwrap();
989
990 let custom_ids: Vec<_> = custom.iter().map(|h| h.entity_id).collect();
991 let rrf_ids: Vec<_> = rrf.iter().map(|h| h.entity_id).collect();
992 assert_ne!(
993 custom_ids, rrf_ids,
994 "custom and RRF must yield different orderings on this fixture"
995 );
996 }
997
998 #[tokio::test]
1000 async fn custom_strategy_unknown_name_fails_closed() {
1001 let rt = KhiveRuntime::memory().unwrap();
1002 let a = Uuid::new_v4();
1003 let text = vec![text_hit(a, 0.9, "a")];
1004 let strategy =
1005 FusionStrategy::try_custom("nonexistent".to_string(), serde_json::Value::Null).unwrap();
1006
1007 let result = rt.fuse_with_strategy(text, vec![], &strategy, 10).await;
1008 assert!(matches!(
1009 result,
1010 Err(RuntimeError::UnknownFusionStrategy(name)) if name == "nonexistent"
1011 ));
1012 }
1013
1014 #[tokio::test]
1017 async fn custom_strategy_unknown_name_fails_closed_even_on_empty_input() {
1018 let rt = KhiveRuntime::memory().unwrap();
1019 let strategy =
1020 FusionStrategy::try_custom("nonexistent".to_string(), serde_json::Value::Null).unwrap();
1021
1022 let result = rt.fuse_with_strategy(vec![], vec![], &strategy, 10).await;
1023 assert!(matches!(
1024 result,
1025 Err(RuntimeError::UnknownFusionStrategy(name)) if name == "nonexistent"
1026 ));
1027 }
1028
1029 #[tokio::test]
1032 async fn custom_strategy_registered_name_empty_input_returns_ok_empty() {
1033 let rt = KhiveRuntime::memory().unwrap();
1034 rt.register_fusion_strategy("reverse", Arc::new(ReverseOrderExecutor));
1035 let strategy =
1036 FusionStrategy::try_custom("reverse".to_string(), serde_json::Value::Null).unwrap();
1037
1038 let result = rt
1039 .fuse_with_strategy(vec![], vec![], &strategy, 10)
1040 .await
1041 .unwrap();
1042 assert!(result.is_empty());
1043 }
1044
1045 #[tokio::test]
1047 async fn registered_custom_strategy_leaves_default_path_unaffected() {
1048 let rt = KhiveRuntime::memory().unwrap();
1049 let a = Uuid::new_v4();
1050 let b = Uuid::new_v4();
1051 let text = vec![text_hit(a, 0.9, "a"), text_hit(b, 0.5, "b")];
1052
1053 rt.register_fusion_strategy("reverse", Arc::new(ReverseOrderExecutor));
1054
1055 let via_rt_with_registration = rt
1056 .fuse_with_strategy(text.clone(), vec![], &FusionStrategy::rrf(), 10)
1057 .await
1058 .unwrap();
1059 let rt2 = KhiveRuntime::memory().unwrap();
1060 let via_rt_without_registration = rt2
1061 .fuse_with_strategy(text, vec![], &FusionStrategy::rrf(), 10)
1062 .await
1063 .unwrap();
1064
1065 let ids_with: Vec<_> = via_rt_with_registration
1066 .iter()
1067 .map(|h| h.entity_id)
1068 .collect();
1069 let ids_without: Vec<_> = via_rt_without_registration
1070 .iter()
1071 .map(|h| h.entity_id)
1072 .collect();
1073 assert_eq!(ids_with, ids_without);
1074 }
1075
1076 #[tokio::test]
1080 async fn custom_executor_output_is_sorted_by_canonical_comparator() {
1081 let rt = KhiveRuntime::memory().unwrap();
1082 let ids: Vec<Uuid> = vec![Uuid::from_u128(3), Uuid::from_u128(1), Uuid::from_u128(2)];
1084 let text: Vec<TextSearchHit> = ids.iter().map(|&id| text_hit(id, 0.5, "tied")).collect();
1085
1086 rt.register_fusion_strategy("equal_score", Arc::new(EqualScoreExecutor));
1087 let strategy =
1088 FusionStrategy::try_custom("equal_score".to_string(), serde_json::Value::Null).unwrap();
1089
1090 let hits = rt
1091 .fuse_with_strategy(text, vec![], &strategy, 10)
1092 .await
1093 .unwrap();
1094
1095 let mut expected = ids.clone();
1096 expected.sort();
1097 let actual: Vec<_> = hits.iter().map(|h| h.entity_id).collect();
1098 assert_eq!(
1099 actual, expected,
1100 "equal-score executor output must be tie-broken by ascending ID"
1101 );
1102 }
1103
1104 #[test]
1106 fn default_strategy_is_rrf_k60() {
1107 assert_eq!(FusionStrategy::default(), FusionStrategy::Rrf { k: 60 });
1108 }
1109
1110 #[tokio::test]
1111 async fn hybrid_default_rrf_alive_filter_refills_below_complete_four_x_prefix() {
1112 let (rt, tok, query_text, query_vector, _text_hits, _vector_hits, live) =
1113 stale_full_prefix_fixture().await;
1114
1115 let hits = rt
1116 .hybrid_search(
1117 &tok,
1118 query_text,
1119 Some(query_vector),
1120 1,
1121 None,
1122 None,
1123 &[],
1124 None,
1125 )
1126 .await
1127 .unwrap();
1128
1129 assert_eq!(hits.len(), 1);
1130 assert!(live.contains(&hits[0].entity_id));
1131 }
1132
1133 #[test]
1135 fn serde_roundtrip() {
1136 let cases = vec![
1137 FusionStrategy::Rrf { k: 60 },
1138 FusionStrategy::Rrf { k: 20 },
1139 FusionStrategy::Weighted {
1140 weights: vec![0.7, 0.3],
1141 },
1142 FusionStrategy::Union,
1143 FusionStrategy::VectorOnly,
1144 FusionStrategy::KeywordOnly,
1145 ];
1146 for strategy in cases {
1147 let json = serde_json::to_string(&strategy).expect("serialize");
1148 let back: FusionStrategy = serde_json::from_str(&json).expect("deserialize");
1149 assert_eq!(strategy, back, "roundtrip failed for {json}");
1150 }
1151 }
1152}