use std::collections::{hash_map::Entry, HashMap};
use uuid::Uuid;
use khive_score::DeterministicScore;
use khive_storage::types::{TextSearchHit, VectorSearchHit};
use crate::error::RuntimeResult;
use crate::retrieval::{RankScoreKind, SearchHit, SearchSignals, SearchSource};
use crate::runtime::KhiveRuntime;
pub use khive_fusion::FusionStrategy;
pub type CandidateStream = Vec<(Uuid, DeterministicScore)>;
pub type RankedHit = (Uuid, DeterministicScore);
#[async_trait::async_trait]
pub trait FusionExecutor: Send + Sync + 'static {
fn rank_score_kind(&self) -> RankScoreKind;
async fn fuse(
&self,
streams: Vec<CandidateStream>,
params: &serde_json::Value,
limit: usize,
) -> RuntimeResult<Vec<RankedHit>>;
}
pub(crate) async fn rrf_fuse_k(
rt: &KhiveRuntime,
text_hits: Vec<TextSearchHit>,
vector_hits: Vec<VectorSearchHit>,
k: usize,
limit: usize,
) -> RuntimeResult<Vec<SearchHit>> {
rt.fuse_with_strategy(text_hits, vector_hits, &FusionStrategy::Rrf { k }, limit)
.await
}
impl KhiveRuntime {
pub(crate) async fn fuse_with_strategy(
&self,
text_hits: Vec<TextSearchHit>,
vector_hits: Vec<VectorSearchHit>,
strategy: &FusionStrategy,
limit: usize,
) -> RuntimeResult<Vec<SearchHit>> {
match strategy {
FusionStrategy::VectorOnly => {
self.fuse_sources(Vec::new(), vector_hits, strategy, limit)
.await
}
FusionStrategy::KeywordOnly => {
self.fuse_sources(text_hits, Vec::new(), strategy, limit)
.await
}
FusionStrategy::Rrf { .. }
| FusionStrategy::Weighted { .. }
| FusionStrategy::Union
| FusionStrategy::Custom { .. } => {
self.fuse_sources(text_hits, vector_hits, strategy, limit)
.await
}
}
}
async fn fuse_sources(
&self,
text_hits: Vec<TextSearchHit>,
vector_hits: Vec<VectorSearchHit>,
strategy: &FusionStrategy,
limit: usize,
) -> RuntimeResult<Vec<SearchHit>> {
let mut metadata: HashMap<Uuid, SearchHit> =
HashMap::with_capacity(text_hits.len() + vector_hits.len());
let prefer_maximum_signal = matches!(
strategy,
FusionStrategy::Weighted { .. } | FusionStrategy::Union
);
let text_source: Vec<(Uuid, DeterministicScore)> = text_hits
.into_iter()
.map(|h| {
let hit = SearchHit {
entity_id: h.subject_id,
score: h.score,
rank_score_kind: RankScoreKind::Keyword,
signals: SearchSignals {
vector_similarity: None,
keyword_score: Some(h.score),
},
source: SearchSource::Text,
title: h.title,
snippet: h.snippet,
};
let id = hit.entity_id;
let score = hit.score;
merge_metadata(&mut metadata, hit, prefer_maximum_signal);
(id, score)
})
.collect();
let vector_source: Vec<(Uuid, DeterministicScore)> = vector_hits
.into_iter()
.map(|h| {
let hit = SearchHit {
entity_id: h.subject_id,
score: h.score,
rank_score_kind: RankScoreKind::Vector,
signals: SearchSignals {
vector_similarity: Some(h.score),
keyword_score: None,
},
source: SearchSource::Vector,
title: None,
snippet: None,
};
let id = hit.entity_id;
let score = hit.score;
merge_metadata(&mut metadata, hit, prefer_maximum_signal);
(id, score)
})
.collect();
let sources: Vec<Vec<(Uuid, DeterministicScore)>> = vec![vector_source, text_source];
let (rank_score_kind, fused) = self.dispatch_fusion(sources, strategy, limit).await?;
Ok(fused
.into_iter()
.filter_map(|(id, score)| {
let mut hit = metadata.remove(&id)?;
hit.score = score;
hit.rank_score_kind = rank_score_kind;
Some(hit)
})
.collect())
}
async fn dispatch_fusion(
&self,
sources: Vec<Vec<(Uuid, DeterministicScore)>>,
strategy: &FusionStrategy,
limit: usize,
) -> RuntimeResult<(RankScoreKind, Vec<RankedHit>)> {
let rank_score_kind = match strategy {
FusionStrategy::Rrf { .. } => RankScoreKind::Rrf,
FusionStrategy::VectorOnly => RankScoreKind::Vector,
FusionStrategy::KeywordOnly => RankScoreKind::Keyword,
FusionStrategy::Weighted { .. } => RankScoreKind::Weighted,
FusionStrategy::Union => RankScoreKind::Union,
FusionStrategy::Custom { name, params } => {
let executor = self.fusion_executor(name)?;
let rank_score_kind = executor.rank_score_kind();
if limit == 0 || sources.iter().all(Vec::is_empty) {
return Ok((rank_score_kind, Vec::new()));
}
let mut hits = executor.fuse(sources, params, limit).await?;
hits.sort_by(khive_fusion::cmp_desc_then_id);
hits.truncate(limit);
return Ok((rank_score_kind, hits));
}
};
Ok((
rank_score_kind,
khive_fusion::fuse(sources, strategy, limit)?,
))
}
}
fn merge_metadata(
metadata: &mut HashMap<Uuid, SearchHit>,
hit: SearchHit,
prefer_maximum_signal: bool,
) {
match metadata.entry(hit.entity_id) {
Entry::Occupied(mut entry) => {
let existing = entry.get_mut();
existing.source = merge_sources(existing.source, hit.source);
existing.signals.vector_similarity = if prefer_maximum_signal {
existing
.signals
.vector_similarity
.max(hit.signals.vector_similarity)
} else {
existing
.signals
.vector_similarity
.or(hit.signals.vector_similarity)
};
existing.signals.keyword_score = if prefer_maximum_signal {
existing
.signals
.keyword_score
.max(hit.signals.keyword_score)
} else {
existing.signals.keyword_score.or(hit.signals.keyword_score)
};
if existing.title.is_none() {
existing.title = hit.title;
}
if existing.snippet.is_none() {
existing.snippet = hit.snippet;
}
}
Entry::Vacant(entry) => {
entry.insert(hit);
}
}
}
fn merge_sources(left: SearchSource, right: SearchSource) -> SearchSource {
match (left, right) {
(SearchSource::Both, _) | (_, SearchSource::Both) => SearchSource::Both,
(SearchSource::Text, SearchSource::Vector) | (SearchSource::Vector, SearchSource::Text) => {
SearchSource::Both
}
(SearchSource::Text, SearchSource::Text) => SearchSource::Text,
(SearchSource::Vector, SearchSource::Vector) => SearchSource::Vector,
}
}
#[cfg(test)]
mod tests {
use super::*;
use chrono::Utc;
use khive_storage::types::{
TextDocument, TextFilter, TextQueryMode, TextSearchHit, TextSearchRequest, VectorSearchHit,
VectorSearchRequest,
};
use khive_storage::Entity;
use khive_types::SubstrateKind;
use lattice_embed::EmbeddingModel;
use std::collections::HashSet;
use std::sync::Arc;
use crate::error::RuntimeError;
use crate::retrieval::CANDIDATE_MULTIPLIER;
use crate::runtime::NamespaceToken;
use crate::RuntimeConfig;
fn text_hit(id: Uuid, score: f64, title: &str) -> TextSearchHit {
TextSearchHit {
subject_id: id,
score: DeterministicScore::from_f64(score),
rank: 1,
title: Some(title.to_string()),
snippet: Some("...".to_string()),
}
}
fn vector_hit(id: Uuid, score: f64) -> VectorSearchHit {
VectorSearchHit {
subject_id: id,
score: DeterministicScore::from_f64(score),
rank: 1,
}
}
fn evidence_runtime() -> KhiveRuntime {
let backend = Arc::new(crate::StorageBackend::memory().expect("in-memory backend"));
backend.prepare_core_schema().expect("core schema");
KhiveRuntime::from_backend(
backend,
RuntimeConfig {
db_path: None,
events_split: None,
actor_id: Some("test:fusion-evidence".into()),
..RuntimeConfig::no_embeddings()
},
)
}
#[tokio::test]
async fn fusion_evidence_labels_builtin_strategies_and_preserves_components() {
let rt = evidence_runtime();
let id = Uuid::from_u128(1);
let keyword = DeterministicScore::from_raw(1_i64 << 30);
let vector = DeterministicScore::from_raw(3_i64 << 30);
for (strategy, kind, raw_score, signals) in [
(
FusionStrategy::Rrf { k: 60 },
RankScoreKind::Rrf,
140_818_600,
SearchSignals {
vector_similarity: Some(vector),
keyword_score: Some(keyword),
},
),
(
FusionStrategy::VectorOnly,
RankScoreKind::Vector,
vector.to_raw(),
SearchSignals {
vector_similarity: Some(vector),
keyword_score: None,
},
),
(
FusionStrategy::KeywordOnly,
RankScoreKind::Keyword,
keyword.to_raw(),
SearchSignals {
vector_similarity: None,
keyword_score: Some(keyword),
},
),
(
FusionStrategy::weighted(vec![0.5, 0.5]),
RankScoreKind::Weighted,
1_i64 << 32,
SearchSignals {
vector_similarity: Some(vector),
keyword_score: Some(keyword),
},
),
(
FusionStrategy::Union,
RankScoreKind::Union,
vector.to_raw(),
SearchSignals {
vector_similarity: Some(vector),
keyword_score: Some(keyword),
},
),
] {
let hits = rt
.fuse_with_strategy(
vec![text_hit(id, 0.25, "candidate")],
vec![vector_hit(id, 0.75)],
&strategy,
10,
)
.await
.unwrap();
assert_eq!(hits.len(), 1);
assert_eq!(hits[0].entity_id, id);
assert_eq!(hits[0].score.to_raw(), raw_score);
assert_eq!(hits[0].rank_score_kind, kind);
assert_eq!(hits[0].signals, signals);
}
assert_eq!(RankScoreKind::Rrf.as_str(), "rrf");
assert_eq!(RankScoreKind::Vector.as_str(), "vector");
assert_eq!(RankScoreKind::Keyword.as_str(), "keyword");
assert_eq!(RankScoreKind::Weighted.as_str(), "weighted");
assert_eq!(RankScoreKind::Union.as_str(), "union");
}
#[tokio::test]
async fn fusion_evidence_distinguishes_absence_from_zero() {
let rt = evidence_runtime();
let id = Uuid::from_u128(1);
for (text, vector, signals) in [
(
vec![text_hit(id, 0.0, "zero keyword")],
vec![],
SearchSignals {
vector_similarity: None,
keyword_score: Some(DeterministicScore::ZERO),
},
),
(
vec![],
vec![vector_hit(id, 0.0)],
SearchSignals {
vector_similarity: Some(DeterministicScore::ZERO),
keyword_score: None,
},
),
] {
let hits = rt
.fuse_with_strategy(text, vector, &FusionStrategy::rrf(), 10)
.await
.unwrap();
assert_eq!(hits.len(), 1);
assert_eq!(hits[0].signals, signals);
}
assert_eq!(
SearchSignals::default(),
SearchSignals {
vector_similarity: None,
keyword_score: None,
}
);
}
#[tokio::test]
async fn fusion_evidence_golden_preserves_true_ties_across_permutations() {
let rt = evidence_runtime();
let a = Uuid::from_u128(1);
let b = Uuid::from_u128(2);
let expected = vec![
(
a,
139_682_966,
RankScoreKind::Rrf,
SearchSignals {
vector_similarity: Some(DeterministicScore::from_raw(1_i64 << 31)),
keyword_score: Some(DeterministicScore::from_raw(1_i64 << 30)),
},
),
(
b,
139_682_966,
RankScoreKind::Rrf,
SearchSignals {
vector_similarity: Some(DeterministicScore::from_raw(1_i64 << 32)),
keyword_score: Some(DeterministicScore::from_raw(3_i64 << 30)),
},
),
];
for (text_ids, vector_ids) in [([a, b], [b, a]), ([b, a], [a, b])] {
for _ in 0..4 {
let text = text_ids
.into_iter()
.map(|id| text_hit(id, if id == a { 0.25 } else { 0.75 }, "candidate"))
.collect();
let vector = vector_ids
.into_iter()
.map(|id| vector_hit(id, if id == a { 0.5 } else { 1.0 }))
.collect();
let hits = rt
.fuse_with_strategy(text, vector, &FusionStrategy::Rrf { k: 60 }, 10)
.await
.unwrap();
assert_eq!(hits.len(), 2);
assert_eq!(hits[0].score, hits[1].score);
assert!(hits.iter().all(|hit| hit.source == SearchSource::Both));
let snapshot: Vec<_> = hits
.iter()
.map(|hit| {
(
hit.entity_id,
hit.score.to_raw(),
hit.rank_score_kind,
hit.signals,
)
})
.collect();
assert_eq!(snapshot, expected);
}
}
}
#[tokio::test]
async fn fusion_evidence_duplicate_selection_follows_strategy() {
let rt = evidence_runtime();
let id = Uuid::from_u128(1);
for (strategy, keyword_raw) in [
(FusionStrategy::rrf(), 1_i64 << 30),
(FusionStrategy::Union, 3_i64 << 30),
(FusionStrategy::weighted(vec![0.5, 0.5]), 3_i64 << 30),
] {
let hits = rt
.fuse_with_strategy(
vec![text_hit(id, 0.25, "first"), text_hit(id, 0.75, "second")],
vec![vector_hit(id, 0.5)],
&strategy,
10,
)
.await
.unwrap();
assert_eq!(hits.len(), 1);
assert_eq!(
hits[0].signals,
SearchSignals {
vector_similarity: Some(DeterministicScore::from_raw(1_i64 << 31)),
keyword_score: Some(DeterministicScore::from_raw(keyword_raw)),
}
);
assert_eq!(hits[0].title.as_deref(), Some("first"));
}
}
#[tokio::test]
async fn custom_fusion_evidence_uses_declared_kind() {
let rt = evidence_runtime();
let id = Uuid::from_u128(1);
rt.register_fusion_strategy("invert", Arc::new(InvertScoreExecutor));
let strategy =
FusionStrategy::try_custom("invert".into(), serde_json::Value::Null).unwrap();
let hits = rt
.fuse_with_strategy(vec![text_hit(id, 0.75, "candidate")], vec![], &strategy, 10)
.await
.unwrap();
assert_eq!(hits.len(), 1);
assert_eq!(hits[0].score.to_raw(), 1_i64 << 30);
assert_eq!(hits[0].rank_score_kind, RankScoreKind::Weighted);
assert_eq!(
hits[0].signals,
SearchSignals {
vector_similarity: None,
keyword_score: Some(DeterministicScore::from_raw(3_i64 << 30)),
}
);
}
fn cosine_fixture_vector(dimensions: usize, x: f32, y: f32) -> Vec<f32> {
let mut vector = vec![0.0; dimensions];
vector[0] = x;
vector[1] = y;
vector
}
async fn stale_full_prefix_fixture() -> (
KhiveRuntime,
NamespaceToken,
&'static str,
Vec<f32>,
Vec<TextSearchHit>,
Vec<VectorSearchHit>,
HashSet<Uuid>,
) {
let model = EmbeddingModel::AllMiniLmL6V2;
let dimensions = model.dimensions();
let rt = KhiveRuntime::new(RuntimeConfig {
db_path: None,
embedding_model: Some(model),
additional_embedding_models: vec![],
..RuntimeConfig::default()
})
.unwrap();
let tok = NamespaceToken::local();
let query_text = "fusionrefillterm";
let query_vector = cosine_fixture_vector(dimensions, 1.0, 0.0);
let common_stale_a = Uuid::from_u128(1);
let common_stale_b = Uuid::from_u128(2);
let text_only_stale = Uuid::from_u128(3);
let vector_only_stale = Uuid::from_u128(4);
let live_text = Entity::new("local", "concept", "live text candidate");
let live_vector = Entity::new("local", "concept", "live vector candidate");
rt.entities(&tok)
.unwrap()
.upsert_entities(vec![live_text.clone(), live_vector.clone()])
.await
.unwrap();
let document = |subject_id, repetitions: usize| TextDocument {
subject_id,
kind: SubstrateKind::Entity,
record_kind: None,
namespace: "local".to_string(),
title: None,
body: std::iter::repeat_n(query_text, repetitions)
.collect::<Vec<_>>()
.join(" "),
tags: vec![],
metadata: None,
updated_at: Utc::now(),
};
rt.text(&tok)
.unwrap()
.upsert_documents(vec![
document(common_stale_a, 12),
document(common_stale_b, 8),
document(text_only_stale, 4),
document(live_text.id, 1),
])
.await
.unwrap();
let vectors = rt.vectors(&tok).unwrap();
for (id, vector) in [
(common_stale_a, cosine_fixture_vector(dimensions, 1.0, 0.0)),
(common_stale_b, cosine_fixture_vector(dimensions, 0.8, 0.6)),
(
vector_only_stale,
cosine_fixture_vector(dimensions, 0.5, 0.866_025_4),
),
(live_vector.id, cosine_fixture_vector(dimensions, -1.0, 0.0)),
] {
vectors
.insert(
id,
SubstrateKind::Entity,
"local",
"entity.body",
vec![vector],
)
.await
.unwrap();
}
let text_hits = rt
.text(&tok)
.unwrap()
.search(TextSearchRequest {
query: query_text.to_string(),
mode: TextQueryMode::Plain,
filter: Some(TextFilter {
namespaces: vec!["local".to_string()],
..TextFilter::default()
}),
top_k: CANDIDATE_MULTIPLIER,
snippet_chars: 0,
})
.await
.unwrap();
let vector_hits = vectors
.search(VectorSearchRequest {
query_vectors: vec![query_vector.clone()],
top_k: CANDIDATE_MULTIPLIER,
namespace: Some("local".to_string()),
kind: Some(SubstrateKind::Entity),
embedding_model: None,
filter: None,
backend_hints: None,
})
.await
.unwrap();
assert_eq!(text_hits.len(), CANDIDATE_MULTIPLIER as usize);
assert_eq!(vector_hits.len(), CANDIDATE_MULTIPLIER as usize);
let live = HashSet::from([live_text.id, live_vector.id]);
(
rt,
tok,
query_text,
query_vector,
text_hits,
vector_hits,
live,
)
}
#[tokio::test]
async fn rrf_custom_k_differs_from_k60() {
let rt = KhiveRuntime::memory().unwrap();
let a = Uuid::new_v4();
let b = Uuid::new_v4();
let text = vec![text_hit(a, 0.9, "a"), text_hit(b, 0.1, "b")];
let hits_k1 = rt
.fuse_with_strategy(text.clone(), vec![], &FusionStrategy::Rrf { k: 1 }, 10)
.await
.unwrap();
let hits_k60 = rt
.fuse_with_strategy(text, vec![], &FusionStrategy::Rrf { k: 60 }, 10)
.await
.unwrap();
assert_eq!(hits_k1[0].entity_id, a);
assert_eq!(hits_k60[0].entity_id, a);
assert!(hits_k1[0].score > hits_k60[0].score);
}
#[tokio::test]
async fn weighted_ordering_depends_on_weights() {
let rt = KhiveRuntime::memory().unwrap();
let a = Uuid::new_v4();
let b = Uuid::new_v4();
let text = vec![text_hit(a, 0.9, "a"), text_hit(b, 0.1, "b")];
let vec_hits = vec![vector_hit(b, 0.9), vector_hit(a, 0.1)];
let heavy_vector = rt
.fuse_with_strategy(
text.clone(),
vec_hits.clone(),
&FusionStrategy::Weighted {
weights: vec![0.7, 0.3],
},
10,
)
.await
.unwrap();
let heavy_keyword = rt
.fuse_with_strategy(
text,
vec_hits,
&FusionStrategy::Weighted {
weights: vec![0.3, 0.7],
},
10,
)
.await
.unwrap();
assert_eq!(heavy_vector[0].entity_id, b);
assert_eq!(heavy_keyword[0].entity_id, a);
}
#[tokio::test]
async fn weighted_scale_invariant() {
let rt = KhiveRuntime::memory().unwrap();
let a = Uuid::new_v4();
let b = Uuid::new_v4();
let text = vec![text_hit(a, 0.9, "a"), text_hit(b, 0.1, "b")];
let vec_hits = vec![vector_hit(b, 0.9), vector_hit(a, 0.1)];
let w1 = rt
.fuse_with_strategy(
text.clone(),
vec_hits.clone(),
&FusionStrategy::Weighted {
weights: vec![0.7, 0.3],
},
10,
)
.await
.unwrap();
let w2 = rt
.fuse_with_strategy(
text,
vec_hits,
&FusionStrategy::Weighted {
weights: vec![7.0, 3.0],
},
10,
)
.await
.unwrap();
assert_eq!(w1[0].entity_id, w2[0].entity_id);
assert_eq!(w1[1].entity_id, w2[1].entity_id);
let diff = (w1[0].score.to_f64() - w2[0].score.to_f64()).abs();
assert!(diff < 1e-9, "scores differ by {diff}");
}
#[tokio::test]
async fn weighted_zero_weights_equal_fallback() {
let rt = KhiveRuntime::memory().unwrap();
let a = Uuid::new_v4();
let b = Uuid::new_v4();
let text = vec![text_hit(a, 0.9, "a"), text_hit(b, 0.1, "b")];
let vec_hits = vec![vector_hit(a, 0.9), vector_hit(b, 0.1)];
let hits = rt
.fuse_with_strategy(
text,
vec_hits,
&FusionStrategy::Weighted {
weights: vec![0.0, 0.0],
},
10,
)
.await
.unwrap();
assert_eq!(hits[0].entity_id, a);
}
#[tokio::test]
async fn weighted_negative_weight_clamped() {
let rt = KhiveRuntime::memory().unwrap();
let a = Uuid::new_v4();
let text = vec![text_hit(a, 0.9, "a")];
let hits = rt
.fuse_with_strategy(
text,
vec![],
&FusionStrategy::Weighted {
weights: vec![-0.5, 1.0],
},
10,
)
.await
.unwrap();
assert_eq!(hits.len(), 1);
assert_eq!(hits[0].entity_id, a);
}
#[tokio::test]
async fn weighted_empty_arm_keeps_canonical_position() {
let rt = KhiveRuntime::memory().unwrap();
let text_only = Uuid::new_v4();
let hits = rt
.fuse_with_strategy(
vec![text_hit(text_only, 0.9, "text")],
vec![],
&FusionStrategy::Weighted {
weights: vec![1.0, 0.0],
},
10,
)
.await
.unwrap();
assert!(
hits.is_empty(),
"dropping the empty vector arm would incorrectly rebind text to its weight"
);
}
#[tokio::test]
async fn union_max_score_per_entity() {
let rt = KhiveRuntime::memory().unwrap();
let a = Uuid::new_v4();
let text = vec![text_hit(a, 0.3, "a")];
let vec_hits = vec![vector_hit(a, 0.9)];
let hits = rt
.fuse_with_strategy(text, vec_hits, &FusionStrategy::Union, 10)
.await
.unwrap();
assert_eq!(hits.len(), 1);
assert!((hits[0].score.to_f64() - 0.9).abs() < 1e-6);
assert_eq!(hits[0].source, SearchSource::Both);
}
#[tokio::test]
async fn vector_only_drops_text() {
let rt = KhiveRuntime::memory().unwrap();
let a = Uuid::new_v4();
let b = Uuid::new_v4();
let text = vec![text_hit(b, 0.9, "b")];
let vec_hits = vec![vector_hit(a, 0.8)];
let hits = rt
.fuse_with_strategy(text, vec_hits, &FusionStrategy::VectorOnly, 10)
.await
.unwrap();
assert_eq!(hits.len(), 1);
assert_eq!(hits[0].entity_id, a);
assert_eq!(hits[0].source, SearchSource::Vector);
assert!(hits[0].title.is_none());
}
#[tokio::test]
async fn keyword_only_drops_vector() {
let rt = KhiveRuntime::memory().unwrap();
let text_id = Uuid::new_v4();
let vector_id = Uuid::new_v4();
let hits = rt
.fuse_with_strategy(
vec![text_hit(text_id, 0.8, "text")],
vec![vector_hit(vector_id, 0.9)],
&FusionStrategy::KeywordOnly,
10,
)
.await
.unwrap();
assert_eq!(hits.len(), 1);
assert_eq!(hits[0].entity_id, text_id);
assert_eq!(hits[0].source, SearchSource::Text);
}
struct ReverseOrderExecutor;
#[async_trait::async_trait]
impl FusionExecutor for ReverseOrderExecutor {
fn rank_score_kind(&self) -> RankScoreKind {
RankScoreKind::Union
}
async fn fuse(
&self,
streams: Vec<CandidateStream>,
_params: &serde_json::Value,
_limit: usize,
) -> RuntimeResult<Vec<RankedHit>> {
let mut flat: Vec<_> = streams.into_iter().flatten().collect();
flat.reverse();
Ok(flat)
}
}
struct InvertScoreExecutor;
#[async_trait::async_trait]
impl FusionExecutor for InvertScoreExecutor {
fn rank_score_kind(&self) -> RankScoreKind {
RankScoreKind::Weighted
}
async fn fuse(
&self,
streams: Vec<CandidateStream>,
_params: &serde_json::Value,
_limit: usize,
) -> RuntimeResult<Vec<RankedHit>> {
Ok(streams
.into_iter()
.flatten()
.map(|(id, score)| (id, DeterministicScore::from_f64(1.0 - score.to_f64())))
.collect())
}
}
struct EqualScoreExecutor;
#[async_trait::async_trait]
impl FusionExecutor for EqualScoreExecutor {
fn rank_score_kind(&self) -> RankScoreKind {
RankScoreKind::Weighted
}
async fn fuse(
&self,
streams: Vec<CandidateStream>,
_params: &serde_json::Value,
_limit: usize,
) -> RuntimeResult<Vec<RankedHit>> {
Ok(streams
.into_iter()
.flatten()
.map(|(id, _)| (id, DeterministicScore::from_f64(1.0)))
.collect())
}
}
#[tokio::test]
async fn custom_strategy_dispatches_through_executor_and_differs_from_rrf() {
let rt = KhiveRuntime::memory().unwrap();
let a = Uuid::new_v4();
let b = Uuid::new_v4();
let text = vec![text_hit(a, 0.9, "a"), text_hit(b, 0.5, "b")];
rt.register_fusion_strategy("invert", Arc::new(InvertScoreExecutor));
let strategy =
FusionStrategy::try_custom("invert".to_string(), serde_json::Value::Null).unwrap();
let custom = rt
.fuse_with_strategy(text.clone(), vec![], &strategy, 10)
.await
.unwrap();
let rrf = rt
.fuse_with_strategy(text, vec![], &FusionStrategy::rrf(), 10)
.await
.unwrap();
let custom_ids: Vec<_> = custom.iter().map(|h| h.entity_id).collect();
let rrf_ids: Vec<_> = rrf.iter().map(|h| h.entity_id).collect();
assert_ne!(
custom_ids, rrf_ids,
"custom and RRF must yield different orderings on this fixture"
);
}
#[tokio::test]
async fn custom_strategy_unknown_name_fails_closed() {
let rt = KhiveRuntime::memory().unwrap();
let a = Uuid::new_v4();
let text = vec![text_hit(a, 0.9, "a")];
let strategy =
FusionStrategy::try_custom("nonexistent".to_string(), serde_json::Value::Null).unwrap();
let result = rt.fuse_with_strategy(text, vec![], &strategy, 10).await;
assert!(matches!(
result,
Err(RuntimeError::UnknownFusionStrategy(name)) if name == "nonexistent"
));
}
#[tokio::test]
async fn custom_strategy_unknown_name_fails_closed_even_on_empty_input() {
let rt = KhiveRuntime::memory().unwrap();
let strategy =
FusionStrategy::try_custom("nonexistent".to_string(), serde_json::Value::Null).unwrap();
let result = rt.fuse_with_strategy(vec![], vec![], &strategy, 10).await;
assert!(matches!(
result,
Err(RuntimeError::UnknownFusionStrategy(name)) if name == "nonexistent"
));
}
#[tokio::test]
async fn custom_strategy_registered_name_empty_input_returns_ok_empty() {
let rt = KhiveRuntime::memory().unwrap();
rt.register_fusion_strategy("reverse", Arc::new(ReverseOrderExecutor));
let strategy =
FusionStrategy::try_custom("reverse".to_string(), serde_json::Value::Null).unwrap();
let result = rt
.fuse_with_strategy(vec![], vec![], &strategy, 10)
.await
.unwrap();
assert!(result.is_empty());
}
#[tokio::test]
async fn registered_custom_strategy_leaves_default_path_unaffected() {
let rt = KhiveRuntime::memory().unwrap();
let a = Uuid::new_v4();
let b = Uuid::new_v4();
let text = vec![text_hit(a, 0.9, "a"), text_hit(b, 0.5, "b")];
rt.register_fusion_strategy("reverse", Arc::new(ReverseOrderExecutor));
let via_rt_with_registration = rt
.fuse_with_strategy(text.clone(), vec![], &FusionStrategy::rrf(), 10)
.await
.unwrap();
let rt2 = KhiveRuntime::memory().unwrap();
let via_rt_without_registration = rt2
.fuse_with_strategy(text, vec![], &FusionStrategy::rrf(), 10)
.await
.unwrap();
let ids_with: Vec<_> = via_rt_with_registration
.iter()
.map(|h| h.entity_id)
.collect();
let ids_without: Vec<_> = via_rt_without_registration
.iter()
.map(|h| h.entity_id)
.collect();
assert_eq!(ids_with, ids_without);
}
#[tokio::test]
async fn custom_executor_output_is_sorted_by_canonical_comparator() {
let rt = KhiveRuntime::memory().unwrap();
let ids: Vec<Uuid> = vec![Uuid::from_u128(3), Uuid::from_u128(1), Uuid::from_u128(2)];
let text: Vec<TextSearchHit> = ids.iter().map(|&id| text_hit(id, 0.5, "tied")).collect();
rt.register_fusion_strategy("equal_score", Arc::new(EqualScoreExecutor));
let strategy =
FusionStrategy::try_custom("equal_score".to_string(), serde_json::Value::Null).unwrap();
let hits = rt
.fuse_with_strategy(text, vec![], &strategy, 10)
.await
.unwrap();
let mut expected = ids.clone();
expected.sort();
let actual: Vec<_> = hits.iter().map(|h| h.entity_id).collect();
assert_eq!(
actual, expected,
"equal-score executor output must be tie-broken by ascending ID"
);
}
#[test]
fn default_strategy_is_rrf_k60() {
assert_eq!(FusionStrategy::default(), FusionStrategy::Rrf { k: 60 });
}
#[tokio::test]
async fn hybrid_default_rrf_alive_filter_refills_below_complete_four_x_prefix() {
let (rt, tok, query_text, query_vector, _text_hits, _vector_hits, live) =
stale_full_prefix_fixture().await;
let hits = rt
.hybrid_search(
&tok,
query_text,
Some(query_vector),
1,
None,
None,
&[],
None,
)
.await
.unwrap();
assert_eq!(hits.len(), 1);
assert!(live.contains(&hits[0].entity_id));
}
#[test]
fn serde_roundtrip() {
let cases = vec![
FusionStrategy::Rrf { k: 60 },
FusionStrategy::Rrf { k: 20 },
FusionStrategy::Weighted {
weights: vec![0.7, 0.3],
},
FusionStrategy::Union,
FusionStrategy::VectorOnly,
FusionStrategy::KeywordOnly,
];
for strategy in cases {
let json = serde_json::to_string(&strategy).expect("serialize");
let back: FusionStrategy = serde_json::from_str(&json).expect("deserialize");
assert_eq!(strategy, back, "roundtrip failed for {json}");
}
}
}