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::error::{RuntimeError, RuntimeResult};
11use crate::runtime::{KhiveRuntime, NamespaceToken};
12use khive_score::{rrf_score, DeterministicScore};
13use khive_storage::types::{
14 PageRequest, TextFilter, TextQueryMode, TextSearchHit, TextSearchRequest, VectorSearchHit,
15 VectorSearchRequest,
16};
17use khive_storage::EntityFilter;
18use khive_types::SubstrateKind;
19
20#[cfg(any(test, feature = "fault-injection"))]
22std::thread_local! {
23 static BACKFILL_READER_FAIL: std::cell::Cell<bool> = const { std::cell::Cell::new(false) };
24}
25
26#[cfg(any(test, feature = "fault-injection"))]
32pub fn arm_backfill_reader_fail() {
33 BACKFILL_READER_FAIL.with(|c| c.set(true));
34}
35
36#[derive(Clone, Copy, Debug, PartialEq, Eq)]
38pub enum RankScoreKind {
39 Rrf,
40 Vector,
41 Keyword,
42 Weighted,
43 Union,
44}
45
46impl RankScoreKind {
47 #[must_use]
48 pub const fn as_str(self) -> &'static str {
49 match self {
50 Self::Rrf => "rrf",
51 Self::Vector => "vector",
52 Self::Keyword => "keyword",
53 Self::Weighted => "weighted",
54 Self::Union => "union",
55 }
56 }
57}
58
59#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
64pub struct SearchSignals {
65 pub vector_similarity: Option<DeterministicScore>,
66 pub keyword_score: Option<DeterministicScore>,
67}
68
69#[derive(Clone, Debug)]
71pub struct SearchHit {
72 pub entity_id: Uuid,
73 pub score: DeterministicScore,
74 pub rank_score_kind: RankScoreKind,
75 pub signals: SearchSignals,
76 pub source: SearchSource,
77 pub title: Option<String>,
78 pub snippet: Option<String>,
79}
80
81#[derive(Clone, Debug)]
85pub struct HybridSearchOutcome {
86 pub hits: Vec<SearchHit>,
87 pub vector_error: Option<String>,
88}
89
90#[derive(Clone, Copy, Debug, PartialEq, Eq)]
92pub enum SearchSource {
93 Vector,
94 Text,
95 Both,
96}
97
98impl SearchSource {
99 #[must_use]
101 pub const fn union(self, other: Self) -> Self {
102 match (self, other) {
103 (Self::Text, Self::Text) => Self::Text,
104 (Self::Vector, Self::Vector) => Self::Vector,
105 _ => Self::Both,
106 }
107 }
108
109 #[must_use]
111 pub const fn as_str(self) -> &'static str {
112 match self {
113 Self::Vector => "vector",
114 Self::Text => "text",
115 Self::Both => "both",
116 }
117 }
118}
119
120const RRF_K: usize = 10;
127
128const CANDIDATE_MULTIPLIER: u32 = 4;
130
131pub const EMBEDDING_INPUT_TRUNCATED_WARNING: &str =
133 "embedding input was truncated to the embedder maximum; full content was stored unchanged";
134
135#[derive(Clone, Debug)]
137pub struct DocumentEmbeddingOutcome {
138 pub model_name: String,
139 pub vector: Vec<f32>,
140 pub source_bytes: usize,
141 pub embedded_bytes: usize,
142 pub truncated: bool,
143}
144
145#[derive(Clone, Debug, Default, PartialEq, Eq, serde::Deserialize, serde::Serialize)]
147pub struct EmbeddingTruncationReport {
148 pub truncated: u64,
149 pub discarded_bytes: u64,
150}
151
152impl EmbeddingTruncationReport {
153 pub fn observe(&mut self, outcome: &DocumentEmbeddingOutcome) {
154 if outcome.truncated {
155 self.truncated += 1;
156 self.discarded_bytes +=
157 outcome.source_bytes.saturating_sub(outcome.embedded_bytes) as u64;
158 }
159 }
160
161 #[must_use]
162 pub const fn any_truncated(&self) -> bool {
163 self.truncated > 0
164 }
165
166 pub fn merge(&mut self, other: Self) {
168 self.truncated = self.truncated.saturating_add(other.truncated);
169 self.discarded_bytes = self.discarded_bytes.saturating_add(other.discarded_bytes);
170 }
171}
172
173pub fn document_embedding_budget(model_name: &str) -> usize {
175 parse_embedding_model_alias(model_name)
176 .and_then(|model| model.document_instruction())
177 .map_or(MAX_TEXT_BYTES, |prefix| {
178 MAX_TEXT_BYTES.saturating_sub(prefix.len())
179 })
180}
181
182pub fn bounded_embedding_input(text: &str, max_bytes: usize) -> (&str, bool) {
184 if text.len() <= max_bytes {
185 return (text, false);
186 }
187
188 let end = text
189 .char_indices()
190 .map(|(index, _)| index)
191 .take_while(|index| *index <= max_bytes)
192 .last()
193 .unwrap_or(0);
194 (&text[..end], true)
195}
196
197impl KhiveRuntime {
198 pub async fn embed(&self, text: &str) -> RuntimeResult<Vec<f32>> {
203 let model_name = self.default_embedder_name();
204 if model_name.is_empty() {
205 return Err(RuntimeError::Unconfigured("embedding_model".into()));
206 }
207 self.embed_with_model(model_name, text).await
208 }
209
210 pub async fn embed_with_model(&self, model_name: &str, text: &str) -> RuntimeResult<Vec<f32>> {
226 let model = parse_embedding_model_alias(model_name);
227 let service = self.embedder(model_name).await?;
228 let emb_model = model.unwrap_or_default();
229 crate::usage::count(crate::usage::UsageUnit::EmbedCalls, 1);
233 let out = service.embed_one(text, emb_model).await;
234 Ok(out?)
235 }
236
237 pub async fn embed_document_with_model(
256 &self,
257 model_name: &str,
258 text: &str,
259 ) -> RuntimeResult<Vec<f32>> {
260 Ok(self
261 .embed_document_with_model_outcome_inner(None, model_name, text)
262 .await?
263 .vector)
264 }
265
266 pub async fn embed_document_with_model_outcome(
267 &self,
268 model_name: &str,
269 text: &str,
270 ) -> RuntimeResult<DocumentEmbeddingOutcome> {
271 self.embed_document_with_model_outcome_inner(None, model_name, text)
272 .await
273 }
274
275 pub(crate) async fn embed_document_with_model_outcome_for_token(
276 &self,
277 token: &NamespaceToken,
278 model_name: &str,
279 text: &str,
280 ) -> RuntimeResult<DocumentEmbeddingOutcome> {
281 self.embed_document_with_model_outcome_inner(Some(token), model_name, text)
282 .await
283 }
284
285 async fn embed_document_with_model_outcome_inner(
286 &self,
287 token: Option<&NamespaceToken>,
288 model_name: &str,
289 text: &str,
290 ) -> RuntimeResult<DocumentEmbeddingOutcome> {
291 let model = parse_embedding_model_alias(model_name);
292 let service = match token {
293 Some(token) => self.embedder_with_token(token, model_name).await?,
294 None => self.embedder(model_name).await?,
295 };
296 let emb_model = model.unwrap_or_default();
297 let source_bytes = text.len();
298 let (text, truncated) =
299 bounded_embedding_input(text, document_embedding_budget(model_name));
300 let embedded_bytes = text.len();
301 crate::usage::count(crate::usage::UsageUnit::EmbedCalls, 1);
303 let embeddings = service.embed_passage(&[text.to_string()], emb_model).await;
304 let mut vectors = embeddings?;
305 if vectors.len() != 1 {
306 return Err(RuntimeError::Internal(format!(
307 "embed_passage returned {} vectors for 1 input",
308 vectors.len()
309 )));
310 }
311 let out = vectors.pop().expect("checked len == 1 above");
312 Ok(DocumentEmbeddingOutcome {
313 model_name: model_name.to_owned(),
314 vector: out,
315 source_bytes,
316 embedded_bytes,
317 truncated,
318 })
319 }
320
321 pub async fn embed_query_with_model(
333 &self,
334 model_name: &str,
335 text: &str,
336 ) -> RuntimeResult<Vec<f32>> {
337 self.embed_query_with_model_inner(None, model_name, text)
338 .await
339 }
340
341 pub(crate) async fn embed_query_with_model_for_token(
342 &self,
343 token: &NamespaceToken,
344 model_name: &str,
345 text: &str,
346 ) -> RuntimeResult<Vec<f32>> {
347 self.embed_query_with_model_inner(Some(token), model_name, text)
348 .await
349 }
350
351 async fn embed_query_with_model_inner(
352 &self,
353 token: Option<&NamespaceToken>,
354 model_name: &str,
355 text: &str,
356 ) -> RuntimeResult<Vec<f32>> {
357 let model = parse_embedding_model_alias(model_name);
358 let service = match token {
359 Some(token) => self.embedder_with_token(token, model_name).await?,
360 None => self.embedder(model_name).await?,
361 };
362 let texts = [text.to_string()];
363 let emb_model = model.unwrap_or_default();
364 crate::usage::count(crate::usage::UsageUnit::EmbedCalls, 1);
366 let embeddings = match emb_model {
367 EmbeddingModel::BgeSmallEnV15
368 | EmbeddingModel::BgeBaseEnV15
369 | EmbeddingModel::BgeLargeEnV15 => service.embed(&texts, emb_model).await,
370 _ => service.embed_query(&texts, emb_model).await,
371 };
372 let out = embeddings?
373 .into_iter()
374 .next()
375 .ok_or_else(|| RuntimeError::Internal("embed_query returned empty vec".into()))?;
376 Ok(out)
377 }
378
379 pub async fn embed_document(&self, text: &str) -> RuntimeResult<Vec<f32>> {
386 let model_name = self.default_embedder_name();
387 if model_name.is_empty() {
388 return Err(RuntimeError::Unconfigured("embedding_model".into()));
389 }
390 self.embed_document_with_model(model_name, text).await
391 }
392
393 pub async fn embed_query(&self, text: &str) -> RuntimeResult<Vec<f32>> {
400 let model_name = self.default_embedder_name();
401 if model_name.is_empty() {
402 return Err(RuntimeError::Unconfigured("embedding_model".into()));
403 }
404 self.embed_query_with_model(model_name, text).await
405 }
406
407 async fn embed_query_for_token(
408 &self,
409 token: &NamespaceToken,
410 text: &str,
411 ) -> RuntimeResult<Vec<f32>> {
412 let model_name = self.default_embedder_name();
413 if model_name.is_empty() {
414 return Err(RuntimeError::Unconfigured("embedding_model".into()));
415 }
416 self.embed_query_with_model_for_token(token, model_name, text)
417 .await
418 }
419
420 pub async fn embed_batch(&self, texts: &[String]) -> RuntimeResult<Vec<Vec<f32>>> {
428 if texts.is_empty() {
429 return Ok(vec![]);
430 }
431 let model_name = self.default_embedder_name();
432 if model_name.is_empty() {
433 return Err(RuntimeError::Unconfigured("embedding_model".into()));
434 }
435 self.embed_batch_with_model(model_name, texts).await
436 }
437
438 pub async fn embed_batch_with_model(
443 &self,
444 model_name: &str,
445 texts: &[String],
446 ) -> RuntimeResult<Vec<Vec<f32>>> {
447 if texts.is_empty() {
448 return Ok(vec![]);
449 }
450 let model = parse_embedding_model_alias(model_name);
451 let service = self.embedder(model_name).await?;
452 let emb_model = model.unwrap_or_default();
453 let out = service.embed(texts, emb_model).await;
454 crate::usage::count(crate::usage::UsageUnit::EmbedCalls, texts.len() as u64);
455 Ok(out?)
456 }
457
458 pub async fn embed_document_batch_with_model(
470 &self,
471 model_name: &str,
472 texts: &[String],
473 ) -> RuntimeResult<Vec<Vec<f32>>> {
474 if texts.is_empty() {
475 return Ok(vec![]);
476 }
477 Ok(self
478 .embed_document_batch_with_model_outcomes(model_name, texts)
479 .await?
480 .into_iter()
481 .map(|outcome| outcome.vector)
482 .collect())
483 }
484
485 pub async fn embed_document_batch_with_model_outcomes(
486 &self,
487 model_name: &str,
488 texts: &[String],
489 ) -> RuntimeResult<Vec<DocumentEmbeddingOutcome>> {
490 self.embed_document_batch_with_model_outcomes_inner(None, model_name, texts)
491 .await
492 }
493
494 pub(crate) async fn embed_document_batch_with_model_outcomes_for_token(
495 &self,
496 token: &NamespaceToken,
497 model_name: &str,
498 texts: &[String],
499 ) -> RuntimeResult<Vec<DocumentEmbeddingOutcome>> {
500 self.embed_document_batch_with_model_outcomes_inner(Some(token), model_name, texts)
501 .await
502 }
503
504 async fn embed_document_batch_with_model_outcomes_inner(
505 &self,
506 token: Option<&NamespaceToken>,
507 model_name: &str,
508 texts: &[String],
509 ) -> RuntimeResult<Vec<DocumentEmbeddingOutcome>> {
510 if texts.is_empty() {
511 return Ok(vec![]);
512 }
513 let model = parse_embedding_model_alias(model_name);
514 let service = match token {
515 Some(token) => self.embedder_with_token(token, model_name).await?,
516 None => self.embedder(model_name).await?,
517 };
518 let emb_model = model.unwrap_or_default();
519 let budget = document_embedding_budget(model_name);
520 if token.is_some() {
521 crate::usage::count(crate::usage::UsageUnit::EmbedCalls, texts.len() as u64);
523 }
524 let out = if texts.iter().all(|text| text.len() <= budget) {
525 service.embed_passage(texts, emb_model).await
526 } else {
527 let bounded_texts: Vec<String> = texts
528 .iter()
529 .map(|text| bounded_embedding_input(text, budget).0.to_owned())
530 .collect();
531 service.embed_passage(&bounded_texts, emb_model).await
532 };
533 if token.is_none() {
534 crate::usage::count(crate::usage::UsageUnit::EmbedCalls, texts.len() as u64);
535 }
536 let vectors = out?;
537 if vectors.len() != texts.len() {
538 return Err(RuntimeError::Internal(format!(
539 "embed_passage returned {} vectors for {} inputs",
540 vectors.len(),
541 texts.len()
542 )));
543 }
544 Ok(texts
545 .iter()
546 .zip(vectors)
547 .map(|(text, vector)| {
548 let (bounded, truncated) = bounded_embedding_input(text, budget);
549 DocumentEmbeddingOutcome {
550 model_name: model_name.to_owned(),
551 vector,
552 source_bytes: text.len(),
553 embedded_bytes: bounded.len(),
554 truncated,
555 }
556 })
557 .collect())
558 }
559
560 pub async fn embed_document_batch(&self, texts: &[String]) -> RuntimeResult<Vec<Vec<f32>>> {
567 if texts.is_empty() {
568 return Ok(vec![]);
569 }
570 let model_name = self.default_embedder_name();
571 if model_name.is_empty() {
572 return Err(RuntimeError::Unconfigured("embedding_model".into()));
573 }
574 self.embed_document_batch_with_model(model_name, texts)
575 .await
576 }
577
578 pub async fn embed_document_batch_outcomes(
580 &self,
581 texts: &[String],
582 ) -> RuntimeResult<Vec<DocumentEmbeddingOutcome>> {
583 if texts.is_empty() {
584 return Ok(vec![]);
585 }
586 let model_name = self.default_embedder_name();
587 if model_name.is_empty() {
588 return Err(RuntimeError::Unconfigured("embedding_model".into()));
589 }
590 self.embed_document_batch_with_model_outcomes(model_name, texts)
591 .await
592 }
593
594 pub async fn embed_query_batch_with_model(
601 &self,
602 model_name: &str,
603 texts: &[String],
604 ) -> RuntimeResult<Vec<Vec<f32>>> {
605 if texts.is_empty() {
606 return Ok(vec![]);
607 }
608 let model = parse_embedding_model_alias(model_name);
609 let service = self.embedder(model_name).await?;
610 let emb_model = model.unwrap_or_default();
611 let out = match emb_model {
612 EmbeddingModel::BgeSmallEnV15
613 | EmbeddingModel::BgeBaseEnV15
614 | EmbeddingModel::BgeLargeEnV15 => service.embed(texts, emb_model).await,
615 _ => service.embed_query(texts, emb_model).await,
616 };
617 crate::usage::count(crate::usage::UsageUnit::EmbedCalls, texts.len() as u64);
618 Ok(out?)
619 }
620
621 pub async fn vector_search(
627 &self,
628 token: &NamespaceToken,
629 query_embedding: Option<Vec<f32>>,
630 query_text: Option<&str>,
631 top_k: u32,
632 kind: Option<SubstrateKind>,
633 ) -> RuntimeResult<Vec<VectorSearchHit>> {
634 let embedding = match query_embedding {
635 Some(vec) => vec,
636 None => {
637 let text = query_text.ok_or_else(|| {
638 RuntimeError::InvalidInput(
639 "vector search requires query_embedding or query_text".into(),
640 )
641 })?;
642 if text.trim().is_empty() {
643 return Err(RuntimeError::InvalidInput(
644 "query_text must not be empty".into(),
645 ));
646 }
647 self.embed_query_for_token(token, text).await?
648 }
649 };
650
651 let ns = token.namespace().as_str().to_owned();
652 let hits = self
653 .vectors(token)?
654 .search(VectorSearchRequest {
655 query_vectors: vec![embedding],
656 top_k,
657 namespace: Some(ns),
658 kind,
659 embedding_model: None,
660 filter: None,
661 backend_hints: None,
662 })
663 .await;
664 crate::usage::count(crate::usage::UsageUnit::VectorPasses, 1);
665 Ok(hits?)
666 }
667
668 #[allow(clippy::too_many_arguments)]
709 pub async fn hybrid_search(
710 &self,
711 token: &NamespaceToken,
712 query_text: &str,
713 query_vector: Option<Vec<f32>>,
714 limit: u32,
715 entity_kind: Option<&str>,
716 entity_type: Option<&str>,
717 tags_any: &[String],
718 properties_filter: Option<&serde_json::Value>,
719 ) -> RuntimeResult<Vec<SearchHit>> {
720 self.hybrid_search_with_text_mode(
721 token,
722 query_text,
723 query_vector,
724 limit,
725 entity_kind,
726 entity_type,
727 tags_any,
728 properties_filter,
729 TextQueryMode::Plain,
730 )
731 .await
732 }
733
734 #[allow(clippy::too_many_arguments)]
736 pub async fn hybrid_search_with_text_mode(
737 &self,
738 token: &NamespaceToken,
739 query_text: &str,
740 query_vector: Option<Vec<f32>>,
741 limit: u32,
742 entity_kind: Option<&str>,
743 entity_type: Option<&str>,
744 tags_any: &[String],
745 properties_filter: Option<&serde_json::Value>,
746 text_mode: TextQueryMode,
747 ) -> RuntimeResult<Vec<SearchHit>> {
748 let (hits, _vector_error) = self
749 .hybrid_search_inner(
750 token,
751 query_text,
752 query_vector,
753 limit,
754 entity_kind,
755 entity_type,
756 tags_any,
757 properties_filter,
758 text_mode,
759 None,
760 false,
761 )
762 .await?;
763 Ok(hits)
764 }
765
766 #[allow(clippy::too_many_arguments)]
771 pub(crate) async fn hybrid_search_with_vector_similarity_floor(
772 &self,
773 token: &NamespaceToken,
774 query_text: &str,
775 query_vector: Option<Vec<f32>>,
776 limit: u32,
777 entity_kind: Option<&str>,
778 entity_type: Option<&str>,
779 tags_any: &[String],
780 properties_filter: Option<&serde_json::Value>,
781 vector_similarity_floor: f64,
782 ) -> RuntimeResult<Vec<SearchHit>> {
783 let (hits, _vector_error) = self
784 .hybrid_search_inner(
785 token,
786 query_text,
787 query_vector,
788 limit,
789 entity_kind,
790 entity_type,
791 tags_any,
792 properties_filter,
793 TextQueryMode::Plain,
794 Some(vector_similarity_floor),
795 false,
796 )
797 .await?;
798 Ok(hits)
799 }
800
801 #[allow(clippy::too_many_arguments)]
811 pub async fn hybrid_search_outcome(
812 &self,
813 token: &NamespaceToken,
814 query_text: &str,
815 limit: u32,
816 entity_kind: Option<&str>,
817 entity_type: Option<&str>,
818 tags_any: &[String],
819 properties_filter: Option<&serde_json::Value>,
820 ) -> RuntimeResult<HybridSearchOutcome> {
821 self.hybrid_search_outcome_with_text_mode(
822 token,
823 query_text,
824 limit,
825 entity_kind,
826 entity_type,
827 tags_any,
828 properties_filter,
829 TextQueryMode::Plain,
830 )
831 .await
832 }
833
834 #[allow(clippy::too_many_arguments)]
836 pub async fn hybrid_search_outcome_with_text_mode(
837 &self,
838 token: &NamespaceToken,
839 query_text: &str,
840 limit: u32,
841 entity_kind: Option<&str>,
842 entity_type: Option<&str>,
843 tags_any: &[String],
844 properties_filter: Option<&serde_json::Value>,
845 text_mode: TextQueryMode,
846 ) -> RuntimeResult<HybridSearchOutcome> {
847 let (hits, vector_error) = self
848 .hybrid_search_inner(
849 token,
850 query_text,
851 None,
852 limit,
853 entity_kind,
854 entity_type,
855 tags_any,
856 properties_filter,
857 text_mode,
858 None,
859 true,
860 )
861 .await?;
862 Ok(HybridSearchOutcome { hits, vector_error })
863 }
864
865 #[allow(clippy::too_many_arguments)]
866 async fn hybrid_search_inner(
867 &self,
868 token: &NamespaceToken,
869 query_text: &str,
870 query_vector: Option<Vec<f32>>,
871 limit: u32,
872 entity_kind: Option<&str>,
873 entity_type: Option<&str>,
874 tags_any: &[String],
875 properties_filter: Option<&serde_json::Value>,
876 text_mode: TextQueryMode,
877 vector_similarity_floor: Option<f64>,
878 tolerate_vector_error: bool,
879 ) -> RuntimeResult<(Vec<SearchHit>, Option<String>)> {
880 let candidates = limit.saturating_mul(CANDIDATE_MULTIPLIER).max(limit);
881
882 let visible_ns: Vec<String> = token
883 .visible_namespaces()
884 .iter()
885 .map(|ns| ns.as_str().to_owned())
886 .collect();
887 let text_search_result = self
891 .text(token)?
892 .search(TextSearchRequest {
893 query: query_text.to_string(),
894 mode: text_mode,
895 filter: Some(TextFilter {
896 namespaces: visible_ns.clone(),
897 ..TextFilter::default()
898 }),
899 top_k: candidates,
900 snippet_chars: 200,
901 })
902 .await;
903 let text_hits = crate::error::fts_text_leg_or_err(
908 text_search_result.map_err(RuntimeError::from),
909 "hybrid_search",
910 query_text,
911 )?;
912
913 let mut vector_error: Option<String> = None;
914 let mut vector_hits = if query_vector.is_some() || self.config().embedding_model.is_some() {
915 match self
916 .vector_search(
917 token,
918 query_vector,
919 Some(query_text),
920 candidates,
921 Some(SubstrateKind::Entity),
922 )
923 .await
924 {
925 Ok(hits) => hits,
926 Err(e) if tolerate_vector_error => {
927 vector_error = Some(e.to_string());
928 Vec::new()
929 }
930 Err(e) => return Err(e),
931 }
932 } else {
933 Vec::new()
934 };
935 if let Some(cosine_floor) = vector_similarity_floor {
936 let score_floor = DeterministicScore::from_f64(cosine_floor);
939 vector_hits.retain(|hit| hit.score >= score_floor);
940 }
941
942 let fusion_limit = text_hits.len().saturating_add(vector_hits.len());
947 let mut fused = rrf_fuse(text_hits, vector_hits, fusion_limit, query_text);
948
949 if !fused.is_empty() {
952 let candidate_ids: Vec<Uuid> = fused.iter().map(|h| h.entity_id).collect();
953 let alive_page = self
954 .entities(token)?
955 .query_entities(
956 token.namespace().as_str(),
957 EntityFilter {
958 ids: candidate_ids,
959 kinds: entity_kind.map(|k| vec![k.to_string()]).unwrap_or_default(),
960 entity_types: entity_type.map(|t| vec![t.to_string()]).unwrap_or_default(),
961 namespaces: visible_ns,
962 tags_any: tags_any.to_vec(),
963 ..EntityFilter::default()
964 },
965 PageRequest {
966 offset: 0,
967 limit: u32::try_from(fused.len()).unwrap_or(u32::MAX),
968 },
969 )
970 .await?;
971 let mut entity_meta: HashMap<Uuid, (String, Option<String>)> = HashMap::new();
972 let mut alive: HashSet<Uuid> = HashSet::new();
973 for e in alive_page.items {
974 if let Some(pf) = properties_filter {
977 if !entity_props_match(e.properties.as_ref(), pf) {
978 continue;
979 }
980 }
981 alive.insert(e.id);
982 entity_meta.insert(e.id, (e.name, e.description));
983 }
984
985 fused.retain(|h| alive.contains(&h.entity_id));
986
987 for hit in &mut fused {
989 if let Some((name, description)) = entity_meta.get(&hit.entity_id) {
990 if hit.title.is_none() {
991 hit.title = Some(name.clone());
992 }
993 if hit.snippet.is_none() {
994 hit.snippet = description.clone();
995 }
996 }
997 }
998 }
999
1000 fused.truncate(limit as usize);
1001 Ok((fused, vector_error))
1002 }
1003
1004 pub async fn knn(
1010 &self,
1011 token: &NamespaceToken,
1012 query_vector: Vec<f32>,
1013 top_k: u32,
1014 ) -> RuntimeResult<Vec<VectorSearchHit>> {
1015 let ns = token.namespace().as_str().to_owned();
1016 Ok(self
1017 .vectors(token)?
1018 .search(VectorSearchRequest {
1019 query_vectors: vec![query_vector],
1020 top_k,
1021 namespace: Some(ns),
1022 kind: Some(SubstrateKind::Entity),
1023 embedding_model: None,
1024 filter: None,
1025 backend_hints: None,
1026 })
1027 .await?)
1028 }
1029
1030 pub async fn rerank(
1036 &self,
1037 token: &NamespaceToken,
1038 query_vector: &[f32],
1039 candidate_ids: &[Uuid],
1040 top_k: u32,
1041 ) -> RuntimeResult<Vec<VectorSearchHit>> {
1042 let candidate_set: HashSet<Uuid> = candidate_ids.iter().copied().collect();
1043 let ns = token.namespace().as_str().to_owned();
1044 let all_hits = self
1045 .vectors(token)?
1046 .search(VectorSearchRequest {
1047 query_vectors: vec![query_vector.to_vec()],
1048 top_k: candidate_ids.len() as u32,
1049 namespace: Some(ns),
1050 kind: Some(SubstrateKind::Entity),
1051 embedding_model: None,
1052 filter: None,
1053 backend_hints: None,
1054 })
1055 .await?;
1056 let mut hits: Vec<VectorSearchHit> = all_hits
1057 .into_iter()
1058 .filter(|h| candidate_set.contains(&h.subject_id))
1059 .collect();
1060 hits.sort_by_key(|hit| std::cmp::Reverse(hit.score));
1061 hits.truncate(top_k as usize);
1062 Ok(hits)
1063 }
1064
1065 pub async fn backfill_missing_embeddings(&self, token: &NamespaceToken) -> RuntimeResult<u64> {
1078 use khive_storage::types::{SqlRow, SqlStatement, SqlValue};
1079
1080 let model_names = self.registered_embedding_model_names();
1081 if model_names.is_empty() {
1082 tracing::debug!(
1083 "backfill_missing_embeddings: no embedding models registered, skipping"
1084 );
1085 return Ok(0);
1086 }
1087
1088 let ns = token.namespace().as_str().to_string();
1089 let mut total_backfilled = 0u64;
1090
1091 for model_name in &model_names {
1092 let mut model_truncation = EmbeddingTruncationReport::default();
1093 let vec_table = format!("vec_{}", sanitize_key(model_name));
1095
1096 const PAGE_SIZE: usize = 500;
1100 let mut entity_total = 0usize;
1101 loop {
1102 let entity_sql = SqlStatement {
1103 sql: format!(
1104 "SELECT id, name, description FROM entities \
1105 WHERE namespace = ?1 AND deleted_at IS NULL \
1106 AND id NOT IN (\
1107 SELECT subject_id FROM {vec_table} \
1108 WHERE namespace = ?1 AND embedding_model = ?2 \
1109 ) LIMIT {PAGE_SIZE}"
1110 ),
1111 params: vec![
1112 SqlValue::Text(ns.clone()),
1113 SqlValue::Text(model_name.clone()),
1114 ],
1115 label: Some("backfill_entities".into()),
1116 };
1117
1118 let entity_rows: Vec<SqlRow> = {
1119 let sql = self.sql();
1120 let reader_result = sql.reader().await;
1121 #[cfg(any(test, feature = "fault-injection"))]
1122 let reader_result = if BACKFILL_READER_FAIL.with(|c| c.get()) {
1123 BACKFILL_READER_FAIL.with(|c| c.set(false));
1124 Err(khive_storage::StorageError::Pool {
1125 operation: "reader".into(),
1126 message: "injected failure".into(),
1127 })
1128 } else {
1129 reader_result
1130 };
1131 let mut reader = reader_result.map_err(RuntimeError::Storage)?;
1132 reader
1133 .query_all(entity_sql)
1134 .await
1135 .map_err(RuntimeError::Storage)?
1136 };
1137
1138 let batch_len = entity_rows.len();
1139 entity_total += batch_len;
1140
1141 for row in &entity_rows {
1142 let id_str = row.columns.first().and_then(|c| {
1143 if let SqlValue::Text(s) = &c.value {
1144 Some(s.clone())
1145 } else {
1146 None
1147 }
1148 });
1149 let description = row.columns.get(2).and_then(|c| {
1150 if let SqlValue::Text(s) = &c.value {
1151 Some(s.clone())
1152 } else if let SqlValue::Null = &c.value {
1153 None
1154 } else {
1155 None
1156 }
1157 });
1158
1159 let (Some(id_str), Some(desc)) = (id_str, description) else {
1160 continue;
1161 };
1162 let Ok(id) = id_str.parse::<Uuid>() else {
1163 continue;
1164 };
1165 if desc.trim().is_empty() {
1166 continue;
1167 }
1168
1169 match self
1170 .embed_document_with_model_outcome_for_token(token, model_name, &desc)
1171 .await
1172 {
1173 Ok(outcome) => {
1174 model_truncation.observe(&outcome);
1175 if let Ok(vs) = self.vectors_for_model(token, model_name) {
1176 match vs
1177 .insert(
1178 id,
1179 SubstrateKind::Entity,
1180 &ns,
1181 "entity.description",
1182 vec![outcome.vector],
1183 )
1184 .await
1185 {
1186 Ok(()) => {
1187 total_backfilled += 1;
1188 }
1189 Err(e) => {
1190 tracing::warn!(
1191 id = %id, model = %model_name,
1192 error = %e,
1193 "backfill_missing_embeddings: entity vector insert failed"
1194 );
1195 }
1196 }
1197 }
1198 }
1199 Err(e) => {
1200 tracing::warn!(
1201 id = %id, model = %model_name,
1202 error = %e,
1203 "backfill_missing_embeddings: entity embed failed"
1204 );
1205 }
1206 }
1207 }
1208
1209 if batch_len < PAGE_SIZE {
1210 break;
1211 }
1212 }
1213
1214 let text_store = self.text_for_notes(token).ok();
1216 let note_store = self.notes(token).ok();
1217 let mut note_total = 0usize;
1218 loop {
1219 let note_sql = SqlStatement {
1222 sql: format!(
1223 "SELECT id FROM notes \
1224 WHERE namespace = ?1 AND deleted_at IS NULL \
1225 AND id NOT IN (\
1226 SELECT subject_id FROM {vec_table} \
1227 WHERE namespace = ?1 AND embedding_model = ?2 \
1228 ) LIMIT {PAGE_SIZE}"
1229 ),
1230 params: vec![
1231 SqlValue::Text(ns.clone()),
1232 SqlValue::Text(model_name.clone()),
1233 ],
1234 label: Some("backfill_notes".into()),
1235 };
1236
1237 let note_rows: Vec<SqlRow> = {
1238 let sql = self.sql();
1239 let reader_result = sql.reader().await;
1240 #[cfg(any(test, feature = "fault-injection"))]
1241 let reader_result = if BACKFILL_READER_FAIL.with(|c| c.get()) {
1242 BACKFILL_READER_FAIL.with(|c| c.set(false));
1243 Err(khive_storage::StorageError::Pool {
1244 operation: "reader".into(),
1245 message: "injected failure".into(),
1246 })
1247 } else {
1248 reader_result
1249 };
1250 let mut reader = reader_result.map_err(RuntimeError::Storage)?;
1251 reader
1252 .query_all(note_sql)
1253 .await
1254 .map_err(RuntimeError::Storage)?
1255 };
1256
1257 let batch_len = note_rows.len();
1258 note_total += batch_len;
1259
1260 for row in ¬e_rows {
1261 let id_str = row.columns.first().and_then(|c| {
1262 if let SqlValue::Text(s) = &c.value {
1263 Some(s.clone())
1264 } else {
1265 None
1266 }
1267 });
1268
1269 let Some(id_str) = id_str else {
1270 continue;
1271 };
1272 let Ok(id) = id_str.parse::<Uuid>() else {
1273 continue;
1274 };
1275
1276 let note = match ¬e_store {
1277 Some(store) => match store.get_note(id).await {
1278 Ok(Some(n)) => n,
1279 _ => continue,
1280 },
1281 None => continue,
1282 };
1283
1284 if note.content.trim().is_empty() {
1285 continue;
1286 }
1287
1288 if model_names.first().map(|n| n.as_str()) == Some(model_name.as_str()) {
1291 if let Some(ref ts) = text_store {
1292 if let Err(e) = ts.upsert_document(note_fts_document(¬e)).await {
1293 tracing::warn!(id = %id, error = %e,
1294 "backfill_missing_embeddings: note FTS upsert failed");
1295 }
1296 }
1297 }
1298
1299 let content = note.content.clone();
1300 match self
1301 .embed_document_with_model_outcome_for_token(token, model_name, &content)
1302 .await
1303 {
1304 Ok(outcome) => {
1305 model_truncation.observe(&outcome);
1306 if let Ok(vs) = self.vectors_for_model(token, model_name) {
1307 match vs
1308 .insert(
1309 id,
1310 SubstrateKind::Note,
1311 &ns,
1312 "note.content",
1313 vec![outcome.vector],
1314 )
1315 .await
1316 {
1317 Ok(()) => {
1318 total_backfilled += 1;
1319 }
1320 Err(e) => {
1321 tracing::warn!(
1322 id = %id, model = %model_name,
1323 error = %e,
1324 "backfill_missing_embeddings: note vector insert failed"
1325 );
1326 }
1327 }
1328 }
1329 }
1330 Err(e) => {
1331 tracing::warn!(
1332 id = %id, model = %model_name,
1333 error = %e,
1334 "backfill_missing_embeddings: note embed failed"
1335 );
1336 }
1337 }
1338 }
1339
1340 if batch_len < PAGE_SIZE {
1341 break;
1342 }
1343 }
1344
1345 tracing::info!(
1346 model = %model_name,
1347 namespace = %ns,
1348 entities = entity_total,
1349 notes = note_total,
1350 truncated = model_truncation.truncated,
1351 discarded_bytes = model_truncation.discarded_bytes,
1352 "backfill_missing_embeddings: model pass complete"
1353 );
1354 }
1355
1356 tracing::info!(
1357 namespace = %ns,
1358 total_backfilled = total_backfilled,
1359 "backfill_missing_embeddings: finished"
1360 );
1361
1362 Ok(total_backfilled)
1363 }
1364
1365 pub async fn sweep_orphan_vectors(
1381 &self,
1382 token: &NamespaceToken,
1383 max_delete_per_model: u32,
1384 dry_run: bool,
1385 ) -> RuntimeResult<u64> {
1386 use khive_storage::types::OrphanSweepConfig;
1387 use khive_storage::StorageError;
1388
1389 let model_names = self.registered_embedding_model_names();
1390 if model_names.is_empty() {
1391 tracing::debug!("sweep_orphan_vectors: no embedding models registered, skipping");
1392 return Ok(0);
1393 }
1394
1395 let ns = token.namespace().as_str().to_string();
1396 let mut total_deleted = 0u64;
1397
1398 for model_name in &model_names {
1399 let store = match self.vectors_for_model(token, model_name) {
1400 Ok(s) => s,
1401 Err(e) => {
1402 tracing::warn!(
1403 model = %model_name,
1404 error = %e,
1405 "sweep_orphan_vectors: failed to get vector store, skipping model"
1406 );
1407 continue;
1408 }
1409 };
1410
1411 let caps = store.capabilities();
1412 if !caps.supports_orphan_sweep {
1413 tracing::debug!(
1414 model = %model_name,
1415 "sweep_orphan_vectors: backend does not support orphan sweep, skipping"
1416 );
1417 continue;
1418 }
1419
1420 let config = OrphanSweepConfig {
1421 subject_id_allowlist: None,
1422 namespaces: vec![ns.clone()],
1423 substrate_kinds: vec![],
1424 max_delete: max_delete_per_model,
1425 dry_run,
1426 };
1427
1428 match store.orphan_sweep(&config).await {
1429 Ok(result) => {
1430 tracing::info!(
1431 model = %model_name,
1432 namespace = %ns,
1433 scanned = result.scanned,
1434 deleted = result.deleted,
1435 would_delete = result.would_delete,
1436 dry_run = dry_run,
1437 "sweep_orphan_vectors: sweep complete"
1438 );
1439 total_deleted += result.deleted;
1440 }
1441 Err(StorageError::Unsupported { .. }) => {
1442 tracing::debug!(
1443 model = %model_name,
1444 "sweep_orphan_vectors: backend returned Unsupported, skipping"
1445 );
1446 }
1447 Err(e) => {
1448 tracing::warn!(
1449 model = %model_name,
1450 error = %e,
1451 "sweep_orphan_vectors: sweep failed, continuing with other models"
1452 );
1453 }
1454 }
1455 }
1456
1457 tracing::info!(
1458 namespace = %ns,
1459 total_deleted = total_deleted,
1460 dry_run = dry_run,
1461 "sweep_orphan_vectors: finished"
1462 );
1463
1464 Ok(total_deleted)
1465 }
1466}
1467
1468fn entity_props_match(
1473 entity_props: Option<&serde_json::Value>,
1474 filter: &serde_json::Value,
1475) -> bool {
1476 let required = match filter.as_object() {
1477 Some(obj) if !obj.is_empty() => obj,
1478 _ => return true,
1479 };
1480 let actual = match entity_props.and_then(serde_json::Value::as_object) {
1481 Some(obj) => obj,
1482 None => return false,
1483 };
1484 required
1485 .iter()
1486 .all(|(k, v)| actual.get(k).is_some_and(|av| av == v))
1487}
1488
1489const EXACT_MATCH_BOOST: f64 = 0.5;
1493
1494fn rrf_fuse(
1502 text_hits: Vec<TextSearchHit>,
1503 vector_hits: Vec<VectorSearchHit>,
1504 limit: usize,
1505 query_text: &str,
1506) -> Vec<SearchHit> {
1507 #[derive(Default)]
1508 struct Bucket {
1509 score: DeterministicScore,
1510 signals: SearchSignals,
1511 source: Option<SearchSource>,
1512 title: Option<String>,
1513 snippet: Option<String>,
1514 }
1515
1516 let mut buckets: HashMap<Uuid, Bucket> = HashMap::new();
1517
1518 let query_lower = query_text.to_lowercase();
1519 let mut text_seen = HashSet::with_capacity(text_hits.len());
1520 for (i, hit) in text_hits.into_iter().enumerate() {
1521 if !text_seen.insert(hit.subject_id) {
1522 continue;
1523 }
1524 let rank = i + 1; let entry = buckets.entry(hit.subject_id).or_default();
1526 entry.score = entry.score + rrf_score(rank, RRF_K);
1527 entry.signals.keyword_score = Some(hit.score);
1528 entry.source = Some(match entry.source {
1529 Some(SearchSource::Vector) => SearchSource::Both,
1530 _ => SearchSource::Text,
1531 });
1532 if entry.title.is_none() {
1533 if let Some(ref title) = hit.title {
1535 if title.to_lowercase() == query_lower {
1536 entry.score = entry.score + DeterministicScore::from_f64(EXACT_MATCH_BOOST);
1537 }
1538 }
1539 entry.title = hit.title;
1540 }
1541 if entry.snippet.is_none() {
1542 entry.snippet = hit.snippet;
1543 }
1544 }
1545
1546 let mut vector_seen = HashSet::with_capacity(vector_hits.len());
1547 for (i, hit) in vector_hits.into_iter().enumerate() {
1548 if !vector_seen.insert(hit.subject_id) {
1549 continue;
1550 }
1551 let rank = i + 1;
1552 let entry = buckets.entry(hit.subject_id).or_default();
1553 entry.score = entry.score + rrf_score(rank, RRF_K);
1554 entry.signals.vector_similarity = Some(hit.score);
1555 entry.source = Some(match entry.source {
1556 Some(SearchSource::Text) => SearchSource::Both,
1557 _ => SearchSource::Vector,
1558 });
1559 }
1560
1561 let mut hits: Vec<SearchHit> = buckets
1562 .into_iter()
1563 .map(|(id, b)| SearchHit {
1564 entity_id: id,
1565 score: b.score,
1566 rank_score_kind: RankScoreKind::Rrf,
1567 signals: b.signals,
1568 source: b.source.expect("each bucket gets a source"),
1569 title: b.title,
1570 snippet: b.snippet,
1571 })
1572 .collect();
1573
1574 hits.sort_by(|a, b| b.score.cmp(&a.score).then(a.entity_id.cmp(&b.entity_id)));
1575 hits.truncate(limit);
1576 hits
1577}
1578
1579#[cfg(test)]
1580mod tests {
1581 use super::*;
1582 use std::sync::Arc;
1583
1584 use crate::runtime::{KhiveRuntime, NamespaceToken, RuntimeConfig};
1585 use khive_storage::types::{TextSearchHit, VectorSearchHit};
1586 use khive_types::namespace::Namespace;
1587 use lattice_embed::{EmbedError, EmbeddingModel};
1588
1589 struct FailingEmbeddingService;
1593
1594 #[async_trait::async_trait]
1595 impl EmbeddingService for FailingEmbeddingService {
1596 async fn embed(
1597 &self,
1598 _texts: &[String],
1599 _model: EmbeddingModel,
1600 ) -> Result<Vec<Vec<f32>>, EmbedError> {
1601 Err(EmbedError::ModelInitialization(
1602 "injected vector-arm failure".to_string(),
1603 ))
1604 }
1605
1606 fn supports_model(&self, _model: EmbeddingModel) -> bool {
1607 true
1608 }
1609
1610 fn name(&self) -> &'static str {
1611 "hybrid-search-test-failing-embedding"
1612 }
1613 }
1614
1615 struct FailingEmbedderProvider {
1616 name: String,
1617 dimensions: usize,
1618 }
1619
1620 #[async_trait::async_trait]
1621 impl EmbedderProvider for FailingEmbedderProvider {
1622 fn name(&self) -> &str {
1623 &self.name
1624 }
1625
1626 fn dimensions(&self) -> usize {
1627 self.dimensions
1628 }
1629
1630 async fn build(&self) -> RuntimeResult<Arc<dyn EmbeddingService>> {
1631 Ok(Arc::new(FailingEmbeddingService))
1632 }
1633 }
1634
1635 fn break_vector_arm(runtime: &KhiveRuntime) {
1642 let model = EmbeddingModel::AllMiniLmL6V2;
1643 runtime.register_embedder(FailingEmbedderProvider {
1644 name: model.to_string(),
1645 dimensions: model.dimensions(),
1646 });
1647 }
1648
1649 struct ConstantEmbeddingService {
1653 dimensions: usize,
1654 }
1655
1656 #[async_trait::async_trait]
1657 impl EmbeddingService for ConstantEmbeddingService {
1658 async fn embed(
1659 &self,
1660 texts: &[String],
1661 _model: EmbeddingModel,
1662 ) -> Result<Vec<Vec<f32>>, EmbedError> {
1663 Ok(texts.iter().map(|_| vec![1.0; self.dimensions]).collect())
1664 }
1665
1666 fn supports_model(&self, _model: EmbeddingModel) -> bool {
1667 true
1668 }
1669
1670 fn name(&self) -> &'static str {
1671 "hybrid-search-test-constant-embedding"
1672 }
1673 }
1674
1675 struct ConstantEmbedderProvider {
1676 name: String,
1677 dimensions: usize,
1678 }
1679
1680 #[async_trait::async_trait]
1681 impl EmbedderProvider for ConstantEmbedderProvider {
1682 fn name(&self) -> &str {
1683 &self.name
1684 }
1685
1686 fn dimensions(&self) -> usize {
1687 self.dimensions
1688 }
1689
1690 async fn build(&self) -> RuntimeResult<Arc<dyn EmbeddingService>> {
1691 Ok(Arc::new(ConstantEmbeddingService {
1692 dimensions: self.dimensions,
1693 }))
1694 }
1695 }
1696
1697 fn runtime_with_constant_embeddings() -> KhiveRuntime {
1700 let model = EmbeddingModel::AllMiniLmL6V2;
1701 let runtime = KhiveRuntime::new(RuntimeConfig {
1702 db_path: None,
1703 embedding_model: Some(model),
1704 packs: vec!["kg".to_string()],
1705 ..RuntimeConfig::no_embeddings()
1706 })
1707 .expect("in-memory runtime");
1708 runtime.register_embedder(ConstantEmbedderProvider {
1709 name: model.to_string(),
1710 dimensions: model.dimensions(),
1711 });
1712 runtime
1713 }
1714
1715 #[test]
1716 fn bounded_embedding_input_reserves_prefix_and_preserves_utf8() {
1717 assert_eq!(
1718 document_embedding_budget("multilingual-e5-base"),
1719 MAX_TEXT_BYTES - "passage: ".len()
1720 );
1721
1722 let input = format!("{}\u{1f980}tail", "a".repeat(MAX_TEXT_BYTES - 1));
1723 let (bounded, truncated) = bounded_embedding_input(&input, MAX_TEXT_BYTES);
1724 assert!(truncated);
1725 assert_eq!(bounded.len(), MAX_TEXT_BYTES - 1);
1726 assert!(bounded.is_char_boundary(bounded.len()));
1727 assert!(!bounded.contains('\u{1f980}'));
1728 }
1729
1730 #[test]
1731 fn bounded_embedding_input_leaves_normal_text_unchanged() {
1732 let input = "normal byte-identical embedding input";
1733 let (bounded, truncated) = bounded_embedding_input(input, MAX_TEXT_BYTES);
1734 assert!(!truncated);
1735 assert_eq!(bounded, input);
1736 assert_eq!(bounded.as_ptr(), input.as_ptr());
1737 }
1738
1739 fn text_hit(id: Uuid, rank: u32, title: &str) -> TextSearchHit {
1740 TextSearchHit {
1741 subject_id: id,
1742 score: DeterministicScore::from_f64(1.0),
1743 rank,
1744 title: Some(title.to_string()),
1745 snippet: Some("...".to_string()),
1746 }
1747 }
1748
1749 fn vector_hit(id: Uuid, rank: u32) -> VectorSearchHit {
1750 VectorSearchHit {
1751 subject_id: id,
1752 score: DeterministicScore::from_f64(0.9),
1753 rank,
1754 }
1755 }
1756
1757 #[test]
1758 fn rrf_evidence_keeps_absence_distinct_from_measured_zero() {
1759 let text_id = Uuid::from_u128(1);
1760 let vector_id = Uuid::from_u128(2);
1761 let both_id = Uuid::from_u128(3);
1762 let quarter = DeterministicScore::from_raw(1_i64 << 30);
1763 let half = DeterministicScore::from_raw(1_i64 << 31);
1764 let text = vec![
1765 TextSearchHit {
1766 score: DeterministicScore::ZERO,
1767 ..text_hit(text_id, 1, "text")
1768 },
1769 TextSearchHit {
1770 score: quarter,
1771 ..text_hit(both_id, 2, "both")
1772 },
1773 text_hit(both_id, 3, "duplicate"),
1774 ];
1775 let vector = vec![
1776 VectorSearchHit {
1777 score: DeterministicScore::ZERO,
1778 ..vector_hit(vector_id, 1)
1779 },
1780 VectorSearchHit {
1781 score: half,
1782 ..vector_hit(both_id, 2)
1783 },
1784 vector_hit(both_id, 3),
1785 ];
1786 let hits = rrf_fuse(text, vector, 10, "unmatched");
1787 assert_eq!(hits.len(), 3);
1788 for (id, source, signals) in [
1789 (
1790 text_id,
1791 SearchSource::Text,
1792 SearchSignals {
1793 vector_similarity: None,
1794 keyword_score: Some(DeterministicScore::ZERO),
1795 },
1796 ),
1797 (
1798 vector_id,
1799 SearchSource::Vector,
1800 SearchSignals {
1801 vector_similarity: Some(DeterministicScore::ZERO),
1802 keyword_score: None,
1803 },
1804 ),
1805 (
1806 both_id,
1807 SearchSource::Both,
1808 SearchSignals {
1809 vector_similarity: Some(half),
1810 keyword_score: Some(quarter),
1811 },
1812 ),
1813 ] {
1814 let hit = hits.iter().find(|hit| hit.entity_id == id).unwrap();
1815 assert_eq!(hit.rank_score_kind, RankScoreKind::Rrf);
1816 assert_eq!(hit.source, source);
1817 assert_eq!(hit.signals, signals);
1818 }
1819 }
1820
1821 #[test]
1822 fn rrf_evidence_golden_preserves_true_ties_across_permutations() {
1823 let a = Uuid::from_u128(1);
1824 let b = Uuid::from_u128(2);
1825 let expected = vec![
1826 (
1827 a,
1828 748_365_513,
1829 RankScoreKind::Rrf,
1830 SearchSignals {
1831 vector_similarity: Some(DeterministicScore::from_raw(1_i64 << 31)),
1832 keyword_score: Some(DeterministicScore::from_raw(1_i64 << 30)),
1833 },
1834 ),
1835 (
1836 b,
1837 748_365_513,
1838 RankScoreKind::Rrf,
1839 SearchSignals {
1840 vector_similarity: Some(DeterministicScore::from_raw(1_i64 << 32)),
1841 keyword_score: Some(DeterministicScore::from_raw(3_i64 << 30)),
1842 },
1843 ),
1844 ];
1845 for (text_ids, vector_ids) in [([a, b], [b, a]), ([b, a], [a, b])] {
1847 for _ in 0..4 {
1848 let text = text_ids
1849 .into_iter()
1850 .enumerate()
1851 .map(|(rank, id)| TextSearchHit {
1852 score: DeterministicScore::from_raw(
1853 if id == a { 1_i64 } else { 3_i64 } << 30,
1854 ),
1855 ..text_hit(id, rank as u32 + 1, "candidate")
1856 })
1857 .collect();
1858 let vector = vector_ids
1859 .into_iter()
1860 .enumerate()
1861 .map(|(rank, id)| VectorSearchHit {
1862 score: DeterministicScore::from_raw(
1863 if id == a { 1_i64 } else { 2_i64 } << 31,
1864 ),
1865 ..vector_hit(id, rank as u32 + 1)
1866 })
1867 .collect();
1868 let hits = rrf_fuse(text, vector, 10, "unmatched");
1869 assert_eq!(hits.len(), 2);
1870 assert_eq!(hits[0].score, hits[1].score);
1871 assert!(hits.iter().all(|hit| hit.source == SearchSource::Both));
1872 let snapshot: Vec<_> = hits
1873 .iter()
1874 .map(|hit| {
1875 (
1876 hit.entity_id,
1877 hit.score.to_raw(),
1878 hit.rank_score_kind,
1879 hit.signals,
1880 )
1881 })
1882 .collect();
1883 assert_eq!(snapshot, expected);
1884 }
1885 }
1886 }
1887
1888 #[test]
1889 fn rrf_fuse_text_only() {
1890 let a = Uuid::new_v4();
1891 let b = Uuid::new_v4();
1892 let text = vec![text_hit(a, 1, "A"), text_hit(b, 2, "B")];
1893 let hits = rrf_fuse(text, vec![], 10, "query");
1894 assert_eq!(hits.len(), 2);
1895 assert_eq!(hits[0].entity_id, a);
1896 assert_eq!(hits[0].source, SearchSource::Text);
1897 assert_eq!(hits[0].title.as_deref(), Some("A"));
1898 }
1899
1900 #[test]
1901 fn rrf_fuse_vector_only() {
1902 let a = Uuid::new_v4();
1903 let hits = rrf_fuse(vec![], vec![vector_hit(a, 1)], 10, "query");
1904 assert_eq!(hits.len(), 1);
1905 assert_eq!(hits[0].source, SearchSource::Vector);
1906 assert!(hits[0].title.is_none());
1907 }
1908
1909 #[test]
1910 fn rrf_fuse_marks_both_when_in_both_lists() {
1911 let id = Uuid::new_v4();
1912 let text = vec![text_hit(id, 1, "A")];
1913 let vec = vec![vector_hit(id, 1)];
1914 let hits = rrf_fuse(text, vec, 10, "query");
1915 assert_eq!(hits.len(), 1);
1916 assert_eq!(hits[0].source, SearchSource::Both);
1917 }
1918
1919 #[test]
1920 fn rrf_fuse_preserves_unique_leg_scores_exactly() {
1921 let text_only = Uuid::new_v4();
1922 let both = Uuid::new_v4();
1923 let vector_only = Uuid::new_v4();
1924 let text = vec![text_hit(text_only, 1, "A"), text_hit(both, 2, "B")];
1925 let vector = vec![vector_hit(both, 1), vector_hit(vector_only, 2)];
1926
1927 let hits = rrf_fuse(text, vector, 10, "query");
1928 let score_for = |id| {
1929 hits.iter()
1930 .find(|hit| hit.entity_id == id)
1931 .expect("expected fused hit")
1932 .score
1933 };
1934
1935 assert_eq!(score_for(text_only), rrf_score(1, RRF_K));
1936 assert_eq!(score_for(both), rrf_score(2, RRF_K) + rrf_score(1, RRF_K));
1937 assert_eq!(score_for(vector_only), rrf_score(2, RRF_K));
1938 }
1939
1940 #[test]
1941 fn rrf_fuse_counts_duplicate_once_per_leg() {
1942 let id = Uuid::new_v4();
1943 let text = vec![text_hit(id, 1, "A"), text_hit(id, 2, "A duplicate")];
1944 let vector = vec![vector_hit(id, 1), vector_hit(id, 2)];
1945
1946 let hits = rrf_fuse(text, vector, 10, "query");
1947
1948 assert_eq!(hits.len(), 1);
1949 assert_eq!(hits[0].source, SearchSource::Both);
1950 assert_eq!(hits[0].score, rrf_score(1, RRF_K) + rrf_score(1, RRF_K));
1951 }
1952
1953 #[test]
1954 fn rrf_fuse_respects_limit() {
1955 let hits: Vec<TextSearchHit> = (0..20)
1956 .map(|i| text_hit(Uuid::new_v4(), i + 1, "x"))
1957 .collect();
1958 let fused = rrf_fuse(hits, vec![], 5, "query");
1959 assert_eq!(fused.len(), 5);
1960 }
1961
1962 #[test]
1963 fn rrf_fuse_orders_higher_score_first() {
1964 let a = Uuid::new_v4();
1966 let b = Uuid::new_v4();
1967 let text = vec![text_hit(a, 1, "A")];
1968 let vec = vec![vector_hit(a, 1), vector_hit(b, 2)];
1969 let hits = rrf_fuse(text, vec, 10, "query");
1970 assert_eq!(hits[0].entity_id, a);
1971 assert_eq!(hits[0].source, SearchSource::Both);
1972 assert!(hits[0].score > hits[1].score);
1973 }
1974
1975 #[test]
1976 fn rrf_fuse_k10_score_spread_exceeds_threshold() {
1977 let ids: Vec<Uuid> = (0..10).map(|_| Uuid::new_v4()).collect();
1980 let text: Vec<TextSearchHit> = ids
1981 .iter()
1982 .enumerate()
1983 .map(|(i, &id)| text_hit(id, (i + 1) as u32, "x"))
1984 .collect();
1985 let hits = rrf_fuse(text, vec![], 10, "query");
1986 assert_eq!(hits.len(), 10);
1987 let top_score = hits[0].score.to_f64();
1988 let bottom_score = hits[9].score.to_f64();
1989 let spread = top_score - bottom_score;
1990 assert!(
1991 spread >= 0.03,
1992 "score spread {spread:.4} between rank 1 and rank 10 must be ≥ 0.03 (was {spread:.4})"
1993 );
1994 }
1995
1996 #[test]
1997 fn rrf_fuse_exact_match_boost_elevates_score() {
1998 let exact_id = Uuid::new_v4();
2001 let other_id = Uuid::new_v4();
2002 let text = vec![
2004 text_hit(other_id, 1, "something else"),
2005 text_hit(exact_id, 2, "FlashAttention"),
2006 ];
2007 let hits = rrf_fuse(text, vec![], 10, "flashattention");
2008 assert_eq!(hits.len(), 2);
2009 assert_eq!(
2010 hits[0].entity_id, exact_id,
2011 "exact match must rank first despite being rank-2 in raw text search"
2012 );
2013 }
2014
2015 #[test]
2018 fn embed_batch_unconfigured_on_memory_runtime() {
2019 let rt = KhiveRuntime::memory().unwrap();
2021 let result = tokio::runtime::Runtime::new()
2022 .unwrap()
2023 .block_on(rt.embed_batch(&[]));
2024 assert!(result.is_ok());
2026 assert!(result.unwrap().is_empty());
2027 }
2028
2029 #[test]
2030 fn embed_batch_empty_input_returns_empty_vec() {
2031 let rt = KhiveRuntime::memory().unwrap();
2033 let result = tokio::runtime::Runtime::new()
2034 .unwrap()
2035 .block_on(rt.embed_batch(&[]));
2036 assert_eq!(result.unwrap(), Vec::<Vec<f32>>::new());
2037 }
2038
2039 #[test]
2040 fn embed_batch_no_model_non_empty_returns_unconfigured() {
2041 let rt = KhiveRuntime::memory().unwrap();
2042 let texts = vec!["hello".to_string()];
2043 let result = tokio::runtime::Runtime::new()
2044 .unwrap()
2045 .block_on(rt.embed_batch(&texts));
2046 match result {
2047 Err(crate::RuntimeError::Unconfigured(s)) => assert_eq!(s, "embedding_model"),
2048 Err(other) => panic!("expected Unconfigured, got {:?}", other),
2049 Ok(_) => panic!("expected Err, got Ok"),
2050 }
2051 }
2052
2053 #[test]
2054 #[ignore = "loads ~80 MB model; run with --include-ignored"]
2055 fn embed_batch_count_matches_input() {
2056 let config = RuntimeConfig {
2057 db_path: None,
2058 default_namespace: Namespace::parse("test").unwrap(),
2059 embedding_model: Some(EmbeddingModel::AllMiniLmL6V2),
2060 packs: vec!["kg".to_string()],
2061 ..RuntimeConfig::default()
2062 };
2063 let rt = KhiveRuntime::new(config).unwrap();
2064 let texts: Vec<String> = vec!["foo".to_string(), "bar".to_string(), "baz".to_string()];
2065 let result = tokio::runtime::Runtime::new()
2066 .unwrap()
2067 .block_on(rt.embed_batch(&texts));
2068 let embeddings = result.unwrap();
2069 assert_eq!(embeddings.len(), texts.len());
2070 }
2071
2072 #[test]
2073 fn vector_search_requires_embedding_or_text() {
2074 let rt = KhiveRuntime::memory().unwrap();
2075 let tok = NamespaceToken::local();
2076 let result = tokio::runtime::Runtime::new()
2077 .unwrap()
2078 .block_on(rt.vector_search(&tok, None, None, 10, Some(SubstrateKind::Entity)));
2079 match result {
2080 Err(crate::RuntimeError::InvalidInput(msg)) => {
2081 assert!(msg.contains("query_embedding or query_text"), "msg: {msg}");
2082 }
2083 other => panic!("expected InvalidInput, got {other:?}"),
2084 }
2085 }
2086
2087 #[test]
2088 fn vector_search_text_without_model_returns_unconfigured() {
2089 let rt = KhiveRuntime::memory().unwrap();
2090 let tok = NamespaceToken::local();
2091 let result = tokio::runtime::Runtime::new()
2092 .unwrap()
2093 .block_on(rt.vector_search(
2094 &tok,
2095 None,
2096 Some("attention"),
2097 10,
2098 Some(SubstrateKind::Entity),
2099 ));
2100 match result {
2101 Err(crate::RuntimeError::Unconfigured(s)) => assert_eq!(s, "embedding_model"),
2102 other => panic!("expected Unconfigured, got {other:?}"),
2103 }
2104 }
2105
2106 #[test]
2107 #[ignore = "loads ~80 MB model; run with --include-ignored"]
2108 fn embed_batch_vectors_have_expected_dimensions() {
2109 let model = EmbeddingModel::AllMiniLmL6V2;
2110 let config = RuntimeConfig {
2111 db_path: None,
2112 default_namespace: Namespace::parse("test").unwrap(),
2113 embedding_model: Some(model),
2114 packs: vec!["kg".to_string()],
2115 ..RuntimeConfig::default()
2116 };
2117 let rt = KhiveRuntime::new(config).unwrap();
2118 let texts = vec!["hello world".to_string()];
2119 let result = tokio::runtime::Runtime::new()
2120 .unwrap()
2121 .block_on(rt.embed_batch(&texts));
2122 let embeddings = result.unwrap();
2123 assert_eq!(embeddings[0].len(), model.dimensions());
2124 }
2125
2126 #[tokio::test]
2134 async fn hybrid_search_still_fails_loud_on_vector_arm_error() {
2135 let rt = runtime_with_constant_embeddings();
2136 let tok = NamespaceToken::local();
2137 rt.create_entity(
2138 &tok,
2139 "concept",
2140 None,
2141 "FlashAttention",
2142 Some("IO-aware exact attention using tiling"),
2143 None,
2144 vec![],
2145 )
2146 .await
2147 .unwrap();
2148 break_vector_arm(&rt);
2149
2150 let result = rt
2151 .hybrid_search(&tok, "FlashAttention", None, 10, None, None, &[], None)
2152 .await;
2153
2154 assert!(
2155 result.is_err(),
2156 "the fail-loud entry point must still propagate a vector-arm failure, got {result:?}"
2157 );
2158 }
2159
2160 #[tokio::test]
2164 async fn hybrid_search_outcome_preserves_text_hits_on_vector_arm_error() {
2165 let rt = runtime_with_constant_embeddings();
2166 let tok = NamespaceToken::local();
2167 rt.create_entity(
2168 &tok,
2169 "concept",
2170 None,
2171 "FlashAttention",
2172 Some("IO-aware exact attention using tiling"),
2173 None,
2174 vec![],
2175 )
2176 .await
2177 .unwrap();
2178 break_vector_arm(&rt);
2179
2180 let outcome = rt
2181 .hybrid_search_outcome(&tok, "FlashAttention", 10, None, None, &[], None)
2182 .await
2183 .expect("text leg must still succeed");
2184
2185 assert!(
2186 !outcome.hits.is_empty(),
2187 "text arm's hit must survive a vector-arm failure"
2188 );
2189 assert!(
2190 outcome.hits[0]
2191 .title
2192 .as_deref()
2193 .unwrap_or_default()
2194 .contains("FlashAttention"),
2195 "surviving hit must be the text match"
2196 );
2197 let vector_error = outcome
2198 .vector_error
2199 .expect("vector arm failure must be reported");
2200 assert!(
2201 vector_error.contains("injected vector-arm failure"),
2202 "vector_error must carry the underlying cause, got {vector_error:?}"
2203 );
2204 }
2205
2206 #[tokio::test]
2209 async fn hybrid_search_outcome_has_no_vector_error_when_vector_arm_healthy() {
2210 let rt = runtime_with_constant_embeddings();
2211 let tok = NamespaceToken::local();
2212 rt.create_entity(
2213 &tok,
2214 "concept",
2215 None,
2216 "FlashAttention",
2217 Some("IO-aware exact attention using tiling"),
2218 None,
2219 vec![],
2220 )
2221 .await
2222 .unwrap();
2223
2224 let outcome = rt
2225 .hybrid_search_outcome(&tok, "FlashAttention", 10, None, None, &[], None)
2226 .await
2227 .expect("hybrid search must succeed");
2228
2229 assert!(!outcome.hits.is_empty(), "should find the entity");
2230 assert!(
2231 outcome.vector_error.is_none(),
2232 "a healthy vector arm must not report an error"
2233 );
2234 }
2235
2236 #[tokio::test]
2239 async fn hybrid_search_entity_hit_has_title() {
2240 let rt = KhiveRuntime::memory().unwrap();
2241 let tok = NamespaceToken::local();
2242 rt.create_entity(
2243 &tok,
2244 "concept",
2245 None,
2246 "FlashAttention",
2247 Some("IO-aware exact attention using tiling"),
2248 None,
2249 vec![],
2250 )
2251 .await
2252 .unwrap();
2253
2254 let hits = rt
2255 .hybrid_search(&tok, "FlashAttention", None, 10, None, None, &[], None)
2256 .await
2257 .unwrap();
2258
2259 assert!(!hits.is_empty(), "should find the entity");
2260 let hit = &hits[0];
2261 assert!(hit.title.is_some(), "title must be populated");
2262 assert!(
2263 hit.title.as_deref().unwrap().contains("FlashAttention"),
2264 "title must contain entity name"
2265 );
2266 }
2267
2268 #[tokio::test]
2274 async fn hybrid_search_with_dollar_sign_query_does_not_error() {
2275 let rt = KhiveRuntime::memory().unwrap();
2276 let tok = NamespaceToken::local();
2277 rt.create_entity(
2278 &tok,
2279 "concept",
2280 None,
2281 "DSL docs",
2282 Some("use $prev.id to chain calls"),
2283 None,
2284 vec![],
2285 )
2286 .await
2287 .unwrap();
2288
2289 let result = rt
2290 .hybrid_search(&tok, "$prev.id", None, 10, None, None, &[], None)
2291 .await;
2292
2293 assert!(
2294 result.is_ok(),
2295 "#388 hybrid_search must not hard-fail on a '$'-bearing query, got: {:?}",
2296 result.err()
2297 );
2298 }
2299
2300 #[tokio::test]
2309 async fn hybrid_search_with_residual_fts5_char_now_sanitized() {
2310 let rt = KhiveRuntime::memory().unwrap();
2311 let tok = NamespaceToken::local();
2312 rt.create_entity(
2313 &tok,
2314 "concept",
2315 None,
2316 "DSL docs",
2317 Some("use foo@bar to chain calls"),
2318 None,
2319 vec![],
2320 )
2321 .await
2322 .unwrap();
2323
2324 let result = rt
2325 .hybrid_search(&tok, "foo@bar", None, 10, None, None, &[], None)
2326 .await;
2327
2328 let hits = result.unwrap_or_else(|e| {
2329 panic!("#916 hybrid_search must not fail on an '@'-bearing query, got: {e:?}")
2330 });
2331 assert!(
2332 !hits.is_empty(),
2333 "#916 '@'-bearing query must still find the seeded 'foo@bar' content via the \
2334 quoted-phrase alternative"
2335 );
2336 }
2337
2338 #[tokio::test]
2346 async fn hybrid_search_with_916_issue_characters_finds_text_leg_hits() {
2347 let rt = KhiveRuntime::memory().unwrap();
2348 let tok = NamespaceToken::local();
2349
2350 rt.create_entity(
2351 &tok,
2352 "concept",
2353 None,
2354 "issue tracker",
2355 Some("tracking #682 Stage 2: MoE expert-cache prefetch work"),
2356 None,
2357 vec![],
2358 )
2359 .await
2360 .unwrap();
2361 rt.create_entity(
2362 &tok,
2363 "concept",
2364 None,
2365 "benchmark notes",
2366 Some("chunkwise B=128 traffic arithmetic simdgroup_matrix DPLR"),
2367 None,
2368 vec![],
2369 )
2370 .await
2371 .unwrap();
2372 rt.create_entity(
2373 &tok,
2374 "concept",
2375 None,
2376 "sampling notes",
2377 Some("evaluated with the Min-K%Prob membership inference method"),
2378 None,
2379 vec![],
2380 )
2381 .await
2382 .unwrap();
2383
2384 for query in ["#682 Stage 2", "B=128", "Min-K%Prob"] {
2385 let result = rt
2386 .hybrid_search(&tok, query, None, 10, None, None, &[], None)
2387 .await;
2388 let hits = result.unwrap_or_else(|e| {
2389 panic!("#916 hybrid_search must not fail on query {query:?}, got: {e:?}")
2390 });
2391 assert!(
2392 hits.iter()
2393 .any(|h| matches!(h.source, SearchSource::Text | SearchSource::Both)),
2394 "#916 query {query:?} must surface a Text/Both-sourced hit \
2395 (the FTS leg must contribute, not just the vector leg); got {hits:?}"
2396 );
2397 }
2398 }
2399
2400 #[tokio::test]
2414 async fn hybrid_search_tag_filter_pushed_before_truncation() {
2415 let rt = KhiveRuntime::memory().unwrap();
2416 let tok = NamespaceToken::local();
2417
2418 rt.create_entity(
2420 &tok,
2421 "concept",
2422 None,
2423 "alpha beta gamma decoy alpha beta gamma",
2424 Some("alpha beta gamma decoy description alpha beta gamma"),
2425 None,
2426 vec!["other-tag".to_string()],
2427 )
2428 .await
2429 .unwrap();
2430
2431 let target = rt
2433 .create_entity(
2434 &tok,
2435 "concept",
2436 None,
2437 "alpha beta gamma target",
2438 Some("alpha beta gamma target description"),
2439 None,
2440 vec!["target-tag".to_string()],
2441 )
2442 .await
2443 .unwrap();
2444
2445 let hits = rt
2449 .hybrid_search(
2450 &tok,
2451 "alpha beta gamma",
2452 None,
2453 1,
2454 None,
2455 None,
2456 &["target-tag".to_string()],
2457 None,
2458 )
2459 .await
2460 .unwrap();
2461
2462 assert_eq!(
2463 hits.len(),
2464 1,
2465 "exactly one hit expected (the tag-matching entity)"
2466 );
2467 assert_eq!(
2468 hits[0].entity_id, target.id,
2469 "the tag-filtered entity must be returned even when ranked below limit in raw fusion"
2470 );
2471 }
2472
2473 #[tokio::test]
2482 async fn hybrid_search_props_filter_pushed_before_truncation() {
2483 let rt = KhiveRuntime::memory().unwrap();
2484 let tok = NamespaceToken::local();
2485
2486 rt.create_entity(
2487 &tok,
2488 "concept",
2489 None,
2490 "delta epsilon zeta decoy delta epsilon zeta",
2491 Some("delta epsilon zeta decoy description delta epsilon zeta"),
2492 Some(serde_json::json!({"domain": "other"})),
2493 vec![],
2494 )
2495 .await
2496 .unwrap();
2497
2498 let target = rt
2499 .create_entity(
2500 &tok,
2501 "concept",
2502 None,
2503 "delta epsilon zeta target",
2504 Some("delta epsilon zeta target description"),
2505 Some(serde_json::json!({"domain": "target"})),
2506 vec![],
2507 )
2508 .await
2509 .unwrap();
2510
2511 let filter = serde_json::json!({"domain": "target"});
2512 let hits = rt
2513 .hybrid_search(
2514 &tok,
2515 "delta epsilon zeta",
2516 None,
2517 1,
2518 None,
2519 None,
2520 &[],
2521 Some(&filter),
2522 )
2523 .await
2524 .unwrap();
2525
2526 assert_eq!(hits.len(), 1, "exactly one hit expected (properties match)");
2527 assert_eq!(
2528 hits[0].entity_id, target.id,
2529 "the properties-filtered entity must be returned even when ranked below limit"
2530 );
2531 }
2532
2533 struct CapturingEmbeddingService {
2536 captured: std::sync::Arc<std::sync::Mutex<Vec<Vec<String>>>>,
2537 }
2538
2539 #[async_trait::async_trait]
2540 impl EmbeddingService for CapturingEmbeddingService {
2541 async fn embed(
2542 &self,
2543 texts: &[String],
2544 _model: EmbeddingModel,
2545 ) -> std::result::Result<Vec<Vec<f32>>, lattice_embed::EmbedError> {
2546 self.captured.lock().unwrap().push(texts.to_vec());
2547 Ok(texts.iter().map(|_| vec![1.0]).collect())
2548 }
2549
2550 fn supports_model(&self, _model: EmbeddingModel) -> bool {
2551 true
2552 }
2553
2554 fn name(&self) -> &'static str {
2555 "capturing-embedding-service"
2556 }
2557 }
2558
2559 struct CapturingEmbedderProvider {
2560 name: String,
2561 captured: std::sync::Arc<std::sync::Mutex<Vec<Vec<String>>>>,
2562 }
2563
2564 struct WrongCardinalityEmbeddingService;
2565
2566 #[async_trait::async_trait]
2567 impl EmbeddingService for WrongCardinalityEmbeddingService {
2568 async fn embed(
2569 &self,
2570 texts: &[String],
2571 _model: EmbeddingModel,
2572 ) -> std::result::Result<Vec<Vec<f32>>, lattice_embed::EmbedError> {
2573 Ok(texts
2574 .iter()
2575 .take(texts.len().saturating_sub(1))
2576 .map(|_| vec![1.0])
2577 .collect())
2578 }
2579
2580 fn supports_model(&self, _model: EmbeddingModel) -> bool {
2581 true
2582 }
2583
2584 fn name(&self) -> &'static str {
2585 "wrong-cardinality-embedding-service"
2586 }
2587 }
2588
2589 struct WrongCardinalityEmbedderProvider;
2590
2591 #[async_trait::async_trait]
2592 impl EmbedderProvider for WrongCardinalityEmbedderProvider {
2593 fn name(&self) -> &str {
2594 "wrong-cardinality-embedding-service"
2595 }
2596
2597 fn dimensions(&self) -> usize {
2598 1
2599 }
2600
2601 async fn build(&self) -> crate::error::RuntimeResult<std::sync::Arc<dyn EmbeddingService>> {
2602 Ok(std::sync::Arc::new(WrongCardinalityEmbeddingService))
2603 }
2604 }
2605
2606 struct SurplusCardinalityEmbeddingService;
2607
2608 #[async_trait::async_trait]
2609 impl EmbeddingService for SurplusCardinalityEmbeddingService {
2610 async fn embed(
2611 &self,
2612 texts: &[String],
2613 _model: EmbeddingModel,
2614 ) -> std::result::Result<Vec<Vec<f32>>, lattice_embed::EmbedError> {
2615 let mut vectors: Vec<Vec<f32>> = texts.iter().map(|_| vec![1.0]).collect();
2616 vectors.push(vec![1.0]);
2617 Ok(vectors)
2618 }
2619
2620 fn supports_model(&self, _model: EmbeddingModel) -> bool {
2621 true
2622 }
2623
2624 fn name(&self) -> &'static str {
2625 "surplus-cardinality-embedding-service"
2626 }
2627 }
2628
2629 struct SurplusCardinalityEmbedderProvider;
2630
2631 #[async_trait::async_trait]
2632 impl EmbedderProvider for SurplusCardinalityEmbedderProvider {
2633 fn name(&self) -> &str {
2634 "surplus-cardinality-embedding-service"
2635 }
2636
2637 fn dimensions(&self) -> usize {
2638 1
2639 }
2640
2641 async fn build(&self) -> crate::error::RuntimeResult<std::sync::Arc<dyn EmbeddingService>> {
2642 Ok(std::sync::Arc::new(SurplusCardinalityEmbeddingService))
2643 }
2644 }
2645
2646 #[async_trait::async_trait]
2647 impl EmbedderProvider for CapturingEmbedderProvider {
2648 fn name(&self) -> &str {
2649 &self.name
2650 }
2651
2652 fn dimensions(&self) -> usize {
2653 1
2654 }
2655
2656 async fn build(&self) -> crate::error::RuntimeResult<std::sync::Arc<dyn EmbeddingService>> {
2657 Ok(std::sync::Arc::new(CapturingEmbeddingService {
2658 captured: std::sync::Arc::clone(&self.captured),
2659 }))
2660 }
2661 }
2662
2663 fn runtime_with_capturing_embedder(
2664 model: EmbeddingModel,
2665 ) -> (
2666 KhiveRuntime,
2667 std::sync::Arc<std::sync::Mutex<Vec<Vec<String>>>>,
2668 ) {
2669 let runtime = KhiveRuntime::memory().unwrap();
2670 let captured = std::sync::Arc::new(std::sync::Mutex::new(Vec::new()));
2671 runtime.register_embedder(CapturingEmbedderProvider {
2672 name: model.to_string(),
2673 captured: std::sync::Arc::clone(&captured),
2674 });
2675 (runtime, captured)
2676 }
2677
2678 #[tokio::test]
2679 async fn bge_query_paths_pass_raw_unprefixed_text() {
2680 const BGE_QUERY_INSTRUCTION: &str =
2681 "Represent this sentence for searching relevant passages: ";
2682 let single = "single raw query";
2683 let batch = vec![
2684 "first raw query".to_string(),
2685 "second raw query".to_string(),
2686 ];
2687
2688 for model in [
2689 EmbeddingModel::BgeSmallEnV15,
2690 EmbeddingModel::BgeBaseEnV15,
2691 EmbeddingModel::BgeLargeEnV15,
2692 ] {
2693 let (runtime, captured) = runtime_with_capturing_embedder(model);
2694 runtime
2695 .embed_query_with_model(&model.to_string(), single)
2696 .await
2697 .unwrap();
2698 runtime
2699 .embed_query_batch_with_model(&model.to_string(), &batch)
2700 .await
2701 .unwrap();
2702
2703 let calls = captured.lock().unwrap().clone();
2704 assert_eq!(
2705 calls,
2706 vec![vec![single.to_string()], batch.clone()],
2707 "{model} must receive raw query text through single and batch paths"
2708 );
2709 assert!(
2710 calls
2711 .iter()
2712 .flatten()
2713 .all(|text| !text.contains(BGE_QUERY_INSTRUCTION)),
2714 "{model} must not receive the BGE retrieval instruction"
2715 );
2716 }
2717 }
2718
2719 #[tokio::test]
2720 async fn e5_query_paths_apply_query_prefix() {
2721 let model = EmbeddingModel::MultilingualE5Small;
2722 let single = "single raw query";
2723 let batch = vec![
2724 "first raw query".to_string(),
2725 "second raw query".to_string(),
2726 ];
2727 let (runtime, captured) = runtime_with_capturing_embedder(model);
2728
2729 runtime
2730 .embed_query_with_model(&model.to_string(), single)
2731 .await
2732 .unwrap();
2733 runtime
2734 .embed_query_batch_with_model(&model.to_string(), &batch)
2735 .await
2736 .unwrap();
2737
2738 assert_eq!(
2739 captured.lock().unwrap().as_slice(),
2740 [
2741 vec!["query: single raw query".to_string()],
2742 vec![
2743 "query: first raw query".to_string(),
2744 "query: second raw query".to_string(),
2745 ],
2746 ],
2747 "E5 must receive its query prefix through single and batch paths"
2748 );
2749 }
2750
2751 #[tokio::test]
2752 async fn mixed_document_batch_stays_one_ordered_provider_call() {
2753 let model = EmbeddingModel::AllMiniLmL6V2;
2754 let (runtime, captured) = runtime_with_capturing_embedder(model);
2755 let texts = vec![
2756 "first normal document".to_string(),
2757 "x".repeat(MAX_TEXT_BYTES + 1),
2758 "second normal document".to_string(),
2759 ];
2760
2761 let outcomes = runtime
2762 .embed_document_batch_with_model_outcomes(&model.to_string(), &texts)
2763 .await
2764 .expect("mixed batch must embed");
2765 assert_eq!(outcomes.len(), texts.len());
2766 assert!(!outcomes[0].truncated);
2767 assert_eq!(outcomes[0].source_bytes, outcomes[0].embedded_bytes);
2768 assert!(outcomes[1].truncated);
2769 assert_eq!(outcomes[1].source_bytes, MAX_TEXT_BYTES + 1);
2770 assert_eq!(outcomes[1].embedded_bytes, MAX_TEXT_BYTES);
2771 assert!(!outcomes[2].truncated);
2772 assert_eq!(
2773 captured.lock().unwrap().as_slice(),
2774 [vec![
2775 texts[0].clone(),
2776 "x".repeat(MAX_TEXT_BYTES),
2777 texts[2].clone(),
2778 ]]
2779 );
2780 }
2781
2782 #[tokio::test]
2783 async fn truncated_document_batch_rejects_provider_cardinality_mismatch() {
2784 let runtime = KhiveRuntime::memory().unwrap();
2785 runtime.register_embedder(WrongCardinalityEmbedderProvider);
2786 let model_name = "wrong-cardinality-embedding-service";
2787 let texts = vec!["x".repeat(MAX_TEXT_BYTES + 1), "normal".to_string()];
2788
2789 let error = runtime
2790 .embed_document_batch_with_model_outcomes(model_name, &texts)
2791 .await
2792 .expect_err("provider cardinality mismatch must fail the whole batch");
2793
2794 assert!(
2795 error
2796 .to_string()
2797 .contains("embed_passage returned 1 vectors for 2 inputs"),
2798 "unexpected error: {error}"
2799 );
2800 }
2801
2802 #[tokio::test]
2803 async fn singleton_document_embed_rejects_zero_vectors() {
2804 let runtime = KhiveRuntime::memory().unwrap();
2805 runtime.register_embedder(WrongCardinalityEmbedderProvider);
2806 let model_name = "wrong-cardinality-embedding-service";
2807
2808 let error = runtime
2809 .embed_document_with_model_outcome(model_name, "single document")
2810 .await
2811 .expect_err("provider returning zero vectors must fail closed");
2812
2813 assert!(
2814 error
2815 .to_string()
2816 .contains("embed_passage returned 0 vectors for 1 input"),
2817 "unexpected error: {error}"
2818 );
2819 }
2820
2821 #[tokio::test]
2822 async fn singleton_document_embed_rejects_surplus_vectors() {
2823 let runtime = KhiveRuntime::memory().unwrap();
2824 runtime.register_embedder(SurplusCardinalityEmbedderProvider);
2825 let model_name = "surplus-cardinality-embedding-service";
2826
2827 let error = runtime
2828 .embed_document_with_model_outcome(model_name, "single document")
2829 .await
2830 .expect_err("provider returning surplus vectors must fail closed");
2831
2832 assert!(
2833 error
2834 .to_string()
2835 .contains("embed_passage returned 2 vectors for 1 input"),
2836 "unexpected error: {error}"
2837 );
2838 }
2839
2840 #[test]
2841 #[ignore = "loads ~80 MB model; run with --include-ignored"]
2842 fn minilm_document_and_query_embed_are_identical_no_prefix_model() {
2843 let model = EmbeddingModel::AllMiniLmL6V2;
2847 let config = RuntimeConfig {
2848 db_path: None,
2849 default_namespace: Namespace::parse("test").unwrap(),
2850 embedding_model: Some(model),
2851 packs: vec!["kg".to_string()],
2852 ..RuntimeConfig::default()
2853 };
2854 let rt = KhiveRuntime::new(config).unwrap();
2855 let text = "attention is all you need".to_string();
2856 let rt_ref = &rt;
2857 let (doc_emb, query_emb) = tokio::runtime::Runtime::new().unwrap().block_on(async {
2858 let d = rt_ref
2859 .embed_document_with_model(&model.to_string(), &text)
2860 .await
2861 .unwrap();
2862 let q = rt_ref
2863 .embed_query_with_model(&model.to_string(), &text)
2864 .await
2865 .unwrap();
2866 (d, q)
2867 });
2868 assert_eq!(
2869 doc_emb, query_emb,
2870 "MiniLM has no instruction prefix: document and query embeds must be identical"
2871 );
2872 }
2873
2874 #[test]
2875 #[ignore = "loads multilingual-e5-small (~90 MB); run with --include-ignored"]
2876 fn e5_document_and_query_embed_differ_instruction_tuned_model() {
2877 let model = EmbeddingModel::MultilingualE5Small;
2882 let config = RuntimeConfig {
2883 db_path: None,
2884 default_namespace: Namespace::parse("test").unwrap(),
2885 embedding_model: Some(model),
2886 packs: vec!["kg".to_string()],
2887 ..RuntimeConfig::default()
2888 };
2889 let rt = KhiveRuntime::new(config).unwrap();
2890 let text = "attention is all you need".to_string();
2891 let rt_ref = &rt;
2892 let (doc_emb, query_emb) = tokio::runtime::Runtime::new().unwrap().block_on(async {
2893 let d = rt_ref
2894 .embed_document_with_model(&model.to_string(), &text)
2895 .await
2896 .unwrap();
2897 let q = rt_ref
2898 .embed_query_with_model(&model.to_string(), &text)
2899 .await
2900 .unwrap();
2901 (d, q)
2902 });
2903 assert_ne!(
2904 doc_emb, query_emb,
2905 "multilingual-e5-small uses asymmetric prefixes: document ('passage: ') \
2906 and query ('query: ') embeds of the same text must differ"
2907 );
2908 }
2909
2910 use crate::embedder_registry::EmbedderProvider;
2913 use lattice_embed::EmbeddingService;
2914
2915 struct StubEmbedderProvider;
2920
2921 #[async_trait::async_trait]
2922 impl EmbedderProvider for StubEmbedderProvider {
2923 fn name(&self) -> &str {
2924 "stub-model-m07"
2925 }
2926
2927 fn dimensions(&self) -> usize {
2928 4
2929 }
2930
2931 async fn build(&self) -> crate::error::RuntimeResult<std::sync::Arc<dyn EmbeddingService>> {
2932 struct StubSvc;
2933 #[async_trait::async_trait]
2934 impl EmbeddingService for StubSvc {
2935 async fn embed(
2936 &self,
2937 _texts: &[String],
2938 _model: lattice_embed::EmbeddingModel,
2939 ) -> std::result::Result<Vec<Vec<f32>>, lattice_embed::EmbedError> {
2940 Ok(vec![])
2941 }
2942
2943 fn supports_model(&self, _model: lattice_embed::EmbeddingModel) -> bool {
2944 true
2945 }
2946
2947 fn name(&self) -> &'static str {
2948 "stub-svc-m07"
2949 }
2950 }
2951 Ok(std::sync::Arc::new(StubSvc))
2952 }
2953 }
2954
2955 #[tokio::test]
2963 async fn backfill_reader_error_is_propagated_not_swallowed() {
2964 let rt = KhiveRuntime::memory().unwrap();
2965 rt.register_embedder(StubEmbedderProvider);
2966 let tok = NamespaceToken::local();
2967
2968 super::arm_backfill_reader_fail();
2971
2972 let result = rt.backfill_missing_embeddings(&tok).await;
2973 assert!(
2974 result.is_err(),
2975 "backfill_missing_embeddings must propagate the reader error (got Ok instead)"
2976 );
2977 let err_msg = result.unwrap_err().to_string();
2978 assert!(
2979 err_msg.contains("injected failure"),
2980 "error must originate from the injected reader failure, got: {err_msg}"
2981 );
2982 }
2983}