1use std::collections::{HashMap, HashSet};
4
5use lattice_embed::{EmbeddingModel, MAX_TEXT_BYTES};
6use uuid::Uuid;
7
8use crate::config::{parse_embedding_model_alias, sanitize_key};
9use crate::curation::note_fts_document;
10use crate::embedder_registry::with_embedding_admission;
11use crate::error::{RuntimeError, RuntimeResult};
12use crate::runtime::{KhiveRuntime, NamespaceToken};
13use khive_retrieval::hybrid::{combine_leg_first_appearance, fuse_labelled, HitLabel};
14use khive_score::DeterministicScore;
15use khive_storage::types::{
16 PageRequest, TextFilter, TextQueryMode, TextSearchHit, TextSearchRequest, VectorRecord,
17 VectorSearchHit, VectorSearchRequest,
18};
19use khive_storage::ContentRef;
20use khive_storage::EntityFilter;
21use khive_types::SubstrateKind;
22
23pub use khive_retrieval::{
24 HybridSearchOutcome, RankScoreKind, SearchHit, SearchSignals, SearchSource,
25};
26
27pub(crate) const EMBEDDING_BATCH_PAGE_SIZE: usize = 256;
29
30#[cfg(any(test, feature = "fault-injection"))]
32std::thread_local! {
33 static BACKFILL_READER_FAIL: std::cell::Cell<bool> = const { std::cell::Cell::new(false) };
34}
35
36#[cfg(any(test, feature = "fault-injection"))]
42pub fn arm_backfill_reader_fail() {
43 BACKFILL_READER_FAIL.with(|c| c.set(true));
44}
45
46const RRF_K: usize = 10;
53
54const CANDIDATE_MULTIPLIER: u32 = 4;
56
57pub const EMBEDDING_INPUT_TRUNCATED_WARNING: &str =
59 "embedding input was truncated to the embedder maximum; full content was stored unchanged";
60
61#[derive(Clone, Debug)]
64pub struct DocumentEmbeddingOutcome {
65 pub model_name: String,
66 pub vector: Vec<f32>,
67 pub prepared_text_fingerprint: Option<ContentRef>,
69 pub source_bytes: usize,
70 pub embedded_bytes: usize,
71 pub truncated: bool,
72}
73
74#[derive(Clone, Debug, Default, PartialEq, Eq, serde::Deserialize, serde::Serialize)]
76pub struct EmbeddingTruncationReport {
77 pub truncated: u64,
78 pub discarded_bytes: u64,
79}
80
81impl EmbeddingTruncationReport {
82 pub fn observe(&mut self, outcome: &DocumentEmbeddingOutcome) {
83 if outcome.truncated {
84 self.truncated += 1;
85 self.discarded_bytes +=
86 outcome.source_bytes.saturating_sub(outcome.embedded_bytes) as u64;
87 }
88 }
89
90 #[must_use]
91 pub const fn any_truncated(&self) -> bool {
92 self.truncated > 0
93 }
94
95 pub fn merge(&mut self, other: Self) {
97 self.truncated = self.truncated.saturating_add(other.truncated);
98 self.discarded_bytes = self.discarded_bytes.saturating_add(other.discarded_bytes);
99 }
100}
101
102pub fn document_embedding_budget(model_name: &str) -> usize {
104 parse_embedding_model_alias(model_name)
105 .and_then(|model| model.document_instruction())
106 .map_or(MAX_TEXT_BYTES, |prefix| {
107 MAX_TEXT_BYTES.saturating_sub(prefix.len())
108 })
109}
110
111pub fn bounded_embedding_input(text: &str, max_bytes: usize) -> (&str, bool) {
113 if text.len() <= max_bytes {
114 return (text, false);
115 }
116
117 let end = text
118 .char_indices()
119 .map(|(index, _)| index)
120 .take_while(|index| *index <= max_bytes)
121 .last()
122 .unwrap_or(0);
123 (&text[..end], true)
124}
125
126fn prepared_document_fingerprint(text: &str, model: EmbeddingModel) -> ContentRef {
127 match model.document_instruction() {
128 Some(prefix) => VectorRecord::fingerprint_text(&format!("{prefix}{text}")),
129 None => VectorRecord::fingerprint_text(text),
130 }
131}
132
133impl KhiveRuntime {
134 fn require_default_embedder(&self) -> RuntimeResult<&str> {
135 let model_name = self.default_embedder_name();
136 if model_name.is_empty() {
137 return Err(RuntimeError::Unconfigured("embedding_model".into()));
138 }
139 Ok(model_name)
140 }
141
142 pub async fn embed(&self, text: &str) -> RuntimeResult<Vec<f32>> {
147 let model_name = self.require_default_embedder()?;
148 self.embed_with_model(model_name, text).await
149 }
150
151 pub async fn embed_with_model(&self, model_name: &str, text: &str) -> RuntimeResult<Vec<f32>> {
167 let model = parse_embedding_model_alias(model_name);
168 let service = self.embedder(model_name).await?;
169 let emb_model = model.unwrap_or_default();
170 crate::usage::count(crate::usage::UsageUnit::EmbedCalls, 1);
174 let out = with_embedding_admission(service.embed_one(text, emb_model)).await;
175 out
176 }
177
178 pub async fn embed_document_with_model_outcome(
197 &self,
198 model_name: &str,
199 text: &str,
200 ) -> RuntimeResult<DocumentEmbeddingOutcome> {
201 self.embed_document_with_model_outcome_inner(None, model_name, text)
202 .await
203 }
204
205 pub(crate) async fn embed_document_with_model_outcome_for_token(
206 &self,
207 token: &NamespaceToken,
208 model_name: &str,
209 text: &str,
210 ) -> RuntimeResult<DocumentEmbeddingOutcome> {
211 self.embed_document_with_model_outcome_inner(Some(token), model_name, text)
212 .await
213 }
214
215 async fn embed_document_with_model_outcome_inner(
216 &self,
217 token: Option<&NamespaceToken>,
218 model_name: &str,
219 text: &str,
220 ) -> RuntimeResult<DocumentEmbeddingOutcome> {
221 let model = parse_embedding_model_alias(model_name);
222 let (service, audited_document_preparation) = self
223 .embedder_with_input_attestation(model_name, token)
224 .await?;
225 let emb_model = model.unwrap_or_default();
226 let source_bytes = text.len();
227 let (text, truncated) =
228 bounded_embedding_input(text, document_embedding_budget(model_name));
229 let embedded_bytes = text.len();
230 crate::usage::count(crate::usage::UsageUnit::EmbedCalls, 1);
232 let embeddings =
233 with_embedding_admission(service.embed_passage(&[text.to_string()], emb_model)).await;
234 let mut vectors = embeddings?;
235 if vectors.len() != 1 {
236 return Err(RuntimeError::Internal(format!(
237 "embed_passage returned {} vectors for 1 input",
238 vectors.len()
239 )));
240 }
241 let out = vectors.pop().expect("checked len == 1 above");
242 Ok(DocumentEmbeddingOutcome {
243 model_name: model_name.to_owned(),
244 vector: out,
245 prepared_text_fingerprint: audited_document_preparation
246 .then(|| prepared_document_fingerprint(text, emb_model)),
247 source_bytes,
248 embedded_bytes,
249 truncated,
250 })
251 }
252
253 pub async fn embed_query_with_model(
265 &self,
266 model_name: &str,
267 text: &str,
268 ) -> RuntimeResult<Vec<f32>> {
269 self.embed_query_with_model_inner(None, model_name, text)
270 .await
271 }
272
273 pub(crate) async fn embed_query_with_model_for_token(
274 &self,
275 token: &NamespaceToken,
276 model_name: &str,
277 text: &str,
278 ) -> RuntimeResult<Vec<f32>> {
279 self.embed_query_with_model_inner(Some(token), model_name, text)
280 .await
281 }
282
283 async fn embed_query_with_model_inner(
284 &self,
285 token: Option<&NamespaceToken>,
286 model_name: &str,
287 text: &str,
288 ) -> RuntimeResult<Vec<f32>> {
289 let model = parse_embedding_model_alias(model_name);
290 let service = match token {
291 Some(token) => self.embedder_with_token(token, model_name).await?,
292 None => self.embedder(model_name).await?,
293 };
294 let texts = [text.to_string()];
295 let emb_model = model.unwrap_or_default();
296 crate::usage::count(crate::usage::UsageUnit::EmbedCalls, 1);
298 let embeddings = match emb_model {
299 EmbeddingModel::BgeSmallEnV15
300 | EmbeddingModel::BgeBaseEnV15
301 | EmbeddingModel::BgeLargeEnV15 => {
302 with_embedding_admission(service.embed(&texts, emb_model)).await
303 }
304 _ => with_embedding_admission(service.embed_query(&texts, emb_model)).await,
305 };
306 let out = embeddings?
307 .into_iter()
308 .next()
309 .ok_or_else(|| RuntimeError::Internal("embed_query returned empty vec".into()))?;
310 Ok(out)
311 }
312
313 pub async fn embed_document(&self, text: &str) -> RuntimeResult<Vec<f32>> {
322 let outcome = self.embed_document_outcome(text).await?;
323 if outcome.truncated {
324 return Err(RuntimeError::InvalidInput(format!(
325 "embedding input truncated from {} to {} bytes; use embed_document_outcome to inspect the bounded vector",
326 outcome.source_bytes, outcome.embedded_bytes
327 )));
328 }
329 Ok(outcome.vector)
330 }
331
332 pub async fn embed_document_outcome(
334 &self,
335 text: &str,
336 ) -> RuntimeResult<DocumentEmbeddingOutcome> {
337 let model_name = self.require_default_embedder()?;
338 self.embed_document_with_model_outcome(model_name, text)
339 .await
340 }
341
342 pub async fn embed_query(&self, text: &str) -> RuntimeResult<Vec<f32>> {
349 let model_name = self.require_default_embedder()?;
350 self.embed_query_with_model(model_name, text).await
351 }
352
353 async fn embed_query_for_token(
354 &self,
355 token: &NamespaceToken,
356 text: &str,
357 ) -> RuntimeResult<Vec<f32>> {
358 let model_name = self.require_default_embedder()?;
359 self.embed_query_with_model_for_token(token, model_name, text)
360 .await
361 }
362
363 pub async fn embed_batch(&self, texts: &[String]) -> RuntimeResult<Vec<Vec<f32>>> {
371 if texts.is_empty() {
372 return Ok(vec![]);
373 }
374 let model_name = self.require_default_embedder()?;
375 self.embed_batch_with_model(model_name, texts).await
376 }
377
378 pub async fn embed_batch_with_model(
383 &self,
384 model_name: &str,
385 texts: &[String],
386 ) -> RuntimeResult<Vec<Vec<f32>>> {
387 if texts.is_empty() {
388 return Ok(vec![]);
389 }
390 let model = parse_embedding_model_alias(model_name);
391 let service = self.embedder(model_name).await?;
392 let emb_model = model.unwrap_or_default();
393 let out = with_embedding_admission(service.embed(texts, emb_model)).await;
394 crate::usage::count(crate::usage::UsageUnit::EmbedCalls, texts.len() as u64);
395 out
396 }
397
398 pub async fn embed_document_batch_with_model(
412 &self,
413 model_name: &str,
414 texts: &[String],
415 ) -> RuntimeResult<Vec<Vec<f32>>> {
416 if texts.is_empty() {
417 return Ok(vec![]);
418 }
419 let outcomes = self
420 .embed_document_batch_with_model_outcomes(model_name, texts)
421 .await?;
422 let mut report = EmbeddingTruncationReport::default();
423 for outcome in &outcomes {
424 report.observe(outcome);
425 }
426 if report.any_truncated() {
427 return Err(RuntimeError::InvalidInput(format!(
428 "embedding input truncated for {} documents ({} discarded bytes); use embed_document_batch_with_model_outcomes to inspect the bounded vectors",
429 report.truncated, report.discarded_bytes
430 )));
431 }
432 Ok(outcomes.into_iter().map(|outcome| outcome.vector).collect())
433 }
434
435 pub async fn embed_document_batch_with_model_outcomes(
436 &self,
437 model_name: &str,
438 texts: &[String],
439 ) -> RuntimeResult<Vec<DocumentEmbeddingOutcome>> {
440 self.embed_document_batch_with_model_outcomes_inner(None, model_name, texts)
441 .await
442 }
443
444 pub(crate) async fn embed_document_batch_with_model_outcomes_for_token(
445 &self,
446 token: &NamespaceToken,
447 model_name: &str,
448 texts: &[String],
449 ) -> RuntimeResult<Vec<DocumentEmbeddingOutcome>> {
450 self.embed_document_batch_with_model_outcomes_inner(Some(token), model_name, texts)
451 .await
452 }
453
454 async fn embed_document_batch_with_model_outcomes_inner(
455 &self,
456 token: Option<&NamespaceToken>,
457 model_name: &str,
458 texts: &[String],
459 ) -> RuntimeResult<Vec<DocumentEmbeddingOutcome>> {
460 if texts.is_empty() {
461 return Ok(vec![]);
462 }
463 let model = parse_embedding_model_alias(model_name);
464 let (service, audited_document_preparation) = self
465 .embedder_with_input_attestation(model_name, token)
466 .await?;
467 let emb_model = model.unwrap_or_default();
468 let budget = document_embedding_budget(model_name);
469 if token.is_some() {
470 crate::usage::count(crate::usage::UsageUnit::EmbedCalls, texts.len() as u64);
472 }
473 let out = if texts.iter().all(|text| text.len() <= budget) {
474 with_embedding_admission(service.embed_passage(texts, emb_model)).await
475 } else {
476 let bounded_texts: Vec<String> = texts
477 .iter()
478 .map(|text| bounded_embedding_input(text, budget).0.to_owned())
479 .collect();
480 with_embedding_admission(service.embed_passage(&bounded_texts, emb_model)).await
481 };
482 if token.is_none() {
483 crate::usage::count(crate::usage::UsageUnit::EmbedCalls, texts.len() as u64);
484 }
485 let vectors = out?;
486 if vectors.len() != texts.len() {
487 return Err(RuntimeError::Internal(format!(
488 "embed_passage returned {} vectors for {} inputs",
489 vectors.len(),
490 texts.len()
491 )));
492 }
493 Ok(texts
494 .iter()
495 .zip(vectors)
496 .map(|(text, vector)| {
497 let (bounded, truncated) = bounded_embedding_input(text, budget);
498 DocumentEmbeddingOutcome {
499 model_name: model_name.to_owned(),
500 vector,
501 prepared_text_fingerprint: audited_document_preparation
502 .then(|| prepared_document_fingerprint(bounded, emb_model)),
503 source_bytes: text.len(),
504 embedded_bytes: bounded.len(),
505 truncated,
506 }
507 })
508 .collect())
509 }
510
511 pub async fn embed_document_batch(&self, texts: &[String]) -> RuntimeResult<Vec<Vec<f32>>> {
520 if texts.is_empty() {
521 return Ok(vec![]);
522 }
523 let model_name = self.require_default_embedder()?;
524 self.embed_document_batch_with_model(model_name, texts)
525 .await
526 }
527
528 pub async fn embed_document_batch_outcomes(
530 &self,
531 texts: &[String],
532 ) -> RuntimeResult<Vec<DocumentEmbeddingOutcome>> {
533 if texts.is_empty() {
534 return Ok(vec![]);
535 }
536 let model_name = self.require_default_embedder()?;
537 self.embed_document_batch_with_model_outcomes(model_name, texts)
538 .await
539 }
540
541 pub async fn embed_query_batch_with_model(
548 &self,
549 model_name: &str,
550 texts: &[String],
551 ) -> RuntimeResult<Vec<Vec<f32>>> {
552 if texts.is_empty() {
553 return Ok(vec![]);
554 }
555 let model = parse_embedding_model_alias(model_name);
556 let service = self.embedder(model_name).await?;
557 let emb_model = model.unwrap_or_default();
558 let out = match emb_model {
559 EmbeddingModel::BgeSmallEnV15
560 | EmbeddingModel::BgeBaseEnV15
561 | EmbeddingModel::BgeLargeEnV15 => {
562 with_embedding_admission(service.embed(texts, emb_model)).await
563 }
564 _ => with_embedding_admission(service.embed_query(texts, emb_model)).await,
565 };
566 crate::usage::count(crate::usage::UsageUnit::EmbedCalls, texts.len() as u64);
567 out
568 }
569
570 pub async fn vector_search(
576 &self,
577 token: &NamespaceToken,
578 query_embedding: Option<Vec<f32>>,
579 query_text: Option<&str>,
580 top_k: u32,
581 kind: Option<SubstrateKind>,
582 ) -> RuntimeResult<Vec<VectorSearchHit>> {
583 let embedding = match query_embedding {
584 Some(vec) => vec,
585 None => {
586 let text = query_text.ok_or_else(|| {
587 RuntimeError::InvalidInput(
588 "vector search requires query_embedding or query_text".into(),
589 )
590 })?;
591 if text.trim().is_empty() {
592 return Err(RuntimeError::InvalidInput(
593 "query_text must not be empty".into(),
594 ));
595 }
596 self.embed_query_for_token(token, text).await?
597 }
598 };
599
600 let ns = token.namespace().as_str().to_owned();
601 let hits = self
602 .vectors(token)?
603 .search(VectorSearchRequest {
604 query_vectors: vec![embedding],
605 top_k,
606 namespace: Some(ns),
607 kind,
608 embedding_model: None,
609 filter: None,
610 backend_hints: None,
611 })
612 .await;
613 crate::usage::count(crate::usage::UsageUnit::VectorPasses, 1);
614 hits.map_err(RuntimeError::from)
615 }
616
617 pub(crate) async fn note_search_vector_search(
621 &self,
622 token: &NamespaceToken,
623 query_embedding: Option<Vec<f32>>,
624 query_text: &str,
625 top_k: u32,
626 ) -> RuntimeResult<Vec<VectorSearchHit>> {
627 let embedding = match query_embedding {
628 Some(embedding) => embedding,
629 None => self.embed_query_for_token(token, query_text).await?,
630 };
631 let model = self.default_embedder_name();
632 if !model.is_empty() {
633 if let Some(provider) = self.note_search_ann_provider()? {
634 if let Some(hits) = provider.search(token, model, &embedding, top_k).await? {
635 crate::note_search_ann::record_ann_route();
636 crate::usage::count(crate::usage::UsageUnit::VectorPasses, 1);
637 return Ok(hits);
638 }
639 }
640 }
641 crate::note_search_ann::record_fallback_route();
642 self.vector_search(
643 token,
644 Some(embedding),
645 None,
646 top_k,
647 Some(SubstrateKind::Note),
648 )
649 .await
650 }
651
652 #[allow(clippy::too_many_arguments)]
698 pub async fn hybrid_search(
699 &self,
700 token: &NamespaceToken,
701 query_text: &str,
702 query_vector: Option<Vec<f32>>,
703 limit: u32,
704 entity_kind: Option<&str>,
705 entity_type: Option<&str>,
706 tags_any: &[String],
707 properties_filter: Option<&serde_json::Value>,
708 ) -> RuntimeResult<Vec<SearchHit>> {
709 self.hybrid_search_with_text_mode(
710 token,
711 query_text,
712 query_vector,
713 limit,
714 entity_kind,
715 entity_type,
716 tags_any,
717 properties_filter,
718 TextQueryMode::Plain,
719 )
720 .await
721 }
722
723 #[allow(clippy::too_many_arguments)]
725 pub async fn hybrid_search_with_text_mode(
726 &self,
727 token: &NamespaceToken,
728 query_text: &str,
729 query_vector: Option<Vec<f32>>,
730 limit: u32,
731 entity_kind: Option<&str>,
732 entity_type: Option<&str>,
733 tags_any: &[String],
734 properties_filter: Option<&serde_json::Value>,
735 text_mode: TextQueryMode,
736 ) -> RuntimeResult<Vec<SearchHit>> {
737 let (hits, _vector_error) = self
738 .hybrid_search_inner(
739 token,
740 query_text,
741 query_vector,
742 limit,
743 entity_kind,
744 entity_type,
745 tags_any,
746 properties_filter,
747 text_mode,
748 None,
749 false,
750 None,
751 )
752 .await?;
753 Ok(hits)
754 }
755
756 pub async fn hybrid_search_each_kind(
767 &self,
768 token: &NamespaceToken,
769 query_text: &str,
770 query_vector: Option<Vec<f32>>,
771 limit: u32,
772 entity_kinds: &[&str],
773 ) -> RuntimeResult<Vec<Vec<SearchHit>>> {
774 if entity_kinds.is_empty() {
775 return Ok(Vec::new());
776 }
777 let candidates = limit.saturating_mul(CANDIDATE_MULTIPLIER).max(limit);
778 let (vector_hits, _vector_error) = self
779 .hybrid_vector_stage(token, query_text, query_vector, candidates, None, false)
780 .await?;
781 let mut per_kind = Vec::with_capacity(entity_kinds.len());
782 for &kind in entity_kinds {
783 let (hits, _vector_error) = self
784 .hybrid_search_inner(
785 token,
786 query_text,
787 None,
788 limit,
789 Some(kind),
790 None,
791 &[],
792 None,
793 TextQueryMode::Plain,
794 None,
795 false,
796 Some(vector_hits.clone()),
797 )
798 .await?;
799 per_kind.push(hits);
800 }
801 Ok(per_kind)
802 }
803
804 #[allow(clippy::too_many_arguments)]
809 pub(crate) async fn hybrid_search_with_vector_similarity_floor(
810 &self,
811 token: &NamespaceToken,
812 query_text: &str,
813 query_vector: Option<Vec<f32>>,
814 limit: u32,
815 entity_kind: Option<&str>,
816 entity_type: Option<&str>,
817 tags_any: &[String],
818 properties_filter: Option<&serde_json::Value>,
819 vector_similarity_floor: f64,
820 ) -> RuntimeResult<Vec<SearchHit>> {
821 let (hits, _vector_error) = self
822 .hybrid_search_inner(
823 token,
824 query_text,
825 query_vector,
826 limit,
827 entity_kind,
828 entity_type,
829 tags_any,
830 properties_filter,
831 TextQueryMode::Plain,
832 Some(vector_similarity_floor),
833 false,
834 None,
835 )
836 .await?;
837 Ok(hits)
838 }
839
840 #[allow(clippy::too_many_arguments)]
850 pub async fn hybrid_search_outcome(
851 &self,
852 token: &NamespaceToken,
853 query_text: &str,
854 limit: u32,
855 entity_kind: Option<&str>,
856 entity_type: Option<&str>,
857 tags_any: &[String],
858 properties_filter: Option<&serde_json::Value>,
859 ) -> RuntimeResult<HybridSearchOutcome> {
860 self.hybrid_search_outcome_with_text_mode(
861 token,
862 query_text,
863 limit,
864 entity_kind,
865 entity_type,
866 tags_any,
867 properties_filter,
868 TextQueryMode::Plain,
869 )
870 .await
871 }
872
873 #[allow(clippy::too_many_arguments)]
875 pub async fn hybrid_search_outcome_with_text_mode(
876 &self,
877 token: &NamespaceToken,
878 query_text: &str,
879 limit: u32,
880 entity_kind: Option<&str>,
881 entity_type: Option<&str>,
882 tags_any: &[String],
883 properties_filter: Option<&serde_json::Value>,
884 text_mode: TextQueryMode,
885 ) -> RuntimeResult<HybridSearchOutcome> {
886 let (hits, vector_error) = self
887 .hybrid_search_inner(
888 token,
889 query_text,
890 None,
891 limit,
892 entity_kind,
893 entity_type,
894 tags_any,
895 properties_filter,
896 text_mode,
897 None,
898 true,
899 None,
900 )
901 .await?;
902 Ok(HybridSearchOutcome { hits, vector_error })
903 }
904
905 #[allow(clippy::too_many_arguments)]
908 async fn hybrid_search_inner(
909 &self,
910 token: &NamespaceToken,
911 query_text: &str,
912 query_vector: Option<Vec<f32>>,
913 limit: u32,
914 entity_kind: Option<&str>,
915 entity_type: Option<&str>,
916 tags_any: &[String],
917 properties_filter: Option<&serde_json::Value>,
918 text_mode: TextQueryMode,
919 vector_similarity_floor: Option<f64>,
920 tolerate_vector_error: bool,
921 vector_pool: Option<Vec<VectorSearchHit>>,
922 ) -> RuntimeResult<(Vec<SearchHit>, Option<String>)> {
923 let candidates = limit.saturating_mul(CANDIDATE_MULTIPLIER).max(limit);
924
925 let visible_ns: Vec<String> = token
926 .visible_namespaces()
927 .iter()
928 .map(|ns| ns.as_str().to_owned())
929 .collect();
930 let text_store = self.text(token)?;
934 let text_fut = text_store.search(TextSearchRequest {
935 query: query_text.to_string(),
936 mode: text_mode,
937 filter: Some(TextFilter {
938 namespaces: visible_ns.clone(),
939 record_kinds: entity_kind
946 .map(|kind| vec![kind.to_string()])
947 .unwrap_or_default(),
948 ..TextFilter::default()
949 }),
950 top_k: candidates,
951 snippet_chars: 200,
952 });
953 let text_fut = crate::stage_seam::text_stage(text_fut);
954 let vector_fut = async {
956 match vector_pool {
957 Some(pool) => Ok((pool, None)),
958 None => {
959 self.hybrid_vector_stage(
960 token,
961 query_text,
962 query_vector,
963 candidates,
964 vector_similarity_floor,
965 tolerate_vector_error,
966 )
967 .await
968 }
969 }
970 };
971 let (text_search_result, vector_result) = tokio::join!(text_fut, vector_fut);
972 let text_hits = crate::error::fts_text_leg_or_err(
977 text_search_result.map_err(RuntimeError::from),
978 "hybrid_search",
979 query_text,
980 )?;
981 let (vector_hits, vector_error) = vector_result?;
982
983 let fusion_limit = text_hits.len().saturating_add(vector_hits.len());
988 let mut fused = rrf_fuse(text_hits, vector_hits, fusion_limit, query_text);
989
990 if !fused.is_empty() {
993 let candidate_ids: Vec<Uuid> = fused.iter().map(|h| h.entity_id).collect();
994 let alive_page = self
995 .entities(token)?
996 .query_entities(
997 token.namespace().as_str(),
998 EntityFilter {
999 ids: candidate_ids,
1000 kinds: entity_kind.map(|k| vec![k.to_string()]).unwrap_or_default(),
1001 entity_types: entity_type.map(|t| vec![t.to_string()]).unwrap_or_default(),
1002 namespaces: visible_ns,
1003 tags_any: tags_any.to_vec(),
1004 ..EntityFilter::default()
1005 },
1006 PageRequest {
1007 offset: 0,
1008 limit: u32::try_from(fused.len()).unwrap_or(u32::MAX),
1009 },
1010 )
1011 .await?;
1012 let mut entity_meta: HashMap<Uuid, (String, Option<String>)> = HashMap::new();
1013 let mut alive: HashSet<Uuid> = HashSet::new();
1014 for e in alive_page.items {
1015 if let Some(pf) = properties_filter {
1018 if !properties_match(e.properties.as_ref(), pf) {
1019 continue;
1020 }
1021 }
1022 alive.insert(e.id);
1023 entity_meta.insert(e.id, (e.name, e.description));
1024 }
1025
1026 fused.retain(|h| alive.contains(&h.entity_id));
1027
1028 for hit in &mut fused {
1030 if let Some((name, description)) = entity_meta.get(&hit.entity_id) {
1031 if hit.title.is_none() {
1032 hit.title = Some(name.clone());
1033 }
1034 if hit.snippet.is_none() {
1035 hit.snippet = description.clone();
1036 }
1037 }
1038 }
1039 }
1040
1041 fused.truncate(limit as usize);
1042 Ok((fused, vector_error))
1043 }
1044
1045 async fn hybrid_vector_stage(
1049 &self,
1050 token: &NamespaceToken,
1051 query_text: &str,
1052 query_vector: Option<Vec<f32>>,
1053 candidates: u32,
1054 vector_similarity_floor: Option<f64>,
1055 tolerate_vector_error: bool,
1056 ) -> RuntimeResult<(Vec<VectorSearchHit>, Option<String>)> {
1057 let mut vector_error: Option<String> = None;
1058 let mut vector_hits = if query_vector.is_some() || self.config().embedding_model.is_some() {
1059 match self
1060 .vector_search(
1061 token,
1062 query_vector,
1063 Some(query_text),
1064 candidates,
1065 Some(SubstrateKind::Entity),
1066 )
1067 .await
1068 {
1069 Ok(hits) => hits,
1070 Err(e) if tolerate_vector_error => {
1071 vector_error = Some(e.to_string());
1072 Vec::new()
1073 }
1074 Err(e) => return Err(e),
1075 }
1076 } else {
1077 Vec::new()
1078 };
1079 if let Some(cosine_floor) = vector_similarity_floor {
1080 let score_floor = DeterministicScore::from_f64(cosine_floor);
1083 vector_hits.retain(|hit| hit.score >= score_floor);
1084 }
1085 Ok((vector_hits, vector_error))
1086 }
1087
1088 pub async fn knn(
1094 &self,
1095 token: &NamespaceToken,
1096 query_vector: Vec<f32>,
1097 top_k: u32,
1098 ) -> RuntimeResult<Vec<VectorSearchHit>> {
1099 let ns = token.namespace().as_str().to_owned();
1100 Ok(self
1101 .vectors(token)?
1102 .search(VectorSearchRequest {
1103 query_vectors: vec![query_vector],
1104 top_k,
1105 namespace: Some(ns),
1106 kind: Some(SubstrateKind::Entity),
1107 embedding_model: None,
1108 filter: None,
1109 backend_hints: None,
1110 })
1111 .await?)
1112 }
1113
1114 pub async fn rerank(
1120 &self,
1121 token: &NamespaceToken,
1122 query_vector: &[f32],
1123 candidate_ids: &[Uuid],
1124 top_k: u32,
1125 ) -> RuntimeResult<Vec<VectorSearchHit>> {
1126 let candidate_set: HashSet<Uuid> = candidate_ids.iter().copied().collect();
1127 let ns = token.namespace().as_str().to_owned();
1128 let all_hits = self
1129 .vectors(token)?
1130 .search(VectorSearchRequest {
1131 query_vectors: vec![query_vector.to_vec()],
1132 top_k: candidate_ids.len() as u32,
1133 namespace: Some(ns),
1134 kind: Some(SubstrateKind::Entity),
1135 embedding_model: None,
1136 filter: None,
1137 backend_hints: None,
1138 })
1139 .await?;
1140 let mut hits: Vec<VectorSearchHit> = all_hits
1141 .into_iter()
1142 .filter(|h| candidate_set.contains(&h.subject_id))
1143 .collect();
1144 hits.sort_by_key(|hit| std::cmp::Reverse(hit.score));
1145 hits.truncate(top_k as usize);
1146 Ok(hits)
1147 }
1148
1149 async fn embed_backfill_page(
1150 &self,
1151 token: &NamespaceToken,
1152 model_name: &str,
1153 inputs: &[(Uuid, String)],
1154 ) -> Vec<Option<DocumentEmbeddingOutcome>> {
1155 let texts: Vec<String> = inputs.iter().map(|(_, text)| text.clone()).collect();
1156 match self
1157 .embed_document_batch_with_model_outcomes_for_token(token, model_name, &texts)
1158 .await
1159 {
1160 Ok(outcomes) => outcomes.into_iter().map(Some).collect(),
1161 Err(error) => {
1162 tracing::warn!(
1163 model = %model_name,
1164 error = %error,
1165 "backfill_missing_embeddings: batch embed failed; retrying records individually"
1166 );
1167 let mut outcomes = Vec::with_capacity(inputs.len());
1168 for (id, text) in inputs {
1169 match self
1170 .embed_document_with_model_outcome_for_token(token, model_name, text)
1171 .await
1172 {
1173 Ok(outcome) => outcomes.push(Some(outcome)),
1174 Err(error) => {
1175 tracing::warn!(
1176 id = %id,
1177 model = %model_name,
1178 error = %error,
1179 "backfill_missing_embeddings: record embed failed"
1180 );
1181 outcomes.push(None);
1182 }
1183 }
1184 }
1185 outcomes
1186 }
1187 }
1188 }
1189
1190 pub async fn backfill_missing_embeddings(&self, token: &NamespaceToken) -> RuntimeResult<u64> {
1203 use khive_storage::types::{SqlRow, SqlStatement, SqlValue};
1204
1205 let model_names = self.registered_embedding_model_names();
1206 if model_names.is_empty() {
1207 tracing::debug!(
1208 "backfill_missing_embeddings: no embedding models registered, skipping"
1209 );
1210 return Ok(0);
1211 }
1212
1213 let ns = token.namespace().as_str().to_string();
1214 let mut total_backfilled = 0u64;
1215
1216 for model_name in &model_names {
1217 let mut model_truncation = EmbeddingTruncationReport::default();
1218 match self.vectors_for_model(token, model_name) {
1219 Ok(_) => {}
1220 Err(error) => {
1221 tracing::warn!(model = %model_name, error = %error,
1222 "backfill_missing_embeddings: vector store unavailable");
1223 continue;
1224 }
1225 };
1226 let vec_table = format!("vec_{}", sanitize_key(model_name));
1228
1229 const PAGE_SIZE: usize = EMBEDDING_BATCH_PAGE_SIZE;
1233 let mut entity_total = 0usize;
1234 let mut entity_cursor = String::new();
1235 loop {
1236 let entity_sql = SqlStatement {
1237 sql: format!(
1238 "SELECT id FROM entities \
1239 WHERE namespace = ?1 AND deleted_at IS NULL AND id > ?3 \
1240 AND id NOT IN (\
1241 SELECT subject_id FROM {vec_table} \
1242 WHERE namespace = ?1 AND embedding_model = ?2 \
1243 ) ORDER BY id LIMIT {PAGE_SIZE}"
1244 ),
1245 params: vec![
1246 SqlValue::Text(ns.clone()),
1247 SqlValue::Text(model_name.clone()),
1248 SqlValue::Text(entity_cursor.clone()),
1249 ],
1250 label: Some("backfill_entities".into()),
1251 };
1252
1253 let entity_rows: Vec<SqlRow> = {
1254 let sql = self.sql();
1255 let reader_result = sql.reader().await;
1256 #[cfg(any(test, feature = "fault-injection"))]
1257 let reader_result = if BACKFILL_READER_FAIL.with(|c| c.get()) {
1258 BACKFILL_READER_FAIL.with(|c| c.set(false));
1259 Err(khive_storage::StorageError::Pool {
1260 operation: "reader".into(),
1261 message: "injected failure".into(),
1262 })
1263 } else {
1264 reader_result
1265 };
1266 let mut reader = reader_result.map_err(RuntimeError::Storage)?;
1267 reader
1268 .query_all(entity_sql)
1269 .await
1270 .map_err(RuntimeError::Storage)?
1271 };
1272
1273 let batch_len = entity_rows.len();
1274 entity_total += batch_len;
1275 if batch_len == 0 {
1276 break;
1277 }
1278 entity_cursor = match entity_rows.last().and_then(|row| row.columns.first()) {
1279 Some(column) => match &column.value {
1280 SqlValue::Text(id) => id.clone(),
1281 _ => {
1282 return Err(RuntimeError::Internal(
1283 "backfill entity ID is not text".into(),
1284 ))
1285 }
1286 },
1287 None => {
1288 return Err(RuntimeError::Internal(
1289 "backfill entity page is empty".into(),
1290 ))
1291 }
1292 };
1293
1294 let entity_store = self.entities(token)?;
1295 let mut entities = Vec::with_capacity(batch_len);
1296 let mut inputs = Vec::with_capacity(batch_len);
1297 for row in &entity_rows {
1298 let Some(SqlValue::Text(id)) = row.columns.first().map(|column| &column.value)
1299 else {
1300 continue;
1301 };
1302 let Ok(id) = id.parse::<Uuid>() else { continue };
1303 let Some(entity) = entity_store.get_entity(id).await? else {
1304 continue;
1305 };
1306 if entity.namespace != ns || entity.deleted_at.is_some() {
1307 continue;
1308 }
1309 let text = crate::curation::entity_embedding_text(&entity);
1310 if text.trim().is_empty() {
1311 continue;
1312 }
1313 inputs.push((id, text));
1314 entities.push(entity);
1315 }
1316 let outcomes = self.embed_backfill_page(token, model_name, &inputs).await;
1317 for (entity, outcome) in entities.into_iter().zip(outcomes) {
1318 let Some(outcome) = outcome else { continue };
1319 model_truncation.observe(&outcome);
1320 match self
1321 .publish_entity_vector_revision(token, &entity, model_name, &outcome.vector)
1322 .await
1323 {
1324 Ok(true) => total_backfilled += 1,
1325 Ok(false) => {}
1326 Err(error) => tracing::warn!(
1327 id = %entity.id, model = %model_name, error = %error,
1328 "backfill_missing_embeddings: entity vector insert failed"
1329 ),
1330 }
1331 }
1332
1333 if batch_len < PAGE_SIZE {
1334 break;
1335 }
1336 }
1337
1338 let text_store = self.text_for_notes(token).ok();
1340 let note_store = self.notes(token).ok();
1341 let mut note_total = 0usize;
1342 let mut note_cursor = String::new();
1345 loop {
1346 let note_sql = SqlStatement {
1349 sql: format!(
1350 "SELECT id FROM notes \
1351 WHERE namespace = ?1 AND deleted_at IS NULL AND id > ?3 \
1352 AND id NOT IN (\
1353 SELECT subject_id FROM {vec_table} \
1354 WHERE namespace = ?1 AND embedding_model = ?2 \
1355 ) ORDER BY id LIMIT {PAGE_SIZE}"
1356 ),
1357 params: vec![
1358 SqlValue::Text(ns.clone()),
1359 SqlValue::Text(model_name.clone()),
1360 SqlValue::Text(note_cursor.clone()),
1361 ],
1362 label: Some("backfill_notes".into()),
1363 };
1364
1365 let note_rows: Vec<SqlRow> = {
1366 let sql = self.sql();
1367 let reader_result = sql.reader().await;
1368 #[cfg(any(test, feature = "fault-injection"))]
1369 let reader_result = if BACKFILL_READER_FAIL.with(|c| c.get()) {
1370 BACKFILL_READER_FAIL.with(|c| c.set(false));
1371 Err(khive_storage::StorageError::Pool {
1372 operation: "reader".into(),
1373 message: "injected failure".into(),
1374 })
1375 } else {
1376 reader_result
1377 };
1378 let mut reader = reader_result.map_err(RuntimeError::Storage)?;
1379 reader
1380 .query_all(note_sql)
1381 .await
1382 .map_err(RuntimeError::Storage)?
1383 };
1384
1385 let batch_len = note_rows.len();
1386 note_total += batch_len;
1387 if batch_len == 0 {
1388 break;
1389 }
1390 note_cursor = match note_rows.last().and_then(|row| row.columns.first()) {
1391 Some(column) => match &column.value {
1392 SqlValue::Text(id) => id.clone(),
1393 _ => {
1394 return Err(RuntimeError::Internal(
1395 "backfill note ID is not text".into(),
1396 ))
1397 }
1398 },
1399 None => {
1400 return Err(RuntimeError::Internal("backfill note page is empty".into()))
1401 }
1402 };
1403
1404 let mut notes_for_batch = Vec::with_capacity(batch_len);
1405 let mut inputs = Vec::with_capacity(batch_len);
1406 for row in ¬e_rows {
1407 let id_str = row.columns.first().and_then(|c| {
1408 if let SqlValue::Text(s) = &c.value {
1409 Some(s.clone())
1410 } else {
1411 None
1412 }
1413 });
1414
1415 let Some(id_str) = id_str else {
1416 continue;
1417 };
1418 let Ok(id) = id_str.parse::<Uuid>() else {
1419 continue;
1420 };
1421
1422 let note = match ¬e_store {
1423 Some(store) => match store.get_note(id).await {
1424 Ok(Some(n)) => n,
1425 _ => continue,
1426 },
1427 None => continue,
1428 };
1429
1430 if note.content.trim().is_empty() {
1431 continue;
1432 }
1433
1434 if model_names.first().map(|n| n.as_str()) == Some(model_name.as_str()) {
1437 if let Some(ref ts) = text_store {
1438 if let Err(e) = ts.upsert_document(note_fts_document(¬e)).await {
1439 tracing::warn!(id = %id, error = %e,
1440 "backfill_missing_embeddings: note FTS upsert failed");
1441 }
1442 }
1443 }
1444
1445 if !self
1446 .embedding_models_for_note_kind(¬e.kind)
1447 .contains(model_name)
1448 {
1449 continue;
1450 }
1451
1452 inputs.push((id, note.content.clone()));
1453 notes_for_batch.push(note);
1454 }
1455
1456 let outcomes = self.embed_backfill_page(token, model_name, &inputs).await;
1457 for (note, outcome) in notes_for_batch.into_iter().zip(outcomes) {
1458 let Some(outcome) = outcome else { continue };
1459 model_truncation.observe(&outcome);
1460 match self
1461 .publish_note_vector_revision(token, ¬e, model_name, &outcome.vector)
1462 .await
1463 {
1464 Ok(true) => total_backfilled += 1,
1465 Ok(false) => {}
1466 Err(error) => tracing::warn!(
1467 id = %note.id, model = %model_name, error = %error,
1468 "backfill_missing_embeddings: note vector insert failed"
1469 ),
1470 }
1471 }
1472
1473 if batch_len < PAGE_SIZE {
1474 break;
1475 }
1476 }
1477
1478 tracing::info!(
1479 model = %model_name,
1480 namespace = %ns,
1481 entities = entity_total,
1482 notes = note_total,
1483 truncated = model_truncation.truncated,
1484 discarded_bytes = model_truncation.discarded_bytes,
1485 "backfill_missing_embeddings: model pass complete"
1486 );
1487 }
1488
1489 tracing::info!(
1490 namespace = %ns,
1491 total_backfilled = total_backfilled,
1492 "backfill_missing_embeddings: finished"
1493 );
1494
1495 Ok(total_backfilled)
1496 }
1497
1498 pub async fn sweep_orphan_vectors(
1514 &self,
1515 token: &NamespaceToken,
1516 max_delete_per_model: u32,
1517 dry_run: bool,
1518 ) -> RuntimeResult<u64> {
1519 use khive_storage::types::OrphanSweepConfig;
1520 use khive_storage::StorageError;
1521
1522 let model_names = self.registered_embedding_model_names();
1523 if model_names.is_empty() {
1524 tracing::debug!("sweep_orphan_vectors: no embedding models registered, skipping");
1525 return Ok(0);
1526 }
1527
1528 let ns = token.namespace().as_str().to_string();
1529 let mut total_deleted = 0u64;
1530
1531 for model_name in &model_names {
1532 let store = match self.vectors_for_model(token, model_name) {
1533 Ok(s) => s,
1534 Err(e) => {
1535 tracing::warn!(
1536 model = %model_name,
1537 error = %e,
1538 "sweep_orphan_vectors: failed to get vector store, skipping model"
1539 );
1540 continue;
1541 }
1542 };
1543
1544 let caps = store.capabilities();
1545 if !caps.supports_orphan_sweep {
1546 tracing::debug!(
1547 model = %model_name,
1548 "sweep_orphan_vectors: backend does not support orphan sweep, skipping"
1549 );
1550 continue;
1551 }
1552
1553 let config = OrphanSweepConfig {
1554 subject_id_allowlist: None,
1555 namespaces: vec![ns.clone()],
1556 substrate_kinds: vec![],
1557 max_delete: max_delete_per_model,
1558 dry_run,
1559 };
1560
1561 match store.orphan_sweep(&config).await {
1562 Ok(result) => {
1563 tracing::info!(
1564 model = %model_name,
1565 namespace = %ns,
1566 scanned = result.scanned,
1567 deleted = result.deleted,
1568 would_delete = result.would_delete,
1569 dry_run = dry_run,
1570 "sweep_orphan_vectors: sweep complete"
1571 );
1572 total_deleted += result.deleted;
1573 }
1574 Err(StorageError::Unsupported { .. }) => {
1575 tracing::debug!(
1576 model = %model_name,
1577 "sweep_orphan_vectors: backend returned Unsupported, skipping"
1578 );
1579 }
1580 Err(e) => {
1581 tracing::warn!(
1582 model = %model_name,
1583 error = %e,
1584 "sweep_orphan_vectors: sweep failed, continuing with other models"
1585 );
1586 }
1587 }
1588 }
1589
1590 tracing::info!(
1591 namespace = %ns,
1592 total_deleted = total_deleted,
1593 dry_run = dry_run,
1594 "sweep_orphan_vectors: finished"
1595 );
1596
1597 Ok(total_deleted)
1598 }
1599}
1600
1601pub(crate) fn properties_match(
1606 properties: Option<&serde_json::Value>,
1607 filter: &serde_json::Value,
1608) -> bool {
1609 let required = match filter.as_object() {
1610 Some(obj) if !obj.is_empty() => obj,
1611 _ => return true,
1612 };
1613 let actual = match properties.and_then(serde_json::Value::as_object) {
1614 Some(obj) => obj,
1615 None => return false,
1616 };
1617 required
1618 .iter()
1619 .all(|(k, v)| actual.get(k).is_some_and(|av| av == v))
1620}
1621
1622const EXACT_MATCH_BOOST: f64 = 0.5;
1626
1627fn rrf_fuse(
1636 text_hits: Vec<TextSearchHit>,
1637 vector_hits: Vec<VectorSearchHit>,
1638 limit: usize,
1639 query_text: &str,
1640) -> Vec<SearchHit> {
1641 let mut text_arm = Vec::new();
1642 for (rank, hit) in text_hits.into_iter().enumerate() {
1643 let label = HitLabel {
1644 rank,
1645 signals: SearchSignals {
1646 vector_similarity: None,
1647 keyword_score: Some(hit.score),
1648 },
1649 source: SearchSource::Text,
1650 title: hit.title,
1651 snippet: hit.snippet,
1652 };
1653 text_arm.push((hit.subject_id, label));
1654 }
1655 let mut vector_arm = Vec::new();
1656 for (rank, hit) in vector_hits.into_iter().enumerate() {
1657 let label = HitLabel {
1658 rank,
1659 signals: SearchSignals {
1660 vector_similarity: Some(hit.score),
1661 keyword_score: None,
1662 },
1663 source: SearchSource::Vector,
1664 title: None,
1665 snippet: None,
1666 };
1667 vector_arm.push((hit.subject_id, label));
1668 }
1669
1670 let fused = fuse_labelled(
1671 vec![text_arm, vector_arm],
1672 RRF_K,
1673 combine_leg_first_appearance,
1674 );
1675
1676 let query_lower = query_text.to_lowercase();
1678 let boost = DeterministicScore::from_f64(EXACT_MATCH_BOOST);
1679 let mut hits = Vec::with_capacity(fused.len());
1680 for (entity_id, score, label) in fused {
1681 let title_lower = label.title.as_deref().map(str::to_lowercase);
1682 let exact = title_lower.as_deref() == Some(query_lower.as_str());
1683 let score = if exact { score + boost } else { score };
1684 hits.push(SearchHit {
1685 entity_id,
1686 score,
1687 rank_score_kind: RankScoreKind::Rrf,
1688 signals: label.signals,
1689 source: label.source,
1690 title: label.title,
1691 snippet: label.snippet,
1692 });
1693 }
1694
1695 hits.sort_by(|a, b| b.score.cmp(&a.score).then(a.entity_id.cmp(&b.entity_id)));
1696 hits.truncate(limit);
1697 hits
1698}
1699
1700#[cfg(test)]
1701mod rrf_fuse_label_tests;
1702
1703#[cfg(test)]
1704mod tests {
1705 use super::*;
1706 use std::sync::atomic::{AtomicUsize, Ordering};
1707 use std::sync::Arc;
1708
1709 use crate::runtime::{KhiveRuntime, NamespaceToken, RuntimeConfig};
1710 use khive_score::rrf_score;
1711 use khive_storage::types::{TextSearchHit, VectorSearchHit};
1712 use khive_types::namespace::Namespace;
1713 use lattice_embed::{EmbedError, EmbeddingModel};
1714
1715 struct FailingEmbeddingService;
1719
1720 #[async_trait::async_trait]
1721 impl EmbeddingService for FailingEmbeddingService {
1722 async fn embed(
1723 &self,
1724 _texts: &[String],
1725 _model: EmbeddingModel,
1726 ) -> Result<Vec<Vec<f32>>, EmbedError> {
1727 Err(EmbedError::ModelInitialization(
1728 "injected vector-arm failure".to_string(),
1729 ))
1730 }
1731
1732 fn supports_model(&self, _model: EmbeddingModel) -> bool {
1733 true
1734 }
1735
1736 fn name(&self) -> &'static str {
1737 "hybrid-search-test-failing-embedding"
1738 }
1739 }
1740
1741 struct FailingEmbedderProvider {
1742 name: String,
1743 dimensions: usize,
1744 }
1745
1746 #[async_trait::async_trait]
1747 impl EmbedderProvider for FailingEmbedderProvider {
1748 fn name(&self) -> &str {
1749 &self.name
1750 }
1751
1752 fn dimensions(&self) -> usize {
1753 self.dimensions
1754 }
1755
1756 async fn build(&self) -> RuntimeResult<Arc<dyn EmbeddingService>> {
1757 Ok(Arc::new(FailingEmbeddingService))
1758 }
1759 }
1760
1761 fn break_vector_arm(runtime: &KhiveRuntime) {
1768 let model = EmbeddingModel::AllMiniLmL6V2;
1769 runtime.register_embedder(FailingEmbedderProvider {
1770 name: model.to_string(),
1771 dimensions: model.dimensions(),
1772 });
1773 }
1774
1775 struct ConstantEmbeddingService {
1779 dimensions: usize,
1780 }
1781
1782 #[async_trait::async_trait]
1783 impl EmbeddingService for ConstantEmbeddingService {
1784 async fn embed(
1785 &self,
1786 texts: &[String],
1787 _model: EmbeddingModel,
1788 ) -> Result<Vec<Vec<f32>>, EmbedError> {
1789 Ok(texts.iter().map(|_| vec![1.0; self.dimensions]).collect())
1790 }
1791
1792 fn supports_model(&self, _model: EmbeddingModel) -> bool {
1793 true
1794 }
1795
1796 fn name(&self) -> &'static str {
1797 "hybrid-search-test-constant-embedding"
1798 }
1799 }
1800
1801 struct ConstantEmbedderProvider {
1802 name: String,
1803 dimensions: usize,
1804 }
1805
1806 #[async_trait::async_trait]
1807 impl EmbedderProvider for ConstantEmbedderProvider {
1808 fn name(&self) -> &str {
1809 &self.name
1810 }
1811
1812 fn dimensions(&self) -> usize {
1813 self.dimensions
1814 }
1815
1816 async fn build(&self) -> RuntimeResult<Arc<dyn EmbeddingService>> {
1817 Ok(Arc::new(ConstantEmbeddingService {
1818 dimensions: self.dimensions,
1819 }))
1820 }
1821 }
1822
1823 fn runtime_with_constant_embeddings() -> KhiveRuntime {
1826 let model = EmbeddingModel::AllMiniLmL6V2;
1827 let runtime = KhiveRuntime::new(RuntimeConfig {
1828 db_path: None,
1829 embedding_model: Some(model),
1830 packs: vec!["kg".to_string()],
1831 ..RuntimeConfig::no_embeddings()
1832 })
1833 .expect("in-memory runtime");
1834 runtime.register_embedder(ConstantEmbedderProvider {
1835 name: model.to_string(),
1836 dimensions: model.dimensions(),
1837 });
1838 runtime
1839 }
1840
1841 #[test]
1842 fn bounded_embedding_input_reserves_prefix_and_preserves_utf8() {
1843 assert_eq!(
1844 document_embedding_budget("multilingual-e5-base"),
1845 MAX_TEXT_BYTES - "passage: ".len()
1846 );
1847
1848 let input = format!("{}\u{1f980}tail", "a".repeat(MAX_TEXT_BYTES - 1));
1849 let (bounded, truncated) = bounded_embedding_input(&input, MAX_TEXT_BYTES);
1850 assert!(truncated);
1851 assert_eq!(bounded.len(), MAX_TEXT_BYTES - 1);
1852 assert!(bounded.is_char_boundary(bounded.len()));
1853 assert!(!bounded.contains('\u{1f980}'));
1854 }
1855
1856 #[test]
1857 fn bounded_embedding_input_leaves_normal_text_unchanged() {
1858 let input = "normal byte-identical embedding input";
1859 let (bounded, truncated) = bounded_embedding_input(input, MAX_TEXT_BYTES);
1860 assert!(!truncated);
1861 assert_eq!(bounded, input);
1862 assert_eq!(bounded.as_ptr(), input.as_ptr());
1863 }
1864
1865 fn text_hit(id: Uuid, rank: u32, title: &str) -> TextSearchHit {
1866 TextSearchHit {
1867 subject_id: id,
1868 score: DeterministicScore::from_f64(1.0),
1869 rank,
1870 title: Some(title.to_string()),
1871 snippet: Some("...".to_string()),
1872 }
1873 }
1874
1875 fn vector_hit(id: Uuid, rank: u32) -> VectorSearchHit {
1876 VectorSearchHit {
1877 subject_id: id,
1878 score: DeterministicScore::from_f64(0.9),
1879 rank,
1880 }
1881 }
1882
1883 #[test]
1884 fn rrf_evidence_keeps_absence_distinct_from_measured_zero() {
1885 let text_id = Uuid::from_u128(1);
1886 let vector_id = Uuid::from_u128(2);
1887 let both_id = Uuid::from_u128(3);
1888 let quarter = DeterministicScore::from_raw(1_i64 << 30);
1889 let half = DeterministicScore::from_raw(1_i64 << 31);
1890 let text = vec![
1891 TextSearchHit {
1892 score: DeterministicScore::ZERO,
1893 ..text_hit(text_id, 1, "text")
1894 },
1895 TextSearchHit {
1896 score: quarter,
1897 ..text_hit(both_id, 2, "both")
1898 },
1899 text_hit(both_id, 3, "duplicate"),
1900 ];
1901 let vector = vec![
1902 VectorSearchHit {
1903 score: DeterministicScore::ZERO,
1904 ..vector_hit(vector_id, 1)
1905 },
1906 VectorSearchHit {
1907 score: half,
1908 ..vector_hit(both_id, 2)
1909 },
1910 vector_hit(both_id, 3),
1911 ];
1912 let hits = rrf_fuse(text, vector, 10, "unmatched");
1913 assert_eq!(hits.len(), 3);
1914 for (id, source, signals) in [
1915 (
1916 text_id,
1917 SearchSource::Text,
1918 SearchSignals {
1919 vector_similarity: None,
1920 keyword_score: Some(DeterministicScore::ZERO),
1921 },
1922 ),
1923 (
1924 vector_id,
1925 SearchSource::Vector,
1926 SearchSignals {
1927 vector_similarity: Some(DeterministicScore::ZERO),
1928 keyword_score: None,
1929 },
1930 ),
1931 (
1932 both_id,
1933 SearchSource::Both,
1934 SearchSignals {
1935 vector_similarity: Some(half),
1936 keyword_score: Some(quarter),
1937 },
1938 ),
1939 ] {
1940 let hit = hits.iter().find(|hit| hit.entity_id == id).unwrap();
1941 assert_eq!(hit.rank_score_kind, RankScoreKind::Rrf);
1942 assert_eq!(hit.source, source);
1943 assert_eq!(hit.signals, signals);
1944 }
1945 }
1946
1947 #[test]
1948 fn rrf_evidence_golden_preserves_true_ties_across_permutations() {
1949 let a = Uuid::from_u128(1);
1950 let b = Uuid::from_u128(2);
1951 let expected = vec![
1952 (
1953 a,
1954 748_365_513,
1955 RankScoreKind::Rrf,
1956 SearchSignals {
1957 vector_similarity: Some(DeterministicScore::from_raw(1_i64 << 31)),
1958 keyword_score: Some(DeterministicScore::from_raw(1_i64 << 30)),
1959 },
1960 ),
1961 (
1962 b,
1963 748_365_513,
1964 RankScoreKind::Rrf,
1965 SearchSignals {
1966 vector_similarity: Some(DeterministicScore::from_raw(1_i64 << 32)),
1967 keyword_score: Some(DeterministicScore::from_raw(3_i64 << 30)),
1968 },
1969 ),
1970 ];
1971 for (text_ids, vector_ids) in [([a, b], [b, a]), ([b, a], [a, b])] {
1973 for _ in 0..4 {
1974 let text = text_ids
1975 .into_iter()
1976 .enumerate()
1977 .map(|(rank, id)| TextSearchHit {
1978 score: DeterministicScore::from_raw(
1979 if id == a { 1_i64 } else { 3_i64 } << 30,
1980 ),
1981 ..text_hit(id, rank as u32 + 1, "candidate")
1982 })
1983 .collect();
1984 let vector = vector_ids
1985 .into_iter()
1986 .enumerate()
1987 .map(|(rank, id)| VectorSearchHit {
1988 score: DeterministicScore::from_raw(
1989 if id == a { 1_i64 } else { 2_i64 } << 31,
1990 ),
1991 ..vector_hit(id, rank as u32 + 1)
1992 })
1993 .collect();
1994 let hits = rrf_fuse(text, vector, 10, "unmatched");
1995 assert_eq!(hits.len(), 2);
1996 assert_eq!(hits[0].score, hits[1].score);
1997 assert!(hits.iter().all(|hit| hit.source == SearchSource::Both));
1998 let snapshot: Vec<_> = hits
1999 .iter()
2000 .map(|hit| {
2001 (
2002 hit.entity_id,
2003 hit.score.to_raw(),
2004 hit.rank_score_kind,
2005 hit.signals,
2006 )
2007 })
2008 .collect();
2009 assert_eq!(snapshot, expected);
2010 }
2011 }
2012 }
2013
2014 #[test]
2015 fn rrf_fuse_text_only() {
2016 let a = Uuid::new_v4();
2017 let b = Uuid::new_v4();
2018 let text = vec![text_hit(a, 1, "A"), text_hit(b, 2, "B")];
2019 let hits = rrf_fuse(text, vec![], 10, "query");
2020 assert_eq!(hits.len(), 2);
2021 assert_eq!(hits[0].entity_id, a);
2022 assert_eq!(hits[0].source, SearchSource::Text);
2023 assert_eq!(hits[0].title.as_deref(), Some("A"));
2024 }
2025
2026 #[test]
2027 fn rrf_fuse_vector_only() {
2028 let a = Uuid::new_v4();
2029 let hits = rrf_fuse(vec![], vec![vector_hit(a, 1)], 10, "query");
2030 assert_eq!(hits.len(), 1);
2031 assert_eq!(hits[0].source, SearchSource::Vector);
2032 assert!(hits[0].title.is_none());
2033 }
2034
2035 #[test]
2036 fn rrf_fuse_marks_both_when_in_both_lists() {
2037 let id = Uuid::new_v4();
2038 let text = vec![text_hit(id, 1, "A")];
2039 let vec = vec![vector_hit(id, 1)];
2040 let hits = rrf_fuse(text, vec, 10, "query");
2041 assert_eq!(hits.len(), 1);
2042 assert_eq!(hits[0].source, SearchSource::Both);
2043 }
2044
2045 #[test]
2046 fn rrf_fuse_preserves_unique_leg_scores_exactly() {
2047 let text_only = Uuid::new_v4();
2048 let both = Uuid::new_v4();
2049 let vector_only = Uuid::new_v4();
2050 let text = vec![text_hit(text_only, 1, "A"), text_hit(both, 2, "B")];
2051 let vector = vec![vector_hit(both, 1), vector_hit(vector_only, 2)];
2052
2053 let hits = rrf_fuse(text, vector, 10, "query");
2054 let score_for = |id| {
2055 hits.iter()
2056 .find(|hit| hit.entity_id == id)
2057 .expect("expected fused hit")
2058 .score
2059 };
2060
2061 assert_eq!(score_for(text_only), rrf_score(1, RRF_K));
2062 assert_eq!(score_for(both), rrf_score(2, RRF_K) + rrf_score(1, RRF_K));
2063 assert_eq!(score_for(vector_only), rrf_score(2, RRF_K));
2064 }
2065
2066 #[test]
2067 fn rrf_fuse_counts_duplicate_once_per_leg() {
2068 let id = Uuid::new_v4();
2069 let text = vec![text_hit(id, 1, "A"), text_hit(id, 2, "A duplicate")];
2070 let vector = vec![vector_hit(id, 1), vector_hit(id, 2)];
2071
2072 let hits = rrf_fuse(text, vector, 10, "query");
2073
2074 assert_eq!(hits.len(), 1);
2075 assert_eq!(hits[0].source, SearchSource::Both);
2076 assert_eq!(hits[0].score, rrf_score(1, RRF_K) + rrf_score(1, RRF_K));
2077 }
2078
2079 #[test]
2080 fn rrf_fuse_respects_limit() {
2081 let hits: Vec<TextSearchHit> = (0..20)
2082 .map(|i| text_hit(Uuid::new_v4(), i + 1, "x"))
2083 .collect();
2084 let fused = rrf_fuse(hits, vec![], 5, "query");
2085 assert_eq!(fused.len(), 5);
2086 }
2087
2088 #[test]
2089 fn rrf_fuse_orders_higher_score_first() {
2090 let a = Uuid::new_v4();
2092 let b = Uuid::new_v4();
2093 let text = vec![text_hit(a, 1, "A")];
2094 let vec = vec![vector_hit(a, 1), vector_hit(b, 2)];
2095 let hits = rrf_fuse(text, vec, 10, "query");
2096 assert_eq!(hits[0].entity_id, a);
2097 assert_eq!(hits[0].source, SearchSource::Both);
2098 assert!(hits[0].score > hits[1].score);
2099 }
2100
2101 #[test]
2102 fn rrf_fuse_k10_score_spread_exceeds_threshold() {
2103 let ids: Vec<Uuid> = (0..10).map(|_| Uuid::new_v4()).collect();
2106 let text: Vec<TextSearchHit> = ids
2107 .iter()
2108 .enumerate()
2109 .map(|(i, &id)| text_hit(id, (i + 1) as u32, "x"))
2110 .collect();
2111 let hits = rrf_fuse(text, vec![], 10, "query");
2112 assert_eq!(hits.len(), 10);
2113 let top_score = hits[0].score.to_f64();
2114 let bottom_score = hits[9].score.to_f64();
2115 let spread = top_score - bottom_score;
2116 assert!(
2117 spread >= 0.03,
2118 "score spread {spread:.4} between rank 1 and rank 10 must be ≥ 0.03 (was {spread:.4})"
2119 );
2120 }
2121
2122 #[test]
2123 fn rrf_fuse_exact_match_boost_elevates_score() {
2124 let exact_id = Uuid::new_v4();
2127 let other_id = Uuid::new_v4();
2128 let text = vec![
2130 text_hit(other_id, 1, "something else"),
2131 text_hit(exact_id, 2, "FlashAttention"),
2132 ];
2133 let hits = rrf_fuse(text, vec![], 10, "flashattention");
2134 assert_eq!(hits.len(), 2);
2135 assert_eq!(
2136 hits[0].entity_id, exact_id,
2137 "exact match must rank first despite being rank-2 in raw text search"
2138 );
2139 }
2140
2141 #[test]
2144 fn embed_batch_unconfigured_on_memory_runtime() {
2145 let rt = KhiveRuntime::memory().unwrap();
2147 let result = tokio::runtime::Runtime::new()
2148 .unwrap()
2149 .block_on(rt.embed_batch(&[]));
2150 assert!(result.is_ok());
2152 assert!(result.unwrap().is_empty());
2153 }
2154
2155 #[test]
2156 fn embed_batch_empty_input_returns_empty_vec() {
2157 let rt = KhiveRuntime::memory().unwrap();
2159 let result = tokio::runtime::Runtime::new()
2160 .unwrap()
2161 .block_on(rt.embed_batch(&[]));
2162 assert_eq!(result.unwrap(), Vec::<Vec<f32>>::new());
2163 }
2164
2165 #[test]
2166 fn embed_batch_no_model_non_empty_returns_unconfigured() {
2167 let rt = KhiveRuntime::memory().unwrap();
2168 let texts = vec!["hello".to_string()];
2169 let result = tokio::runtime::Runtime::new()
2170 .unwrap()
2171 .block_on(rt.embed_batch(&texts));
2172 match result {
2173 Err(crate::RuntimeError::Unconfigured(s)) => assert_eq!(s, "embedding_model"),
2174 Err(other) => panic!("expected Unconfigured, got {:?}", other),
2175 Ok(_) => panic!("expected Err, got Ok"),
2176 }
2177 }
2178
2179 #[test]
2180 #[ignore = "loads ~80 MB model; run with --include-ignored"]
2181 fn embed_batch_count_matches_input() {
2182 let config = RuntimeConfig {
2183 db_path: None,
2184 default_namespace: Namespace::parse("test").unwrap(),
2185 embedding_model: Some(EmbeddingModel::AllMiniLmL6V2),
2186 packs: vec!["kg".to_string()],
2187 ..RuntimeConfig::default()
2188 };
2189 let rt = KhiveRuntime::new(config).unwrap();
2190 let texts: Vec<String> = vec!["foo".to_string(), "bar".to_string(), "baz".to_string()];
2191 let result = tokio::runtime::Runtime::new()
2192 .unwrap()
2193 .block_on(rt.embed_batch(&texts));
2194 let embeddings = result.unwrap();
2195 assert_eq!(embeddings.len(), texts.len());
2196 }
2197
2198 #[test]
2199 fn vector_search_requires_embedding_or_text() {
2200 let rt = KhiveRuntime::memory().unwrap();
2201 let tok = NamespaceToken::local();
2202 let result = tokio::runtime::Runtime::new()
2203 .unwrap()
2204 .block_on(rt.vector_search(&tok, None, None, 10, Some(SubstrateKind::Entity)));
2205 match result {
2206 Err(crate::RuntimeError::InvalidInput(msg)) => {
2207 assert!(msg.contains("query_embedding or query_text"), "msg: {msg}");
2208 }
2209 other => panic!("expected InvalidInput, got {other:?}"),
2210 }
2211 }
2212
2213 #[test]
2214 fn vector_search_text_without_model_returns_unconfigured() {
2215 let rt = KhiveRuntime::memory().unwrap();
2216 let tok = NamespaceToken::local();
2217 let result = tokio::runtime::Runtime::new()
2218 .unwrap()
2219 .block_on(rt.vector_search(
2220 &tok,
2221 None,
2222 Some("attention"),
2223 10,
2224 Some(SubstrateKind::Entity),
2225 ));
2226 match result {
2227 Err(crate::RuntimeError::Unconfigured(s)) => assert_eq!(s, "embedding_model"),
2228 other => panic!("expected Unconfigured, got {other:?}"),
2229 }
2230 }
2231
2232 #[test]
2233 #[ignore = "loads ~80 MB model; run with --include-ignored"]
2234 fn embed_batch_vectors_have_expected_dimensions() {
2235 let model = EmbeddingModel::AllMiniLmL6V2;
2236 let config = RuntimeConfig {
2237 db_path: None,
2238 default_namespace: Namespace::parse("test").unwrap(),
2239 embedding_model: Some(model),
2240 packs: vec!["kg".to_string()],
2241 ..RuntimeConfig::default()
2242 };
2243 let rt = KhiveRuntime::new(config).unwrap();
2244 let texts = vec!["hello world".to_string()];
2245 let result = tokio::runtime::Runtime::new()
2246 .unwrap()
2247 .block_on(rt.embed_batch(&texts));
2248 let embeddings = result.unwrap();
2249 assert_eq!(embeddings[0].len(), model.dimensions());
2250 }
2251
2252 #[tokio::test]
2260 async fn hybrid_search_still_fails_loud_on_vector_arm_error() {
2261 let rt = runtime_with_constant_embeddings();
2262 let tok = NamespaceToken::local();
2263 rt.create_entity(
2264 &tok,
2265 "concept",
2266 None,
2267 "FlashAttention",
2268 Some("IO-aware exact attention using tiling"),
2269 None,
2270 vec![],
2271 )
2272 .await
2273 .unwrap();
2274 break_vector_arm(&rt);
2275
2276 let result = rt
2277 .hybrid_search(&tok, "FlashAttention", None, 10, None, None, &[], None)
2278 .await;
2279
2280 assert!(
2281 result.is_err(),
2282 "the fail-loud entry point must still propagate a vector-arm failure, got {result:?}"
2283 );
2284 }
2285
2286 #[tokio::test]
2290 async fn hybrid_search_outcome_preserves_text_hits_on_vector_arm_error() {
2291 let rt = runtime_with_constant_embeddings();
2292 let tok = NamespaceToken::local();
2293 rt.create_entity(
2294 &tok,
2295 "concept",
2296 None,
2297 "FlashAttention",
2298 Some("IO-aware exact attention using tiling"),
2299 None,
2300 vec![],
2301 )
2302 .await
2303 .unwrap();
2304 break_vector_arm(&rt);
2305
2306 let outcome = rt
2307 .hybrid_search_outcome(&tok, "FlashAttention", 10, None, None, &[], None)
2308 .await
2309 .expect("text leg must still succeed");
2310
2311 assert!(
2312 !outcome.hits.is_empty(),
2313 "text arm's hit must survive a vector-arm failure"
2314 );
2315 assert!(
2316 outcome.hits[0]
2317 .title
2318 .as_deref()
2319 .unwrap_or_default()
2320 .contains("FlashAttention"),
2321 "surviving hit must be the text match"
2322 );
2323 let vector_error = outcome
2324 .vector_error
2325 .expect("vector arm failure must be reported");
2326 assert!(
2327 vector_error.contains("injected vector-arm failure"),
2328 "vector_error must carry the underlying cause, got {vector_error:?}"
2329 );
2330 }
2331
2332 #[tokio::test]
2335 async fn hybrid_search_outcome_has_no_vector_error_when_vector_arm_healthy() {
2336 let rt = runtime_with_constant_embeddings();
2337 let tok = NamespaceToken::local();
2338 rt.create_entity(
2339 &tok,
2340 "concept",
2341 None,
2342 "FlashAttention",
2343 Some("IO-aware exact attention using tiling"),
2344 None,
2345 vec![],
2346 )
2347 .await
2348 .unwrap();
2349
2350 let outcome = rt
2351 .hybrid_search_outcome(&tok, "FlashAttention", 10, None, None, &[], None)
2352 .await
2353 .expect("hybrid search must succeed");
2354
2355 assert!(!outcome.hits.is_empty(), "should find the entity");
2356 assert!(
2357 outcome.vector_error.is_none(),
2358 "a healthy vector arm must not report an error"
2359 );
2360 }
2361
2362 #[tokio::test]
2365 async fn hybrid_search_entity_hit_has_title() {
2366 let rt = KhiveRuntime::memory().unwrap();
2367 let tok = NamespaceToken::local();
2368 rt.create_entity(
2369 &tok,
2370 "concept",
2371 None,
2372 "FlashAttention",
2373 Some("IO-aware exact attention using tiling"),
2374 None,
2375 vec![],
2376 )
2377 .await
2378 .unwrap();
2379
2380 let hits = rt
2381 .hybrid_search(&tok, "FlashAttention", None, 10, None, None, &[], None)
2382 .await
2383 .unwrap();
2384
2385 assert!(!hits.is_empty(), "should find the entity");
2386 let hit = &hits[0];
2387 assert!(hit.title.is_some(), "title must be populated");
2388 assert!(
2389 hit.title.as_deref().unwrap().contains("FlashAttention"),
2390 "title must contain entity name"
2391 );
2392 }
2393
2394 #[tokio::test]
2400 async fn hybrid_search_with_dollar_sign_query_does_not_error() {
2401 let rt = KhiveRuntime::memory().unwrap();
2402 let tok = NamespaceToken::local();
2403 rt.create_entity(
2404 &tok,
2405 "concept",
2406 None,
2407 "DSL docs",
2408 Some("use $prev.id to chain calls"),
2409 None,
2410 vec![],
2411 )
2412 .await
2413 .unwrap();
2414
2415 let result = rt
2416 .hybrid_search(&tok, "$prev.id", None, 10, None, None, &[], None)
2417 .await;
2418
2419 assert!(
2420 result.is_ok(),
2421 "#388 hybrid_search must not hard-fail on a '$'-bearing query, got: {:?}",
2422 result.err()
2423 );
2424 }
2425
2426 #[tokio::test]
2435 async fn hybrid_search_with_residual_fts5_char_now_sanitized() {
2436 let rt = KhiveRuntime::memory().unwrap();
2437 let tok = NamespaceToken::local();
2438 rt.create_entity(
2439 &tok,
2440 "concept",
2441 None,
2442 "DSL docs",
2443 Some("use foo@bar to chain calls"),
2444 None,
2445 vec![],
2446 )
2447 .await
2448 .unwrap();
2449
2450 let result = rt
2451 .hybrid_search(&tok, "foo@bar", None, 10, None, None, &[], None)
2452 .await;
2453
2454 let hits = result.unwrap_or_else(|e| {
2455 panic!("#916 hybrid_search must not fail on an '@'-bearing query, got: {e:?}")
2456 });
2457 assert!(
2458 !hits.is_empty(),
2459 "#916 '@'-bearing query must still find the seeded 'foo@bar' content via the \
2460 quoted-phrase alternative"
2461 );
2462 }
2463
2464 #[tokio::test]
2472 async fn hybrid_search_with_916_issue_characters_finds_text_leg_hits() {
2473 let rt = KhiveRuntime::memory().unwrap();
2474 let tok = NamespaceToken::local();
2475
2476 rt.create_entity(
2477 &tok,
2478 "concept",
2479 None,
2480 "issue tracker",
2481 Some("tracking #682 Stage 2: MoE expert-cache prefetch work"),
2482 None,
2483 vec![],
2484 )
2485 .await
2486 .unwrap();
2487 rt.create_entity(
2488 &tok,
2489 "concept",
2490 None,
2491 "benchmark notes",
2492 Some("chunkwise B=128 traffic arithmetic simdgroup_matrix DPLR"),
2493 None,
2494 vec![],
2495 )
2496 .await
2497 .unwrap();
2498 rt.create_entity(
2499 &tok,
2500 "concept",
2501 None,
2502 "sampling notes",
2503 Some("evaluated with the Min-K%Prob membership inference method"),
2504 None,
2505 vec![],
2506 )
2507 .await
2508 .unwrap();
2509
2510 for query in ["#682 Stage 2", "B=128", "Min-K%Prob"] {
2511 let result = rt
2512 .hybrid_search(&tok, query, None, 10, None, None, &[], None)
2513 .await;
2514 let hits = result.unwrap_or_else(|e| {
2515 panic!("#916 hybrid_search must not fail on query {query:?}, got: {e:?}")
2516 });
2517 assert!(
2518 hits.iter()
2519 .any(|h| matches!(h.source, SearchSource::Text | SearchSource::Both)),
2520 "#916 query {query:?} must surface a Text/Both-sourced hit \
2521 (the FTS leg must contribute, not just the vector leg); got {hits:?}"
2522 );
2523 }
2524 }
2525
2526 #[tokio::test]
2540 async fn hybrid_search_tag_filter_pushed_before_truncation() {
2541 let rt = KhiveRuntime::memory().unwrap();
2542 let tok = NamespaceToken::local();
2543
2544 rt.create_entity(
2546 &tok,
2547 "concept",
2548 None,
2549 "alpha beta gamma decoy alpha beta gamma",
2550 Some("alpha beta gamma decoy description alpha beta gamma"),
2551 None,
2552 vec!["other-tag".to_string()],
2553 )
2554 .await
2555 .unwrap();
2556
2557 let target = rt
2559 .create_entity(
2560 &tok,
2561 "concept",
2562 None,
2563 "alpha beta gamma target",
2564 Some("alpha beta gamma target description"),
2565 None,
2566 vec!["target-tag".to_string()],
2567 )
2568 .await
2569 .unwrap();
2570
2571 let hits = rt
2575 .hybrid_search(
2576 &tok,
2577 "alpha beta gamma",
2578 None,
2579 1,
2580 None,
2581 None,
2582 &["target-tag".to_string()],
2583 None,
2584 )
2585 .await
2586 .unwrap();
2587
2588 assert_eq!(
2589 hits.len(),
2590 1,
2591 "exactly one hit expected (the tag-matching entity)"
2592 );
2593 assert_eq!(
2594 hits[0].entity_id, target.id,
2595 "the tag-filtered entity must be returned even when ranked below limit in raw fusion"
2596 );
2597 }
2598
2599 #[tokio::test]
2608 async fn hybrid_search_props_filter_pushed_before_truncation() {
2609 let rt = KhiveRuntime::memory().unwrap();
2610 let tok = NamespaceToken::local();
2611
2612 rt.create_entity(
2613 &tok,
2614 "concept",
2615 None,
2616 "delta epsilon zeta decoy delta epsilon zeta",
2617 Some("delta epsilon zeta decoy description delta epsilon zeta"),
2618 Some(serde_json::json!({"domain": "other"})),
2619 vec![],
2620 )
2621 .await
2622 .unwrap();
2623
2624 let target = rt
2625 .create_entity(
2626 &tok,
2627 "concept",
2628 None,
2629 "delta epsilon zeta target",
2630 Some("delta epsilon zeta target description"),
2631 Some(serde_json::json!({"domain": "target"})),
2632 vec![],
2633 )
2634 .await
2635 .unwrap();
2636
2637 let filter = serde_json::json!({"domain": "target"});
2638 let hits = rt
2639 .hybrid_search(
2640 &tok,
2641 "delta epsilon zeta",
2642 None,
2643 1,
2644 None,
2645 None,
2646 &[],
2647 Some(&filter),
2648 )
2649 .await
2650 .unwrap();
2651
2652 assert_eq!(hits.len(), 1, "exactly one hit expected (properties match)");
2653 assert_eq!(
2654 hits[0].entity_id, target.id,
2655 "the properties-filtered entity must be returned even when ranked below limit"
2656 );
2657 }
2658
2659 #[tokio::test]
2669 async fn hybrid_search_entity_kind_filter_pushed_into_text_arm() {
2670 let rt = KhiveRuntime::memory().unwrap();
2671 let tok = NamespaceToken::local();
2672
2673 for i in 0..12 {
2674 rt.create_entity(
2675 &tok,
2676 "concept",
2677 None,
2678 &format!("quillfeather decoy {i}"),
2679 Some("quillfeather quillfeather quillfeather"),
2680 None,
2681 vec![],
2682 )
2683 .await
2684 .unwrap();
2685 }
2686
2687 let description = format!(
2688 "A long administrative description that mentions quillfeather once. {}",
2689 "Unrelated filing, scheduling and review words. ".repeat(20)
2690 );
2691 let target = rt
2692 .create_entity(
2693 &tok,
2694 "document",
2695 None,
2696 "Archive filing report",
2697 Some(description.as_str()),
2698 None,
2699 vec![],
2700 )
2701 .await
2702 .unwrap();
2703
2704 let unfiltered = rt
2707 .hybrid_search(&tok, "quillfeather", None, 13, None, None, &[], None)
2708 .await
2709 .unwrap();
2710 let target_rank = unfiltered
2711 .iter()
2712 .position(|hit| hit.entity_id == target.id)
2713 .expect("the document is found without a kind filter");
2714 assert!(
2715 target_rank >= 4,
2716 "the document must rank behind the candidate window, got rank {target_rank}"
2717 );
2718
2719 let hits = rt
2720 .hybrid_search(
2721 &tok,
2722 "quillfeather",
2723 None,
2724 1,
2725 Some("document"),
2726 None,
2727 &[],
2728 None,
2729 )
2730 .await
2731 .unwrap();
2732
2733 assert_eq!(
2734 hits.len(),
2735 1,
2736 "the kind-filtered search must return the matching document"
2737 );
2738 assert_eq!(
2739 hits[0].entity_id, target.id,
2740 "the returned hit must be the document, not a concept"
2741 );
2742 }
2743
2744 #[tokio::test]
2752 async fn hybrid_search_each_kind_matches_per_kind_search_when_one_kind_dominates() {
2753 let rt = KhiveRuntime::memory().unwrap();
2754 let tok = NamespaceToken::local();
2755
2756 for i in 0..20 {
2757 rt.create_entity(
2758 &tok,
2759 "concept",
2760 None,
2761 &format!("quillfeather decoy {i}"),
2762 Some("quillfeather quillfeather quillfeather"),
2763 None,
2764 vec![],
2765 )
2766 .await
2767 .unwrap();
2768 }
2769 for i in 0..3 {
2770 let description = format!(
2771 "A long administrative description {i} that mentions quillfeather once. {}",
2772 "Unrelated filing, scheduling and review words. ".repeat(20)
2773 );
2774 rt.create_entity(
2775 &tok,
2776 "document",
2777 None,
2778 &format!("Archive filing report {i}"),
2779 Some(description.as_str()),
2780 None,
2781 vec![],
2782 )
2783 .await
2784 .unwrap();
2785 }
2786
2787 let kinds = ["concept", "document"];
2788 let mut expected: Vec<Vec<Uuid>> = Vec::new();
2789 for kind in kinds {
2790 let hits = rt
2791 .hybrid_search(&tok, "quillfeather", None, 2, Some(kind), None, &[], None)
2792 .await
2793 .unwrap();
2794 expected.push(hits.iter().map(|hit| hit.entity_id).collect());
2795 }
2796 assert_eq!(
2797 expected[0].len(),
2798 2,
2799 "premise: the concept list is cut to limit"
2800 );
2801 assert_eq!(
2802 expected[1].len(),
2803 2,
2804 "premise: the document list is not starved"
2805 );
2806
2807 let per_kind = rt
2808 .hybrid_search_each_kind(&tok, "quillfeather", None, 2, &kinds)
2809 .await
2810 .unwrap();
2811 let mut actual: Vec<Vec<Uuid>> = Vec::new();
2812 for hits in &per_kind {
2813 actual.push(hits.iter().map(|hit| hit.entity_id).collect());
2814 }
2815 assert_eq!(
2816 actual, expected,
2817 "each kind must get the ids and order its own per-kind search returns"
2818 );
2819 }
2820
2821 struct CapturingEmbeddingService {
2824 captured: std::sync::Arc<std::sync::Mutex<Vec<Vec<String>>>>,
2825 }
2826
2827 #[async_trait::async_trait]
2828 impl EmbeddingService for CapturingEmbeddingService {
2829 async fn embed(
2830 &self,
2831 texts: &[String],
2832 _model: EmbeddingModel,
2833 ) -> std::result::Result<Vec<Vec<f32>>, lattice_embed::EmbedError> {
2834 self.captured.lock().unwrap().push(texts.to_vec());
2835 Ok(texts.iter().map(|_| vec![1.0]).collect())
2836 }
2837
2838 fn supports_model(&self, _model: EmbeddingModel) -> bool {
2839 true
2840 }
2841
2842 fn name(&self) -> &'static str {
2843 "capturing-embedding-service"
2844 }
2845 }
2846
2847 struct CapturingEmbedderProvider {
2848 name: String,
2849 captured: std::sync::Arc<std::sync::Mutex<Vec<Vec<String>>>>,
2850 }
2851
2852 struct RewritingPassageService {
2853 captured: std::sync::Arc<std::sync::Mutex<Vec<Vec<String>>>>,
2854 }
2855
2856 #[async_trait::async_trait]
2857 impl EmbeddingService for RewritingPassageService {
2858 async fn embed(
2859 &self,
2860 texts: &[String],
2861 _model: EmbeddingModel,
2862 ) -> std::result::Result<Vec<Vec<f32>>, lattice_embed::EmbedError> {
2863 self.captured.lock().unwrap().push(texts.to_vec());
2864 Ok(texts.iter().map(|_| vec![1.0]).collect())
2865 }
2866
2867 async fn embed_passage(
2868 &self,
2869 texts: &[String],
2870 model: EmbeddingModel,
2871 ) -> std::result::Result<Vec<Vec<f32>>, lattice_embed::EmbedError> {
2872 let prepared: Vec<String> = texts.iter().map(|text| format!("custom:{text}")).collect();
2873 self.embed(&prepared, model).await
2874 }
2875
2876 fn supports_model(&self, _model: EmbeddingModel) -> bool {
2877 true
2878 }
2879
2880 fn name(&self) -> &'static str {
2881 "rewriting-passage-service"
2882 }
2883 }
2884
2885 struct RewritingPassageProvider {
2886 name: String,
2887 captured: std::sync::Arc<std::sync::Mutex<Vec<Vec<String>>>>,
2888 }
2889
2890 #[async_trait::async_trait]
2891 impl EmbedderProvider for RewritingPassageProvider {
2892 fn name(&self) -> &str {
2893 &self.name
2894 }
2895
2896 fn dimensions(&self) -> usize {
2897 1
2898 }
2899
2900 async fn build(&self) -> crate::error::RuntimeResult<std::sync::Arc<dyn EmbeddingService>> {
2901 Ok(std::sync::Arc::new(RewritingPassageService {
2902 captured: std::sync::Arc::clone(&self.captured),
2903 }))
2904 }
2905 }
2906
2907 struct WrongCardinalityEmbeddingService;
2908
2909 #[async_trait::async_trait]
2910 impl EmbeddingService for WrongCardinalityEmbeddingService {
2911 async fn embed(
2912 &self,
2913 texts: &[String],
2914 _model: EmbeddingModel,
2915 ) -> std::result::Result<Vec<Vec<f32>>, lattice_embed::EmbedError> {
2916 Ok(texts
2917 .iter()
2918 .take(texts.len().saturating_sub(1))
2919 .map(|_| vec![1.0])
2920 .collect())
2921 }
2922
2923 fn supports_model(&self, _model: EmbeddingModel) -> bool {
2924 true
2925 }
2926
2927 fn name(&self) -> &'static str {
2928 "wrong-cardinality-embedding-service"
2929 }
2930 }
2931
2932 struct WrongCardinalityEmbedderProvider;
2933
2934 #[async_trait::async_trait]
2935 impl EmbedderProvider for WrongCardinalityEmbedderProvider {
2936 fn name(&self) -> &str {
2937 "wrong-cardinality-embedding-service"
2938 }
2939
2940 fn dimensions(&self) -> usize {
2941 1
2942 }
2943
2944 async fn build(&self) -> crate::error::RuntimeResult<std::sync::Arc<dyn EmbeddingService>> {
2945 Ok(std::sync::Arc::new(WrongCardinalityEmbeddingService))
2946 }
2947 }
2948
2949 struct SurplusCardinalityEmbeddingService;
2950
2951 #[async_trait::async_trait]
2952 impl EmbeddingService for SurplusCardinalityEmbeddingService {
2953 async fn embed(
2954 &self,
2955 texts: &[String],
2956 _model: EmbeddingModel,
2957 ) -> std::result::Result<Vec<Vec<f32>>, lattice_embed::EmbedError> {
2958 let mut vectors: Vec<Vec<f32>> = texts.iter().map(|_| vec![1.0]).collect();
2959 vectors.push(vec![1.0]);
2960 Ok(vectors)
2961 }
2962
2963 fn supports_model(&self, _model: EmbeddingModel) -> bool {
2964 true
2965 }
2966
2967 fn name(&self) -> &'static str {
2968 "surplus-cardinality-embedding-service"
2969 }
2970 }
2971
2972 struct SurplusCardinalityEmbedderProvider;
2973
2974 #[async_trait::async_trait]
2975 impl EmbedderProvider for SurplusCardinalityEmbedderProvider {
2976 fn name(&self) -> &str {
2977 "surplus-cardinality-embedding-service"
2978 }
2979
2980 fn dimensions(&self) -> usize {
2981 1
2982 }
2983
2984 async fn build(&self) -> crate::error::RuntimeResult<std::sync::Arc<dyn EmbeddingService>> {
2985 Ok(std::sync::Arc::new(SurplusCardinalityEmbeddingService))
2986 }
2987 }
2988
2989 #[async_trait::async_trait]
2990 impl EmbedderProvider for CapturingEmbedderProvider {
2991 fn name(&self) -> &str {
2992 &self.name
2993 }
2994
2995 fn dimensions(&self) -> usize {
2996 1
2997 }
2998
2999 async fn build(&self) -> crate::error::RuntimeResult<std::sync::Arc<dyn EmbeddingService>> {
3000 Ok(std::sync::Arc::new(CapturingEmbeddingService {
3001 captured: std::sync::Arc::clone(&self.captured),
3002 }))
3003 }
3004 }
3005
3006 fn runtime_with_capturing_embedder(
3007 model: EmbeddingModel,
3008 ) -> (
3009 KhiveRuntime,
3010 std::sync::Arc<std::sync::Mutex<Vec<Vec<String>>>>,
3011 ) {
3012 let runtime = KhiveRuntime::memory().unwrap();
3013 let captured = std::sync::Arc::new(std::sync::Mutex::new(Vec::new()));
3014 runtime.register_embedder(CapturingEmbedderProvider {
3015 name: model.to_string(),
3016 captured: std::sync::Arc::clone(&captured),
3017 });
3018 (runtime, captured)
3019 }
3020
3021 #[tokio::test]
3022 async fn bge_query_paths_pass_raw_unprefixed_text() {
3023 const BGE_QUERY_INSTRUCTION: &str =
3024 "Represent this sentence for searching relevant passages: ";
3025 let single = "single raw query";
3026 let batch = vec![
3027 "first raw query".to_string(),
3028 "second raw query".to_string(),
3029 ];
3030
3031 for model in [
3032 EmbeddingModel::BgeSmallEnV15,
3033 EmbeddingModel::BgeBaseEnV15,
3034 EmbeddingModel::BgeLargeEnV15,
3035 ] {
3036 let (runtime, captured) = runtime_with_capturing_embedder(model);
3037 runtime
3038 .embed_query_with_model(&model.to_string(), single)
3039 .await
3040 .unwrap();
3041 runtime
3042 .embed_query_batch_with_model(&model.to_string(), &batch)
3043 .await
3044 .unwrap();
3045
3046 let calls = captured.lock().unwrap().clone();
3047 assert_eq!(
3048 calls,
3049 vec![vec![single.to_string()], batch.clone()],
3050 "{model} must receive raw query text through single and batch paths"
3051 );
3052 assert!(
3053 calls
3054 .iter()
3055 .flatten()
3056 .all(|text| !text.contains(BGE_QUERY_INSTRUCTION)),
3057 "{model} must not receive the BGE retrieval instruction"
3058 );
3059 }
3060 }
3061
3062 #[tokio::test]
3063 async fn e5_query_paths_apply_query_prefix() {
3064 let model = EmbeddingModel::MultilingualE5Small;
3065 let single = "single raw query";
3066 let batch = vec![
3067 "first raw query".to_string(),
3068 "second raw query".to_string(),
3069 ];
3070 let (runtime, captured) = runtime_with_capturing_embedder(model);
3071
3072 runtime
3073 .embed_query_with_model(&model.to_string(), single)
3074 .await
3075 .unwrap();
3076 runtime
3077 .embed_query_batch_with_model(&model.to_string(), &batch)
3078 .await
3079 .unwrap();
3080
3081 assert_eq!(
3082 captured.lock().unwrap().as_slice(),
3083 [
3084 vec!["query: single raw query".to_string()],
3085 vec![
3086 "query: first raw query".to_string(),
3087 "query: second raw query".to_string(),
3088 ],
3089 ],
3090 "E5 must receive its query prefix through single and batch paths"
3091 );
3092 }
3093
3094 #[tokio::test]
3095 async fn mixed_document_batch_stays_one_ordered_provider_call() {
3096 let model = EmbeddingModel::AllMiniLmL6V2;
3097 let (runtime, captured) = runtime_with_capturing_embedder(model);
3098 let texts = vec![
3099 "first normal document".to_string(),
3100 "x".repeat(MAX_TEXT_BYTES + 1),
3101 "second normal document".to_string(),
3102 ];
3103
3104 let outcomes = runtime
3105 .embed_document_batch_with_model_outcomes(&model.to_string(), &texts)
3106 .await
3107 .expect("mixed batch must embed");
3108 assert_eq!(outcomes.len(), texts.len());
3109 assert!(!outcomes[0].truncated);
3110 assert_eq!(outcomes[0].source_bytes, outcomes[0].embedded_bytes);
3111 assert!(outcomes[1].truncated);
3112 assert_eq!(outcomes[1].source_bytes, MAX_TEXT_BYTES + 1);
3113 assert_eq!(outcomes[1].embedded_bytes, MAX_TEXT_BYTES);
3114 assert!(!outcomes[2].truncated);
3115 assert_eq!(
3116 captured.lock().unwrap().as_slice(),
3117 [vec![
3118 texts[0].clone(),
3119 "x".repeat(MAX_TEXT_BYTES),
3120 texts[2].clone(),
3121 ]]
3122 );
3123 }
3124
3125 #[tokio::test]
3126 async fn canonical_name_override_has_no_exact_input_fingerprint() {
3127 let model = EmbeddingModel::MultilingualE5Small;
3128 let model_name = model.to_string();
3129 let (runtime, captured) = runtime_with_capturing_embedder(model);
3130 let input = "x".repeat(document_embedding_budget(&model_name) + 3);
3131 let outcomes = runtime
3132 .embed_document_batch_with_model_outcomes(&model_name, std::slice::from_ref(&input))
3133 .await
3134 .unwrap();
3135 let calls = captured.lock().unwrap();
3136 let actual_input = &calls[0][0];
3137 assert!(actual_input.starts_with("passage: "));
3138 assert_eq!(actual_input.len(), MAX_TEXT_BYTES);
3139 assert!(outcomes[0].truncated);
3140 assert_eq!(outcomes[0].prepared_text_fingerprint, None);
3141 assert_eq!(
3142 prepared_document_fingerprint(
3143 bounded_embedding_input(&input, document_embedding_budget(&model_name)).0,
3144 model,
3145 ),
3146 VectorRecord::fingerprint_text(actual_input),
3147 "the audited lattice path hashes the bounded text plus passage prefix"
3148 );
3149 }
3150
3151 #[tokio::test]
3152 async fn ordinary_custom_document_provider_has_no_exact_input_fingerprint() {
3153 let runtime = KhiveRuntime::memory().unwrap();
3154 let captured = std::sync::Arc::new(std::sync::Mutex::new(Vec::new()));
3155 runtime.register_embedder(CapturingEmbedderProvider {
3156 name: "custom-document-provider".to_string(),
3157 captured: std::sync::Arc::clone(&captured),
3158 });
3159
3160 let outcome = runtime
3161 .embed_document_with_model_outcome("custom-document-provider", "document")
3162 .await
3163 .expect("custom provider must embed");
3164 assert_eq!(outcome.prepared_text_fingerprint, None);
3165 assert_eq!(captured.lock().unwrap().len(), 1);
3166 }
3167
3168 #[tokio::test]
3169 async fn custom_passage_override_under_builtin_name_cannot_claim_exact_input() {
3170 let runtime = KhiveRuntime::memory().unwrap();
3171 let captured = std::sync::Arc::new(std::sync::Mutex::new(Vec::new()));
3172 let model = EmbeddingModel::MultilingualE5Small;
3173 let model_name = model.to_string();
3174 runtime.register_embedder(RewritingPassageProvider {
3175 name: model_name.clone(),
3176 captured: std::sync::Arc::clone(&captured),
3177 });
3178
3179 let outcome = runtime
3180 .embed_document_with_model_outcome(&model_name, "document")
3181 .await
3182 .expect("custom override must embed");
3183 assert_eq!(outcome.prepared_text_fingerprint, None);
3184 assert_eq!(
3185 captured.lock().unwrap()[0],
3186 vec!["custom:document".to_string()]
3187 );
3188 assert_ne!(
3189 prepared_document_fingerprint("document", model),
3190 VectorRecord::fingerprint_text("custom:document")
3191 );
3192 }
3193
3194 #[tokio::test]
3195 async fn truncated_document_batch_rejects_provider_cardinality_mismatch() {
3196 let runtime = KhiveRuntime::memory().unwrap();
3197 runtime.register_embedder(WrongCardinalityEmbedderProvider);
3198 let model_name = "wrong-cardinality-embedding-service";
3199 let texts = vec!["x".repeat(MAX_TEXT_BYTES + 1), "normal".to_string()];
3200
3201 let error = runtime
3202 .embed_document_batch_with_model_outcomes(model_name, &texts)
3203 .await
3204 .expect_err("provider cardinality mismatch must fail the whole batch");
3205
3206 assert!(
3207 error
3208 .to_string()
3209 .contains("embed_passage returned 1 vectors for 2 inputs"),
3210 "unexpected error: {error}"
3211 );
3212 }
3213
3214 #[tokio::test]
3215 async fn singleton_document_embed_rejects_zero_vectors() {
3216 let runtime = KhiveRuntime::memory().unwrap();
3217 runtime.register_embedder(WrongCardinalityEmbedderProvider);
3218 let model_name = "wrong-cardinality-embedding-service";
3219
3220 let error = runtime
3221 .embed_document_with_model_outcome(model_name, "single document")
3222 .await
3223 .expect_err("provider returning zero vectors must fail closed");
3224
3225 assert!(
3226 error
3227 .to_string()
3228 .contains("embed_passage returned 0 vectors for 1 input"),
3229 "unexpected error: {error}"
3230 );
3231 }
3232
3233 #[tokio::test]
3234 async fn singleton_document_embed_rejects_surplus_vectors() {
3235 let runtime = KhiveRuntime::memory().unwrap();
3236 runtime.register_embedder(SurplusCardinalityEmbedderProvider);
3237 let model_name = "surplus-cardinality-embedding-service";
3238
3239 let error = runtime
3240 .embed_document_with_model_outcome(model_name, "single document")
3241 .await
3242 .expect_err("provider returning surplus vectors must fail closed");
3243
3244 assert!(
3245 error
3246 .to_string()
3247 .contains("embed_passage returned 2 vectors for 1 input"),
3248 "unexpected error: {error}"
3249 );
3250 }
3251
3252 #[test]
3253 #[ignore = "loads ~80 MB model; run with --include-ignored"]
3254 fn minilm_document_and_query_embed_are_identical_no_prefix_model() {
3255 let model = EmbeddingModel::AllMiniLmL6V2;
3259 let config = RuntimeConfig {
3260 db_path: None,
3261 default_namespace: Namespace::parse("test").unwrap(),
3262 embedding_model: Some(model),
3263 packs: vec!["kg".to_string()],
3264 ..RuntimeConfig::default()
3265 };
3266 let rt = KhiveRuntime::new(config).unwrap();
3267 let text = "attention is all you need".to_string();
3268 let rt_ref = &rt;
3269 let (doc_emb, query_emb) = tokio::runtime::Runtime::new().unwrap().block_on(async {
3270 let d = rt_ref
3271 .embed_document_with_model_outcome(&model.to_string(), &text)
3272 .await
3273 .unwrap()
3274 .vector;
3275 let q = rt_ref
3276 .embed_query_with_model(&model.to_string(), &text)
3277 .await
3278 .unwrap();
3279 (d, q)
3280 });
3281 assert_eq!(
3282 doc_emb, query_emb,
3283 "MiniLM has no instruction prefix: document and query embeds must be identical"
3284 );
3285 }
3286
3287 #[test]
3288 #[ignore = "loads multilingual-e5-small (~90 MB); run with --include-ignored"]
3289 fn e5_document_and_query_embed_differ_instruction_tuned_model() {
3290 let model = EmbeddingModel::MultilingualE5Small;
3295 let config = RuntimeConfig {
3296 db_path: None,
3297 default_namespace: Namespace::parse("test").unwrap(),
3298 embedding_model: Some(model),
3299 packs: vec!["kg".to_string()],
3300 ..RuntimeConfig::default()
3301 };
3302 let rt = KhiveRuntime::new(config).unwrap();
3303 let text = "attention is all you need".to_string();
3304 let rt_ref = &rt;
3305 let (doc_emb, query_emb) = tokio::runtime::Runtime::new().unwrap().block_on(async {
3306 let d = rt_ref
3307 .embed_document_with_model_outcome(&model.to_string(), &text)
3308 .await
3309 .unwrap()
3310 .vector;
3311 let q = rt_ref
3312 .embed_query_with_model(&model.to_string(), &text)
3313 .await
3314 .unwrap();
3315 (d, q)
3316 });
3317 assert_ne!(
3318 doc_emb, query_emb,
3319 "multilingual-e5-small uses asymmetric prefixes: document ('passage: ') \
3320 and query ('query: ') embeds of the same text must differ"
3321 );
3322 }
3323
3324 use crate::embedder_registry::EmbedderProvider;
3327 use lattice_embed::EmbeddingService;
3328
3329 struct StubEmbedderProvider;
3334
3335 #[async_trait::async_trait]
3336 impl EmbedderProvider for StubEmbedderProvider {
3337 fn name(&self) -> &str {
3338 "stub-model-m07"
3339 }
3340
3341 fn dimensions(&self) -> usize {
3342 4
3343 }
3344
3345 async fn build(&self) -> crate::error::RuntimeResult<std::sync::Arc<dyn EmbeddingService>> {
3346 struct StubSvc;
3347 #[async_trait::async_trait]
3348 impl EmbeddingService for StubSvc {
3349 async fn embed(
3350 &self,
3351 _texts: &[String],
3352 _model: lattice_embed::EmbeddingModel,
3353 ) -> std::result::Result<Vec<Vec<f32>>, lattice_embed::EmbedError> {
3354 Ok(vec![])
3355 }
3356
3357 fn supports_model(&self, _model: lattice_embed::EmbeddingModel) -> bool {
3358 true
3359 }
3360
3361 fn name(&self) -> &'static str {
3362 "stub-svc-m07"
3363 }
3364 }
3365 Ok(std::sync::Arc::new(StubSvc))
3366 }
3367 }
3368
3369 #[tokio::test]
3377 async fn backfill_reader_error_is_propagated_not_swallowed() {
3378 let rt = KhiveRuntime::memory().unwrap();
3379 rt.register_embedder(StubEmbedderProvider);
3380 let tok = NamespaceToken::local();
3381
3382 super::arm_backfill_reader_fail();
3385
3386 let result = rt.backfill_missing_embeddings(&tok).await;
3387 assert!(
3388 result.is_err(),
3389 "backfill_missing_embeddings must propagate the reader error (got Ok instead)"
3390 );
3391 let err_msg = result.unwrap_err().to_string();
3392 assert!(
3393 err_msg.contains("injected failure"),
3394 "error must originate from the injected reader failure, got: {err_msg}"
3395 );
3396 }
3397
3398 #[tokio::test]
3399 async fn backfill_skips_a_full_page_of_excluded_messages_and_reaches_eligible_tail() {
3400 use crate::{NoteEmbeddingPolicy, NoteEmbeddingPolicySpec};
3401 use khive_storage::note::Note;
3402
3403 const PAGE: u128 = EMBEDDING_BATCH_PAGE_SIZE as u128;
3404 let primary = EmbeddingModel::AllMiniLmL6V2;
3405 let primary_name = primary.to_string();
3406 let secondary_name = "zz-backfill-secondary";
3407 let rt = KhiveRuntime::new(RuntimeConfig {
3408 db_path: None,
3409 embedding_model: Some(primary),
3410 packs: vec![],
3411 ..RuntimeConfig::no_embeddings()
3412 })
3413 .unwrap();
3414 rt.register_embedder(ConstantEmbedderProvider {
3415 name: primary_name.clone(),
3416 dimensions: primary.dimensions(),
3417 });
3418 rt.register_embedder(ConstantEmbedderProvider {
3419 name: secondary_name.into(),
3420 dimensions: 4,
3421 });
3422 rt.install_note_embedding_policies(&[NoteEmbeddingPolicySpec {
3423 kind: "message",
3424 policy: NoteEmbeddingPolicy::DefaultModel,
3425 }]);
3426 let tok = NamespaceToken::local();
3427 let primary_store = rt.vectors_for_model(&tok, &primary_name).unwrap();
3428 let secondary_store = rt.vectors_for_model(&tok, secondary_name).unwrap();
3429 let notes = rt.notes(&tok).unwrap();
3430 let mut seeded = Vec::new();
3431 for ordinal in 1..=PAGE + 1 {
3432 let mut message = Note::new("local", "message", "excluded from secondary");
3433 message.id = Uuid::from_u128(ordinal);
3434 seeded.push(message);
3435 }
3436 let mut ordinary = Note::new("local", "observation", "eligible after full page");
3437 ordinary.id = Uuid::from_u128(u128::MAX);
3438 seeded.push(ordinary.clone());
3439 notes.upsert_notes(seeded).await.unwrap();
3440
3441 let backfilled = tokio::time::timeout(
3442 std::time::Duration::from_secs(60),
3443 rt.backfill_missing_embeddings(&tok),
3444 )
3445 .await
3446 .expect("backfill must advance past a full excluded page")
3447 .unwrap();
3448 assert_eq!(backfilled, (PAGE + 2) as u64 + 1);
3449 assert_eq!(primary_store.count().await.unwrap(), (PAGE + 2) as u64);
3450 assert_eq!(secondary_store.count().await.unwrap(), 1);
3451 let excluded_document = rt
3452 .text_for_notes(&tok)
3453 .unwrap()
3454 .get_document("local", Uuid::from_u128(1))
3455 .await
3456 .unwrap()
3457 .expect("first-model pass must index an excluded message");
3458 assert_eq!(excluded_document.body, "excluded from secondary");
3459 assert!(
3460 rt.text_for_notes(&tok)
3461 .unwrap()
3462 .get_document("local", ordinary.id)
3463 .await
3464 .unwrap()
3465 .is_some(),
3466 "the first-model pass must still repopulate FTS"
3467 );
3468 }
3469
3470 const BACKFILL_BATCH_MODEL: &str = "backfill-batch-model";
3471
3472 struct CountingBackfillService {
3473 calls: Arc<AtomicUsize>,
3474 reject_poison: bool,
3475 }
3476
3477 #[async_trait::async_trait]
3478 impl EmbeddingService for CountingBackfillService {
3479 async fn embed(
3480 &self,
3481 texts: &[String],
3482 _model: EmbeddingModel,
3483 ) -> std::result::Result<Vec<Vec<f32>>, EmbedError> {
3484 self.calls.fetch_add(1, Ordering::SeqCst);
3485 if self.reject_poison
3486 && (texts.len() > 1 || texts.iter().any(|text| text.contains("poison")))
3487 {
3488 return Err(EmbedError::InferenceFailed("poison input".into()));
3489 }
3490 Ok(texts
3491 .iter()
3492 .map(|text| {
3493 vec![
3494 text.len() as f32,
3495 text.bytes().map(u32::from).sum::<u32>() as f32,
3496 text.as_bytes().first().copied().unwrap_or_default() as f32,
3497 text.as_bytes().last().copied().unwrap_or_default() as f32,
3498 ]
3499 })
3500 .collect())
3501 }
3502
3503 fn supports_model(&self, _model: EmbeddingModel) -> bool {
3504 true
3505 }
3506
3507 fn name(&self) -> &'static str {
3508 "backfill-batch-counting-service"
3509 }
3510 }
3511
3512 struct CountingBackfillProvider {
3513 calls: Arc<AtomicUsize>,
3514 reject_poison: bool,
3515 }
3516
3517 #[async_trait::async_trait]
3518 impl EmbedderProvider for CountingBackfillProvider {
3519 fn name(&self) -> &str {
3520 BACKFILL_BATCH_MODEL
3521 }
3522
3523 fn dimensions(&self) -> usize {
3524 4
3525 }
3526
3527 async fn build(&self) -> crate::error::RuntimeResult<Arc<dyn EmbeddingService>> {
3528 Ok(Arc::new(CountingBackfillService {
3529 calls: Arc::clone(&self.calls),
3530 reject_poison: self.reject_poison,
3531 }))
3532 }
3533 }
3534
3535 #[tokio::test]
3536 async fn backfill_batches_provider_calls_and_matches_single_record_vectors() {
3537 let token = NamespaceToken::local();
3538 let runtime = KhiveRuntime::memory().unwrap();
3539 let mut entity_ids = Vec::new();
3540 for index in 0..257 {
3541 let entity = runtime
3542 .create_entity(
3543 &token,
3544 "concept",
3545 None,
3546 &format!("Backfill entity {index}"),
3547 (index != 0).then_some("body"),
3548 None,
3549 vec![],
3550 )
3551 .await
3552 .unwrap();
3553 entity_ids.push(entity.id);
3554 }
3555 let mut note_ids = Vec::new();
3556 for index in 0..3 {
3557 let note = runtime
3558 .create_note(
3559 &token,
3560 "observation",
3561 None,
3562 &format!("Backfill note {index}"),
3563 None,
3564 None,
3565 vec![],
3566 )
3567 .await
3568 .unwrap();
3569 note_ids.push(note.id);
3570 }
3571 let calls = Arc::new(AtomicUsize::new(0));
3572 runtime.register_embedder(CountingBackfillProvider {
3573 calls: Arc::clone(&calls),
3574 reject_poison: false,
3575 });
3576 let usage = crate::usage::UsageContext::new();
3577 let backfilled =
3578 crate::usage::scope(usage.clone(), runtime.backfill_missing_embeddings(&token))
3579 .await
3580 .unwrap();
3581 assert_eq!(backfilled, 260);
3582 assert_eq!(
3583 calls.load(Ordering::SeqCst),
3584 3,
3585 "two entity pages and one note page"
3586 );
3587 assert_eq!(
3588 usage.snapshot()["embed_calls"],
3589 260,
3590 "ADR-103 counts texts, not provider calls"
3591 );
3592
3593 let singles = KhiveRuntime::memory().unwrap();
3594 let single_calls = Arc::new(AtomicUsize::new(0));
3595 singles.register_embedder(CountingBackfillProvider {
3596 calls: Arc::clone(&single_calls),
3597 reject_poison: false,
3598 });
3599 for id in &entity_ids {
3600 let entity = runtime.get_entity(&token, *id).await.unwrap();
3601 singles
3602 .entities(&token)
3603 .unwrap()
3604 .upsert_entity(entity.clone())
3605 .await
3606 .unwrap();
3607 singles.reindex_entity(&token, &entity).await.unwrap();
3608 }
3609 for id in ¬e_ids {
3610 let note = runtime
3611 .notes(&token)
3612 .unwrap()
3613 .get_note(*id)
3614 .await
3615 .unwrap()
3616 .unwrap();
3617 singles
3618 .notes(&token)
3619 .unwrap()
3620 .upsert_note(note.clone())
3621 .await
3622 .unwrap();
3623 singles.reindex_note(&token, ¬e).await.unwrap();
3624 }
3625 assert_eq!(single_calls.load(Ordering::SeqCst), 260);
3626 let batched_entities = runtime
3627 .vectors_for_model(&token, BACKFILL_BATCH_MODEL)
3628 .unwrap()
3629 .get_vectors(&entity_ids, "local", "entity.body")
3630 .await
3631 .unwrap();
3632 let single_entities = singles
3633 .vectors_for_model(&token, BACKFILL_BATCH_MODEL)
3634 .unwrap()
3635 .get_vectors(&entity_ids, "local", "entity.body")
3636 .await
3637 .unwrap();
3638 assert_eq!(batched_entities.len(), 257);
3639 assert_eq!(batched_entities, single_entities);
3640 let batched_notes = runtime
3641 .vectors_for_model(&token, BACKFILL_BATCH_MODEL)
3642 .unwrap()
3643 .get_vectors(¬e_ids, "local", "note.content")
3644 .await
3645 .unwrap();
3646 let single_notes = singles
3647 .vectors_for_model(&token, BACKFILL_BATCH_MODEL)
3648 .unwrap()
3649 .get_vectors(¬e_ids, "local", "note.content")
3650 .await
3651 .unwrap();
3652 assert_eq!(batched_notes.len(), 3);
3653 assert_eq!(batched_notes, single_notes);
3654 }
3655
3656 #[tokio::test]
3657 async fn backfill_failed_pages_retry_singly_without_duplicate_index_rows() {
3658 use khive_storage::types::{SqlStatement, SqlValue};
3659
3660 let token = NamespaceToken::local();
3661 let runtime = KhiveRuntime::memory().unwrap();
3662 for name in ["good one", "poison", "good two"] {
3663 runtime
3664 .create_entity(&token, "concept", None, name, Some("body"), None, vec![])
3665 .await
3666 .unwrap();
3667 runtime
3668 .create_note(&token, "observation", None, name, None, None, vec![])
3669 .await
3670 .unwrap();
3671 }
3672 let calls = Arc::new(AtomicUsize::new(0));
3673 runtime.register_embedder(CountingBackfillProvider {
3674 calls: Arc::clone(&calls),
3675 reject_poison: true,
3676 });
3677 let backfilled = runtime.backfill_missing_embeddings(&token).await.unwrap();
3678 assert_eq!(backfilled, 4);
3679 assert_eq!(
3680 calls.load(Ordering::SeqCst),
3681 8,
3682 "two failed pages plus six singleton retries"
3683 );
3684 let mut reader = runtime.sql().reader().await.unwrap();
3685 for (table, expected) in [("ann_write_log", 4), ("vector_provenance", 0)] {
3686 let count = reader
3687 .query_scalar(SqlStatement {
3688 sql: format!("SELECT COUNT(*) FROM {table} WHERE namespace = ?1"),
3689 params: vec![SqlValue::Text("local".into())],
3690 label: Some("backfill-batch-no-duplicate-rows".into()),
3691 })
3692 .await
3693 .unwrap();
3694 assert!(
3695 matches!(count, Some(SqlValue::Integer(value)) if value == expected),
3696 "{table} must match the guarded per-record writer"
3697 );
3698 }
3699 assert_eq!(
3700 runtime.backfill_missing_embeddings(&token).await.unwrap(),
3701 0
3702 );
3703 }
3704}