diff --git a/crates/ruvector-core/src/agenticdb.rs b/crates/ruvector-core/src/agenticdb.rs index 9a17a4a94c..11360f06b5 100644 --- a/crates/ruvector-core/src/agenticdb.rs +++ b/crates/ruvector-core/src/agenticdb.rs @@ -315,7 +315,7 @@ impl AgenticDB { k: usize, ) -> Result> { // Generate embedding for query - let query_embedding = self.generate_text_embedding(query)?; + let query_embedding = self.generate_query_embedding(query)?; // Search in vector DB let results = self.vector_db.search(SearchQuery { @@ -406,7 +406,7 @@ impl AgenticDB { /// Search skills by description pub fn search_skills(&self, query_description: &str, k: usize) -> Result> { - let query_embedding = self.generate_text_embedding(query_description)?; + let query_embedding = self.generate_query_embedding(query_description)?; let results = self.vector_db.search(SearchQuery { vector: query_embedding, @@ -527,7 +527,7 @@ impl AgenticDB { gamma: f64, ) -> Result> { let start_time = std::time::Instant::now(); - let query_embedding = self.generate_text_embedding(query)?; + let query_embedding = self.generate_query_embedding(query)?; // Get all causal edges let results = self.vector_db.search(SearchQuery { @@ -759,6 +759,16 @@ impl AgenticDB { fn generate_text_embedding(&self, text: &str) -> Result> { self.embedding_provider.embed(text) } + + /// Embed text that is a **search query** rather than stored content. + /// + /// Asymmetric providers apply a query-side instruction here that they must + /// not apply when embedding the documents being searched. Providers without + /// one fall back to [`EmbeddingProvider::embed`], so this is identical to + /// [`Self::generate_text_embedding`] for them. + fn generate_query_embedding(&self, text: &str) -> Result> { + self.embedding_provider.embed_query(text) + } } // Helper functions @@ -1044,7 +1054,7 @@ impl<'a> SessionStateIndex<'a> { /// Find relevant past turns based on current context pub fn find_relevant_turns(&self, query: &str, k: usize) -> Result> { - let query_embedding = self.db.generate_text_embedding(query)?; + let query_embedding = self.db.generate_query_embedding(query)?; let current_time = chrono::Utc::now().timestamp(); let results = self.db.vector_db.search(SearchQuery { @@ -1246,7 +1256,7 @@ impl<'a> WitnessLog<'a> { /// Search witness log semantically pub fn search(&self, query: &str, k: usize) -> Result> { - let query_embedding = self.db.generate_text_embedding(query)?; + let query_embedding = self.db.generate_query_embedding(query)?; let results = self.db.vector_db.search(SearchQuery { vector: query_embedding, @@ -1355,6 +1365,8 @@ impl AgenticDB { #[cfg(test)] mod tests { use super::*; + use crate::embeddings::EmbeddingProvider; + use parking_lot::Mutex; use tempfile::tempdir; fn create_test_db() -> Result { @@ -1494,4 +1506,106 @@ mod tests { Ok(()) } + + /// Records which side of the embedding provider each call landed on. + /// + /// Both sides return the same vector, so nothing downstream can tell them + /// apart by value. Only the recorded call log distinguishes them, which is + /// what makes the test below fail if a query path is routed back through + /// the passage method. + struct SideRecordingProvider { + dimensions: usize, + calls: Arc>>, + } + + impl EmbeddingProvider for SideRecordingProvider { + fn embed(&self, _text: &str) -> Result> { + self.calls.lock().push("passage"); + Ok(vec![0.1; self.dimensions]) + } + + fn embed_query(&self, _text: &str) -> Result> { + self.calls.lock().push("query"); + Ok(vec![0.1; self.dimensions]) + } + + fn dimensions(&self) -> usize { + self.dimensions + } + + fn name(&self) -> &str { + "side-recording" + } + } + + #[test] + fn searches_embed_their_query_on_the_query_side() -> Result<()> { + let calls: Arc>> = Arc::new(Mutex::new(Vec::new())); + let dir = tempdir().unwrap(); + let mut options = DbOptions::default(); + options.storage_path = dir.path().join("sides.db").to_string_lossy().to_string(); + options.dimensions = 128; + let db = AgenticDB::with_embedding_provider( + options, + Arc::new(SideRecordingProvider { + dimensions: 128, + calls: Arc::clone(&calls), + }), + )?; + + let drain = || -> Vec<&'static str> { std::mem::take(&mut *calls.lock()) }; + + // Stored content is a passage: it must never carry a query instruction. + db.store_episode( + "task".to_string(), + vec!["action".to_string()], + vec!["outcome".to_string()], + "critique".to_string(), + )?; + assert_eq!(drain(), ["passage"], "store_episode stores a passage"); + + db.session_index("s1", 3600).add_turn(1, "user", "hello")?; + assert_eq!(drain(), ["passage"], "add_turn stores a passage"); + + db.witness_log().append("agent", "act", "details")?; + assert_eq!(drain(), ["passage"], "witness append stores a passage"); + + db.create_skill( + "skill".to_string(), + "description".to_string(), + HashMap::new(), + vec!["example".to_string()], + )?; + assert_eq!(drain(), ["passage"], "create_skill stores a passage"); + + db.add_causal_edge( + vec!["cause".to_string()], + vec!["effect".to_string()], + 0.9, + "context".to_string(), + )?; + assert_eq!(drain(), ["passage"], "add_causal_edge stores a passage"); + + // Every search embeds its query text on the query side. + db.retrieve_similar_episodes("q", 5)?; + assert_eq!( + drain(), + ["query"], + "retrieve_similar_episodes embeds a query" + ); + + db.search_skills("q", 5)?; + assert_eq!(drain(), ["query"], "search_skills embeds a query"); + + db.query_with_utility("q", 5, 1.0, 0.0, 0.0)?; + assert_eq!(drain(), ["query"], "query_with_utility embeds a query"); + + db.session_index("s1", 3600).find_relevant_turns("q", 5)?; + assert_eq!(drain(), ["query"], "find_relevant_turns embeds a query"); + + db.witness_log().search("q", 5)?; + assert_eq!(drain(), ["query"], "witness search embeds a query"); + + Ok(()) + } } diff --git a/crates/ruvector-core/src/embeddings.rs b/crates/ruvector-core/src/embeddings.rs index 1f6cfd4cf6..2e81f686e5 100644 --- a/crates/ruvector-core/src/embeddings.rs +++ b/crates/ruvector-core/src/embeddings.rs @@ -47,6 +47,24 @@ pub trait EmbeddingProvider: Send + Sync { /// Generate embedding vector for the given text fn embed(&self, text: &str) -> Result>; + /// Generate an embedding for text that is a **search query**. + /// + /// Asymmetric embedding models encode queries and passages differently: + /// the query side carries an instruction prefix that the passage side must + /// not have. `bge-small-en-v1.5` prefixes queries with `"Represent this + /// sentence for searching relevant passages: "`, the E5 family uses + /// `"query: "` against `"passage: "`. Embedding a query through + /// [`embed`](EmbeddingProvider::embed) on such a model lowers + /// query-to-passage similarity: nothing errors, retrieval just gets worse. + /// + /// The default forwards to [`embed`](EmbeddingProvider::embed), which is + /// what a symmetric model wants, so existing providers keep their current + /// behaviour without changes. Providers backed by an asymmetric model + /// should override it. + fn embed_query(&self, text: &str) -> Result> { + self.embed(text) + } + /// Get the dimensionality of embeddings produced by this provider fn dimensions(&self) -> usize; @@ -877,6 +895,12 @@ pub mod lattice_native { // `request_tx` closes the channel, which ends the worker's `recv` // loop and lets the thread exit on its own. _worker: thread::JoinHandle<()>, + // Test-only observation seam: records which side (`"query"` / + // `"passage"`) each `send_request` call was for, so a test can prove + // that dispatch through `Arc::embed_query` + // reaches the query side instead of only comparing output vectors. + #[cfg(test)] + requested_sides: Arc>>, } impl LatticeEmbedding { @@ -971,6 +995,8 @@ pub mod lattice_native { dimensions: model.dimensions(), request_tx: Mutex::new(request_tx), _worker: worker, + #[cfg(test)] + requested_sides: Arc::new(Mutex::new(Vec::new())), }) } @@ -1001,6 +1027,15 @@ pub mod lattice_native { /// reply. Never calls `block_on` on the caller's thread, so this is /// safe to invoke from inside an existing async runtime. fn send_request(&self, kind: EmbedKind, text: &str) -> Result> { + #[cfg(test)] + { + let side = match kind { + EmbedKind::Query => "query", + EmbedKind::Passage => "passage", + }; + self.requested_sides.lock().unwrap().push(side); + } + let (reply_tx, reply_rx) = mpsc::channel(); let request = EmbedRequest { kind, @@ -1048,6 +1083,17 @@ pub mod lattice_native { self.send_request(EmbedKind::Passage, text) } + /// Embed **query** text, applying the model's query instruction when it + /// has one. + /// + /// This is what carries the asymmetry across the trait boundary. The + /// inherent [`LatticeEmbedding::embed_query`] is unreachable through an + /// `Arc`, so without this override every holder + /// of a boxed provider embeds queries as passages. + fn embed_query(&self, text: &str) -> Result> { + LatticeEmbedding::embed_query(self, text) + } + fn dimensions(&self) -> usize { self.dimensions } @@ -1217,6 +1263,36 @@ pub mod lattice_native { ); } } + + /// Regression test for the trait-object bridge: a caller holding only + /// `Arc` (no concrete `LatticeEmbedding` type) + /// must still reach the query side when it calls `embed_query`. The + /// `SideRecordingProvider`-style tests elsewhere in the crate use a + /// fake provider, so they'd keep passing even if `LatticeEmbedding`'s + /// own `embed_query` override were deleted -- they never touch this + /// impl. This test uses the real provider and the `requested_sides` + /// observation seam so it fails if that override goes away and + /// `embed_query` silently falls back to the trait default (which + /// forwards to `embed`, the passage side). + #[test] + fn dyn_provider_embed_query_reaches_query_side() { + let lattice = LatticeEmbedding::from_pretrained("bge-small-en-v1.5") + .expect("bge-small-en-v1.5 is a native local model"); + let requested_sides = Arc::clone(&lattice.requested_sides); + + let provider: Arc = Arc::new(lattice); + provider + .embed_query("a trait-object dispatch regression test") + .expect("embed_query must succeed through the trait object"); + + assert_eq!( + *requested_sides.lock().unwrap(), + ["query"], + "Arc::embed_query must dispatch through \ + LatticeEmbedding's query-side override; if this fails, the override was \ + removed and calls are falling back to the passage side" + ); + } } } @@ -1258,6 +1334,21 @@ mod tests { ); } + #[test] + fn embed_query_defaults_to_embed_for_symmetric_providers() { + // A provider that implements only `embed` must be unaffected by the + // addition of `embed_query`: the default forwards rather than leaving + // a hole. This is what makes the new trait method non-breaking for + // every existing implementor, in this crate and downstream. + let provider = HashEmbedding::new(128); + + assert_eq!( + provider.embed("the cat sat on the mat").unwrap(), + provider.embed_query("the cat sat on the mat").unwrap(), + "the default embed_query must return exactly what embed returns" + ); + } + #[cfg(feature = "real-embeddings")] #[test] #[ignore] // Requires model download