Skip to main content

lc_rag/
self_query.rs

1// lc-rag/src/self_query.rs
2//! SelfQueryRetriever — a self-querying retriever
3//!
4//! Lets an LLM split a natural-language query into `{ query, filter }`: the cleaned query
5//! goes to vector retrieval, and the parsed [`MetadataFilter`] is handed to
6//! `vector_store.similarity_search_with_filter` for metadata filtering (relying on S3's
7//! unified filtering capability). The split goes through [`lc_core::judge::structured_call`]
8//! (binds a tool to get structured arguments, falling back to text parsing when the model
9//! does not support it), the same execution path as Guardrails / Evaluation; the
10//! `allowed_attributes` whitelist blocks the LLM from filtering on fields that do not exist.
11//!
12//! **No silent degradation**: when a filter references a field outside the whitelist, it
13//! explicitly returns [`RetrieverError::InvalidFilter`] rather than dropping the filter and
14//! falling back to an unfiltered search — that would return data that should have been
15//! filtered out (data-plane over-exposure). An empty whitelist = filtering is entirely
16//! disabled: filters are always ignored with a warning (this is the established "disable
17//! filtering" pattern, not silent degradation).
18
19use std::sync::Arc;
20
21use async_trait::async_trait;
22use lc_core::judge::{structured_call, StructuredJudgeError};
23use lc_core::language_models::BaseChatModel;
24use lc_core::tools::ToolDefinition;
25use lc_embeddings::Embeddings;
26use lc_schema::Message;
27use lc_vector_stores::{Document, MetadataFilter, SearchResult, VectorStore};
28use serde::Deserialize;
29
30use crate::retriever::{RetrieverError, RetrieverTrait};
31
32/// Structured parameters parsed by the LLM: the cleaned query + an optional metadata filter.
33///
34/// `filter` is deserialized directly as [`MetadataFilter`] (filter.rs is lenient about
35/// the JSON shape to tolerate LLM output variance); the default is no filter.
36#[derive(Debug, Clone, PartialEq, Deserialize)]
37pub struct SelfQueryArgs {
38    /// The pure semantic query with the filtering constraints stripped out.
39    pub query: String,
40    /// Metadata filter (optional).
41    #[serde(default)]
42    pub filter: Option<MetadataFilter>,
43}
44
45/// Self-querying retriever: LLM splits the query -> whitelist validation -> filtered similarity search.
46///
47/// Implements [`RetrieverTrait`] and can be wrapped into an LCEL chain by [`crate::RetrieverRunnable`].
48pub struct SelfQueryRetriever<M: BaseChatModel> {
49    llm: Arc<M>,
50    store: Arc<dyn VectorStore>,
51    embeddings: Arc<dyn Embeddings>,
52    allowed_attributes: Vec<String>,
53}
54
55impl<M: BaseChatModel> SelfQueryRetriever<M> {
56    /// Creates a self-querying retriever.
57    ///
58    /// - `llm`: the model responsible for splitting the natural-language query (supports
59    ///   structured output, or falls back to text).
60    /// - `store` / `embeddings`: the vector store and embedding model used for the filtered
61    ///   similarity search.
62    /// - `allowed_attributes`: the whitelist of attributes allowed in filter fields;
63    ///   **an empty whitelist disables filtering entirely** (LLM-returned filters are always
64    ///   ignored with a warning); **with a non-empty whitelist, a filter referencing a field
65    ///   outside the whitelist explicitly returns [`RetrieverError::InvalidFilter`]** rather
66    ///   than silently dropping it and falling back to an unfiltered search.
67    pub fn new(
68        llm: impl Into<Arc<M>>,
69        store: Arc<dyn VectorStore>,
70        embeddings: Arc<dyn Embeddings>,
71        allowed_attributes: Vec<String>,
72    ) -> Self {
73        Self {
74            llm: llm.into(),
75            store,
76            embeddings,
77            allowed_attributes,
78        }
79    }
80
81    /// Self-query tool definition: lets the LLM return `{ query, filter }` structurally.
82    fn self_query_tool() -> ToolDefinition {
83        ToolDefinition::new(
84            "self_query",
85            "把用户的自然语言查询拆成纯语义查询词和可选的元数据过滤条件。",
86        )
87        .with_parameters(serde_json::json!({
88            "type": "object",
89            "properties": {
90                "query": {
91                    "type": "string",
92                    "description": "清洗掉过滤约束后的纯语义查询词"
93                },
94                "filter": {
95                    "type": ["object", "null"],
96                    "description": "元数据过滤条件(MetadataFilter JSON):单条件 {\"Field\": {\"key\", \"op\", \"value\"}},组合 {\"And\": [...]} / {\"Or\": [...]};op 取 Eq Ne Gt Gte Lt Lte In Nin(In/Nin 的 value 为数组);无过滤时为 null"
97                }
98            },
99            "required": ["query"]
100        }))
101    }
102
103    /// Builds the prompt: tells the LLM the available fields and the output format.
104    fn build_prompt(&self, query: &str) -> String {
105        let allowed = if self.allowed_attributes.is_empty() {
106            "无(本检索器不启用元数据过滤,filter 必须为 null)".to_string()
107        } else {
108            self.allowed_attributes.join(", ")
109        };
110        format!(
111            "把下面的自然语言查询拆成两部分:纯语义查询词(query)和可选的元数据过滤条件(filter)。\n\
112             filter 的 key 只能取以下允许字段之一: {allowed}\n\
113             filter 的 JSON 形状:单条件 {{\"Field\": {{\"key\": ..., \"op\": ..., \"value\": ...}}}};\n\
114             组合条件 {{\"And\": [...]}} / {{\"Or\": [...]}};op 取 Eq Ne Gt Gte Lt Lte In Nin(In/Nin 的 value 为数组)。\n\
115             没有过滤需求时 filter 为 null。\n\
116             用户查询: {query}"
117        )
118    }
119
120    /// Calls the LLM to split the query (structured or text fallback).
121    async fn parse_query(&self, query: &str) -> Result<SelfQueryArgs, RetrieverError> {
122        let messages = vec![Message::human(self.build_prompt(query))];
123        structured_call(
124            &*self.llm,
125            Self::self_query_tool(),
126            messages,
127            parse_text_fallback,
128        )
129        .await
130        .map_err(|e| RetrieverError::LlmError(e.to_string()))
131    }
132
133    /// Whitelist validation.
134    ///
135    /// - No filter -> `Ok(None)`.
136    /// - Empty whitelist = filtering is entirely disabled: filters are always ignored (this
137    ///   is the established "disable filtering" pattern, not silent degradation), logging a
138    ///   warning and returning `Ok(None)`.
139    /// - With a non-empty whitelist, a filter referencing a field outside the whitelist ->
140    ///   `Err(InvalidFilter)`. It never drops the filter to fall back to an unfiltered search,
141    ///   otherwise data that should have been filtered out would be returned (data-plane
142    ///   over-exposure).
143    fn validated_filter(
144        &self,
145        filter: &Option<MetadataFilter>,
146    ) -> Result<Option<MetadataFilter>, RetrieverError> {
147        let Some(f) = filter else {
148            return Ok(None);
149        };
150        if self.allowed_attributes.is_empty() {
151            log::warn!(
152                "SelfQuery: filtering is disabled (empty allowed_attributes); ignoring filter"
153            );
154            return Ok(None);
155        }
156        if let Some(field) = Self::disallowed_field(f, &self.allowed_attributes) {
157            return Err(RetrieverError::InvalidFilter(format!(
158                "filter references `{field}`, which is not in allowed_attributes [{}]; \
159                 refusing to degrade to an unfiltered search",
160                self.allowed_attributes.join(", ")
161            )));
162        }
163        Ok(filter.clone())
164    }
165
166    /// The first key referencing a field outside the whitelist, traversing nested And/Or
167    /// combinations; returns `None` when all fields are listed.
168    fn disallowed_field<'a>(filter: &'a MetadataFilter, allowed: &[String]) -> Option<&'a String> {
169        match filter {
170            MetadataFilter::Field { key, .. } => {
171                if allowed.iter().any(|a| a == key) {
172                    None
173                } else {
174                    Some(key)
175                }
176            }
177            MetadataFilter::And(items) | MetadataFilter::Or(items) => items
178                .iter()
179                .find_map(|f| Self::disallowed_field(f, allowed)),
180        }
181    }
182}
183
184/// Text fallback parsing: tries the whole JSON first (when the LLM outputs structured);
185/// otherwise treats the whole text as the query.
186fn parse_text_fallback(raw: &str) -> Result<SelfQueryArgs, StructuredJudgeError> {
187    let trimmed = raw.trim();
188    if !trimmed.is_empty() {
189        if let Ok(args) = serde_json::from_str::<SelfQueryArgs>(trimmed) {
190            return Ok(args);
191        }
192    }
193    let query = raw.trim().to_string();
194    if query.is_empty() {
195        return Err(StructuredJudgeError::Parse(
196            "self-query fallback produced an empty query".to_string(),
197        ));
198    }
199    Ok(SelfQueryArgs {
200        query,
201        filter: None,
202    })
203}
204
205#[async_trait]
206impl<M: BaseChatModel> RetrieverTrait for SelfQueryRetriever<M> {
207    async fn retrieve(&self, query: &str, k: usize) -> Result<Vec<Document>, RetrieverError> {
208        let results = self.retrieve_with_scores(query, k).await?;
209        Ok(results.into_iter().map(|r| r.document).collect())
210    }
211
212    async fn retrieve_with_scores(
213        &self,
214        query: &str,
215        k: usize,
216    ) -> Result<Vec<SearchResult>, RetrieverError> {
217        let args = self.parse_query(query).await?;
218        let filter = self.validated_filter(&args.filter)?;
219
220        let query_embedding = self
221            .embeddings
222            .embed_query(&args.query)
223            .await
224            .map_err(|e| RetrieverError::EmbeddingError(e.to_string()))?;
225
226        self.store
227            .similarity_search_with_filter(&query_embedding, k, filter.as_ref())
228            .await
229            .map_err(RetrieverError::from)
230    }
231
232    async fn add_documents(&self, documents: Vec<Document>) -> Result<(), RetrieverError> {
233        let texts: Vec<&str> = documents.iter().map(|d| d.content.as_str()).collect();
234        let embeddings = self
235            .embeddings
236            .embed_documents(&texts)
237            .await
238            .map_err(|e| RetrieverError::EmbeddingError(e.to_string()))?;
239        self.store.add_documents(documents, embeddings).await?;
240        Ok(())
241    }
242}
243
244#[cfg(test)]
245mod tests {
246    use super::*;
247    use async_trait::async_trait;
248    use futures_util::Stream;
249    use lc_core::language_models::{LLMResult, StreamChunk};
250    use lc_core::runnables::RunnableConfig;
251    use lc_core::{BaseLanguageModel, Runnable};
252    use lc_embeddings::MockEmbeddings;
253    use lc_vector_stores::InMemoryVectorStore;
254    use std::collections::HashSet;
255    use std::pin::Pin;
256    use std::sync::atomic::{AtomicUsize, Ordering};
257    use std::sync::Arc;
258
259    /// A mock chat model with a fixed reply (exercises the text fallback path; does not implement bind_tools).
260    struct MockChatModel {
261        reply: String,
262        calls: AtomicUsize,
263    }
264
265    impl MockChatModel {
266        fn new(reply: &str) -> Self {
267            Self {
268                reply: reply.to_string(),
269                calls: AtomicUsize::new(0),
270            }
271        }
272    }
273
274    #[async_trait]
275    impl Runnable<Vec<Message>, LLMResult> for MockChatModel {
276        type Error = MockChatError;
277        async fn invoke(
278            &self,
279            _input: Vec<Message>,
280            _config: Option<RunnableConfig>,
281        ) -> Result<LLMResult, Self::Error> {
282            Err(MockChatError)
283        }
284    }
285
286    #[async_trait]
287    impl BaseLanguageModel<Vec<Message>, LLMResult> for MockChatModel {
288        fn model_name(&self) -> &str {
289            "self-query-mock"
290        }
291        fn get_num_tokens(&self, t: &str) -> usize {
292            t.len()
293        }
294        fn with_temperature(self, _: f32) -> Self {
295            self
296        }
297        fn with_max_tokens(self, _: usize) -> Self {
298            self
299        }
300    }
301
302    #[derive(Debug)]
303    struct MockChatError;
304    impl std::fmt::Display for MockChatError {
305        fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
306            write!(f, "mock chat error")
307        }
308    }
309    impl std::error::Error for MockChatError {}
310
311    #[async_trait]
312    impl BaseChatModel for MockChatModel {
313        async fn chat(
314            &self,
315            _messages: Vec<Message>,
316            _config: Option<RunnableConfig>,
317        ) -> Result<LLMResult, Self::Error> {
318            self.calls.fetch_add(1, Ordering::SeqCst);
319            Ok(LLMResult {
320                content: self.reply.clone(),
321                model: "self-query-mock".to_string(),
322                token_usage: None,
323                tool_calls: None,
324                thinking_content: None,
325            })
326        }
327        async fn stream_chat(
328            &self,
329            _messages: Vec<Message>,
330            _config: Option<RunnableConfig>,
331        ) -> Result<Pin<Box<dyn Stream<Item = Result<StreamChunk, Self::Error>> + Send>>, Self::Error>
332        {
333            Err(MockChatError)
334        }
335    }
336
337    /// Builds an in-memory vector store with metadata + a mock embedding, returning the store for assertions.
338    async fn store_with_docs() -> Arc<InMemoryVectorStore> {
339        let store = Arc::new(InMemoryVectorStore::new());
340        let embeddings = Arc::new(MockEmbeddings::new(64));
341        let docs = vec![
342            Document::new("Rust systems programming").with_metadata("source", "docs"),
343            Document::new("Rust borrow checker").with_metadata("source", "docs"),
344            Document::new("Python scripting").with_metadata("source", "blog"),
345        ];
346        let texts: Vec<&str> = docs.iter().map(|d| d.content.as_str()).collect();
347        let vecs = embeddings.embed_documents(&texts).await.unwrap();
348        store.add_documents(docs, vecs).await.unwrap();
349        store
350    }
351
352    fn build_retriever(
353        llm: MockChatModel,
354        store: Arc<dyn VectorStore>,
355        allowed: &[&str],
356    ) -> SelfQueryRetriever<MockChatModel> {
357        SelfQueryRetriever::new(
358            Arc::new(llm),
359            store,
360            Arc::new(MockEmbeddings::new(64)),
361            allowed.iter().map(|s| s.to_string()).collect(),
362        )
363    }
364
365    /// S4: the parsed filter correctly reaches the search — only docs matching source=docs are returned.
366    #[tokio::test]
367    async fn test_self_query_filter_reaches_search() {
368        let store = store_with_docs().await;
369        let llm = MockChatModel::new(
370            r#"{"query": "rust", "filter": {"key": "source", "op": "eq", "value": "docs"}}"#,
371        );
372        let retriever = build_retriever(llm, store.clone(), &["source"]);
373
374        let results = retriever
375            .retrieve("告诉我关于 Rust 的文档", 10)
376            .await
377            .unwrap();
378        let contents: HashSet<&str> = results.iter().map(|d| d.content.as_str()).collect();
379        assert_eq!(
380            contents,
381            HashSet::from(["Rust systems programming", "Rust borrow checker"])
382        );
383    }
384
385    /// S2: a disallowed field is blocked by the whitelist — explicit error, not a silent filter drop into an unfiltered search.
386    #[tokio::test]
387    async fn test_self_query_rejects_disallowed_attribute() {
388        let store = store_with_docs().await;
389        let llm = MockChatModel::new(
390            r#"{"query": "rust", "filter": {"key": "nonexistent", "op": "eq", "value": 1}}"#,
391        );
392        let retriever = build_retriever(llm, store.clone(), &["source"]);
393
394        let err = retriever.retrieve("rust", 10).await.unwrap_err();
395        assert!(matches!(err, RetrieverError::InvalidFilter(_)));
396        assert!(err.to_string().contains("nonexistent"));
397    }
398
399    /// S2: a whitelist-outside field inside a nested And/Or also errors out explicitly.
400    #[tokio::test]
401    async fn test_self_query_rejects_disallowed_attribute_in_nested_and() {
402        let store = store_with_docs().await;
403        let llm = MockChatModel::new(
404            r#"{"query": "rust", "filter": {"And": [{"key": "source", "op": "eq", "value": "docs"}, {"key": "private", "op": "eq", "value": false}]}}"#,
405        );
406        let retriever = build_retriever(llm, store.clone(), &["source"]);
407
408        let err = retriever.retrieve("rust", 10).await.unwrap_err();
409        assert!(matches!(err, RetrieverError::InvalidFilter(_)));
410        assert!(err.to_string().contains("private"));
411    }
412
413    /// S2: empty whitelist = filtering entirely disabled — filter ignored and all docs returned, no error.
414    #[tokio::test]
415    async fn test_self_query_empty_whitelist_ignores_filter() {
416        let store = store_with_docs().await;
417        let llm = MockChatModel::new(
418            r#"{"query": "rust", "filter": {"key": "source", "op": "eq", "value": "docs"}}"#,
419        );
420        let retriever = build_retriever(llm, store.clone(), &[]);
421
422        let results = retriever.retrieve("rust", 10).await.unwrap();
423        assert_eq!(results.len(), 3, "filtering disabled -> all docs returned");
424    }
425
426    /// S4: text fallback — when the model outputs plain text, the whole text is used as the query, no filter.
427    #[tokio::test]
428    async fn test_self_query_text_fallback_query_only() {
429        let store = store_with_docs().await;
430        let llm = MockChatModel::new("rust programming");
431        let retriever = build_retriever(llm, store.clone(), &["source"]);
432
433        let results = retriever.retrieve("rust", 10).await.unwrap();
434        assert_eq!(
435            results.len(),
436            3,
437            "plain-text fallback must search without filter"
438        );
439    }
440
441    /// S4: derived shapes + nested combinations deserialize too (LLM outputs And/Or).
442    #[tokio::test]
443    async fn test_self_query_nested_filter_parses() {
444        let store = store_with_docs().await;
445        let llm = MockChatModel::new(
446            r#"{"query": "rust", "filter": {"And": [{"key": "source", "op": "eq", "value": "docs"}]}}"#,
447        );
448        let retriever = build_retriever(llm, store.clone(), &["source"]);
449
450        let results = retriever.retrieve("rust", 10).await.unwrap();
451        assert_eq!(results.len(), 2);
452    }
453
454    /// S4: `RetrieverRunnable` composes into an LCEL chain that compiles and runs.
455    #[tokio::test]
456    async fn test_self_query_pipes_into_retriever_runnable() {
457        use crate::RetrieverRunnable;
458        use lc_core::runnables::RunnableExt;
459
460        let store = store_with_docs().await;
461        let llm = MockChatModel::new(
462            r#"{"query": "rust", "filter": {"key": "source", "op": "eq", "value": "docs"}}"#,
463        );
464        let retriever: Arc<dyn RetrieverTrait> =
465            Arc::new(build_retriever(llm, store.clone(), &["source"]));
466
467        let step = RetrieverRunnable::new(retriever, 10);
468        let docs = step
469            .invoke("告诉我 Rust 的文档".to_string(), None)
470            .await
471            .unwrap();
472        assert_eq!(docs.len(), 2);
473
474        // Chain one more step: verify the type chain (Vec<Document> -> usize).
475        let count = step
476            .pipe(lc_core::runnables::RunnableLambda::new_sync(
477                |docs: Vec<Document>| docs.len(),
478            ))
479            .invoke("rust 文档".to_string(), None)
480            .await
481            .unwrap();
482        assert_eq!(count, 2);
483    }
484}