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