Skip to main content

lc_rag/
self_query.rs

1// lc-rag/src/self_query.rs
2//! SelfQueryRetriever — 自查询检索器
3//!
4//! 让 LLM 把自然语言查询拆成 `{ query, filter }`:清洗后的查询词走向量检索,
5//! 解析出的 [`MetadataFilter`] 交给 `vector_store.similarity_search_with_filter`
6//! 做元数据过滤(依赖 S3 的统一过滤能力)。拆解走 [`lc_core::judge::structured_call`]
7//! (绑定工具拿结构化参数,模型不支持时回落文本解析),与 Guardrails / Evaluation
8//! 同一执行路径;`allowed_attributes` 白名单拦截 LLM 用不存在的字段过滤。
9//!
10//! **不静默降级**:过滤器引用了白名单外字段时显式返回
11//! [`RetrieverError::InvalidFilter`],绝不丢弃过滤器回落无过滤检索——那会让
12//! 本该被过滤排除的数据被返回(数据面过曝)。空白名单 = 过滤整体禁用,过滤器
13//! 一律忽略并记 warning(这是「禁用过滤」的既定模式,不是静默降级)。
14
15use std::sync::Arc;
16
17use async_trait::async_trait;
18use lc_core::judge::{structured_call, StructuredJudgeError};
19use lc_core::language_models::BaseChatModel;
20use lc_core::tools::ToolDefinition;
21use lc_embeddings::Embeddings;
22use lc_schema::Message;
23use lc_vector_stores::{Document, MetadataFilter, SearchResult, VectorStore};
24use serde::Deserialize;
25
26use crate::retriever::{RetrieverError, RetrieverTrait};
27
28/// LLM 拆解出的结构化参数:清洗后的查询词 + 可选元数据过滤。
29///
30/// `filter` 直接以 [`MetadataFilter`] 反序列化(filter.rs 对 JSON 形状做了宽松
31/// 处理,兼容 LLM 输出差异);缺省为无过滤。
32#[derive(Debug, Clone, PartialEq, Deserialize)]
33pub struct SelfQueryArgs {
34    /// 清洗掉过滤约束后的纯语义查询词。
35    pub query: String,
36    /// 元数据过滤条件(可选)。
37    #[serde(default)]
38    pub filter: Option<MetadataFilter>,
39}
40
41/// 自查询检索器:LLM 拆解查询 → 白名单校验 → 过滤相似度检索。
42///
43/// 实现 [`RetrieverTrait`],可被 [`crate::RetrieverRunnable`] 包进 LCEL 链。
44pub struct SelfQueryRetriever<M: BaseChatModel> {
45    llm: Arc<M>,
46    store: Arc<dyn VectorStore>,
47    embeddings: Arc<dyn Embeddings>,
48    allowed_attributes: Vec<String>,
49}
50
51impl<M: BaseChatModel> SelfQueryRetriever<M> {
52    /// 创建自查询检索器。
53    ///
54    /// - `llm`:负责拆解自然语言查询的模型(支持结构化输出,或可文本回落)。
55    /// - `store` / `embeddings`:过滤相似度检索用的向量存储与嵌入模型。
56    /// - `allowed_attributes`:允许出现在过滤字段的白名单;**空白名单 = 过滤整体
57    ///   禁用**(LLM 返回的 filter 一律忽略并记 warning);**非空白名单下,过滤器
58    ///   引用白名单外字段会显式返回 [`RetrieverError::InvalidFilter`]**,不静默
59    ///   丢弃后回落无过滤检索。
60    pub fn new(
61        llm: impl Into<Arc<M>>,
62        store: Arc<dyn VectorStore>,
63        embeddings: Arc<dyn Embeddings>,
64        allowed_attributes: Vec<String>,
65    ) -> Self {
66        Self {
67            llm: llm.into(),
68            store,
69            embeddings,
70            allowed_attributes,
71        }
72    }
73
74    /// 自查询工具定义:让 LLM 以 `{ query, filter }` 结构化返回。
75    fn self_query_tool() -> ToolDefinition {
76        ToolDefinition::new(
77            "self_query",
78            "把用户的自然语言查询拆成纯语义查询词和可选的元数据过滤条件。",
79        )
80        .with_parameters(serde_json::json!({
81            "type": "object",
82            "properties": {
83                "query": {
84                    "type": "string",
85                    "description": "清洗掉过滤约束后的纯语义查询词"
86                },
87                "filter": {
88                    "type": ["object", "null"],
89                    "description": "元数据过滤条件(MetadataFilter JSON):单条件 {\"Field\": {\"key\", \"op\", \"value\"}},组合 {\"And\": [...]} / {\"Or\": [...]};op 取 Eq Ne Gt Gte Lt Lte In Nin(In/Nin 的 value 为数组);无过滤时为 null"
90                }
91            },
92            "required": ["query"]
93        }))
94    }
95
96    /// 构建提示词:告诉 LLM 可用字段与输出格式。
97    fn build_prompt(&self, query: &str) -> String {
98        let allowed = if self.allowed_attributes.is_empty() {
99            "无(本检索器不启用元数据过滤,filter 必须为 null)".to_string()
100        } else {
101            self.allowed_attributes.join(", ")
102        };
103        format!(
104            "把下面的自然语言查询拆成两部分:纯语义查询词(query)和可选的元数据过滤条件(filter)。\n\
105             filter 的 key 只能取以下允许字段之一: {allowed}\n\
106             filter 的 JSON 形状:单条件 {{\"Field\": {{\"key\": ..., \"op\": ..., \"value\": ...}}}};\n\
107             组合条件 {{\"And\": [...]}} / {{\"Or\": [...]}};op 取 Eq Ne Gt Gte Lt Lte In Nin(In/Nin 的 value 为数组)。\n\
108             没有过滤需求时 filter 为 null。\n\
109             用户查询: {query}"
110        )
111    }
112
113    /// 调用 LLM 拆解查询(结构化或文本回落)。
114    async fn parse_query(&self, query: &str) -> Result<SelfQueryArgs, RetrieverError> {
115        let messages = vec![Message::human(self.build_prompt(query))];
116        structured_call(
117            &*self.llm,
118            Self::self_query_tool(),
119            messages,
120            parse_text_fallback,
121        )
122        .await
123        .map_err(|e| RetrieverError::LlmError(e.to_string()))
124    }
125
126    /// 白名单校验。
127    ///
128    /// - 无过滤器 → `Ok(None)`。
129    /// - 空白名单 = 过滤整体禁用:过滤器一律忽略(这是「禁用过滤」的既定模式,
130    ///   不是静默降级),记 warning 后 `Ok(None)`。
131    /// - 非空白名单下,过滤器引用白名单外字段 → `Err(InvalidFilter)`。绝不丢弃
132    ///   过滤器回落无过滤检索,否则本该被过滤排除的数据会被返回(数据面过曝)。
133    fn validated_filter(
134        &self,
135        filter: &Option<MetadataFilter>,
136    ) -> Result<Option<MetadataFilter>, RetrieverError> {
137        let Some(f) = filter else {
138            return Ok(None);
139        };
140        if self.allowed_attributes.is_empty() {
141            log::warn!(
142                "SelfQuery: filtering is disabled (empty allowed_attributes); ignoring filter"
143            );
144            return Ok(None);
145        }
146        if let Some(field) = Self::disallowed_field(f, &self.allowed_attributes) {
147            return Err(RetrieverError::InvalidFilter(format!(
148                "filter references `{field}`, which is not in allowed_attributes [{}]; \
149                 refusing to degrade to an unfiltered search",
150                self.allowed_attributes.join(", ")
151            )));
152        }
153        Ok(filter.clone())
154    }
155
156    /// 第一个引用白名单外字段的 key,遍历嵌套的 And/Or 组合;全部在列则 `None`。
157    fn disallowed_field<'a>(filter: &'a MetadataFilter, allowed: &[String]) -> Option<&'a String> {
158        match filter {
159            MetadataFilter::Field { key, .. } => {
160                if allowed.iter().any(|a| a == key) {
161                    None
162                } else {
163                    Some(key)
164                }
165            }
166            MetadataFilter::And(items) | MetadataFilter::Or(items) => items
167                .iter()
168                .find_map(|f| Self::disallowed_field(f, allowed)),
169        }
170    }
171}
172
173/// 文本回落解析:优先尝试整段 JSON(LLM 输出结构时);否则把整段文本当查询词。
174fn parse_text_fallback(raw: &str) -> Result<SelfQueryArgs, StructuredJudgeError> {
175    let trimmed = raw.trim();
176    if !trimmed.is_empty() {
177        if let Ok(args) = serde_json::from_str::<SelfQueryArgs>(trimmed) {
178            return Ok(args);
179        }
180    }
181    let query = raw.trim().to_string();
182    if query.is_empty() {
183        return Err(StructuredJudgeError::Parse(
184            "self-query fallback produced an empty query".to_string(),
185        ));
186    }
187    Ok(SelfQueryArgs {
188        query,
189        filter: None,
190    })
191}
192
193#[async_trait]
194impl<M: BaseChatModel> RetrieverTrait for SelfQueryRetriever<M> {
195    async fn retrieve(&self, query: &str, k: usize) -> Result<Vec<Document>, RetrieverError> {
196        let results = self.retrieve_with_scores(query, k).await?;
197        Ok(results.into_iter().map(|r| r.document).collect())
198    }
199
200    async fn retrieve_with_scores(
201        &self,
202        query: &str,
203        k: usize,
204    ) -> Result<Vec<SearchResult>, RetrieverError> {
205        let args = self.parse_query(query).await?;
206        let filter = self.validated_filter(&args.filter)?;
207
208        let query_embedding = self
209            .embeddings
210            .embed_query(&args.query)
211            .await
212            .map_err(|e| RetrieverError::EmbeddingError(e.to_string()))?;
213
214        self.store
215            .similarity_search_with_filter(&query_embedding, k, filter.as_ref())
216            .await
217            .map_err(RetrieverError::from)
218    }
219
220    async fn add_documents(&self, documents: Vec<Document>) -> Result<(), RetrieverError> {
221        let texts: Vec<&str> = documents.iter().map(|d| d.content.as_str()).collect();
222        let embeddings = self
223            .embeddings
224            .embed_documents(&texts)
225            .await
226            .map_err(|e| RetrieverError::EmbeddingError(e.to_string()))?;
227        self.store.add_documents(documents, embeddings).await?;
228        Ok(())
229    }
230}
231
232#[cfg(test)]
233mod tests {
234    use super::*;
235    use async_trait::async_trait;
236    use futures_util::Stream;
237    use lc_core::language_models::{LLMResult, StreamChunk};
238    use lc_core::runnables::RunnableConfig;
239    use lc_core::{BaseLanguageModel, Runnable};
240    use lc_embeddings::MockEmbeddings;
241    use lc_vector_stores::InMemoryVectorStore;
242    use std::collections::HashSet;
243    use std::pin::Pin;
244    use std::sync::atomic::{AtomicUsize, Ordering};
245    use std::sync::Arc;
246
247    /// 固定回复的 mock 聊天模型(走文本回落路径,不实现 bind_tools)。
248    struct MockChatModel {
249        reply: String,
250        calls: AtomicUsize,
251    }
252
253    impl MockChatModel {
254        fn new(reply: &str) -> Self {
255            Self {
256                reply: reply.to_string(),
257                calls: AtomicUsize::new(0),
258            }
259        }
260    }
261
262    #[async_trait]
263    impl Runnable<Vec<Message>, LLMResult> for MockChatModel {
264        type Error = MockChatError;
265        async fn invoke(
266            &self,
267            _input: Vec<Message>,
268            _config: Option<RunnableConfig>,
269        ) -> Result<LLMResult, Self::Error> {
270            Err(MockChatError)
271        }
272    }
273
274    #[async_trait]
275    impl BaseLanguageModel<Vec<Message>, LLMResult> for MockChatModel {
276        fn model_name(&self) -> &str {
277            "self-query-mock"
278        }
279        fn get_num_tokens(&self, t: &str) -> usize {
280            t.len()
281        }
282        fn with_temperature(self, _: f32) -> Self {
283            self
284        }
285        fn with_max_tokens(self, _: usize) -> Self {
286            self
287        }
288    }
289
290    #[derive(Debug)]
291    struct MockChatError;
292    impl std::fmt::Display for MockChatError {
293        fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
294            write!(f, "mock chat error")
295        }
296    }
297    impl std::error::Error for MockChatError {}
298
299    #[async_trait]
300    impl BaseChatModel for MockChatModel {
301        async fn chat(
302            &self,
303            _messages: Vec<Message>,
304            _config: Option<RunnableConfig>,
305        ) -> Result<LLMResult, Self::Error> {
306            self.calls.fetch_add(1, Ordering::SeqCst);
307            Ok(LLMResult {
308                content: self.reply.clone(),
309                model: "self-query-mock".to_string(),
310                token_usage: None,
311                tool_calls: None,
312                thinking_content: None,
313            })
314        }
315        async fn stream_chat(
316            &self,
317            _messages: Vec<Message>,
318            _config: Option<RunnableConfig>,
319        ) -> Result<Pin<Box<dyn Stream<Item = Result<StreamChunk, Self::Error>> + Send>>, Self::Error>
320        {
321            Err(MockChatError)
322        }
323    }
324
325    /// 造一个带 metadata 的内存向量存储 + mock 嵌入,返回 store 供断言。
326    async fn store_with_docs() -> Arc<InMemoryVectorStore> {
327        let store = Arc::new(InMemoryVectorStore::new());
328        let embeddings = Arc::new(MockEmbeddings::new(64));
329        let docs = vec![
330            Document::new("Rust systems programming").with_metadata("source", "docs"),
331            Document::new("Rust borrow checker").with_metadata("source", "docs"),
332            Document::new("Python scripting").with_metadata("source", "blog"),
333        ];
334        let texts: Vec<&str> = docs.iter().map(|d| d.content.as_str()).collect();
335        let vecs = embeddings.embed_documents(&texts).await.unwrap();
336        store.add_documents(docs, vecs).await.unwrap();
337        store
338    }
339
340    fn build_retriever(
341        llm: MockChatModel,
342        store: Arc<dyn VectorStore>,
343        allowed: &[&str],
344    ) -> SelfQueryRetriever<MockChatModel> {
345        SelfQueryRetriever::new(
346            Arc::new(llm),
347            store,
348            Arc::new(MockEmbeddings::new(64)),
349            allowed.iter().map(|s| s.to_string()).collect(),
350        )
351    }
352
353    /// S4: 拆出的 filter 正确落到检索 —— 只返回匹配 source=docs 的文档。
354    #[tokio::test]
355    async fn test_self_query_filter_reaches_search() {
356        let store = store_with_docs().await;
357        let llm = MockChatModel::new(
358            r#"{"query": "rust", "filter": {"key": "source", "op": "eq", "value": "docs"}}"#,
359        );
360        let retriever = build_retriever(llm, store.clone(), &["source"]);
361
362        let results = retriever
363            .retrieve("告诉我关于 Rust 的文档", 10)
364            .await
365            .unwrap();
366        let contents: HashSet<&str> = results.iter().map(|d| d.content.as_str()).collect();
367        assert_eq!(
368            contents,
369            HashSet::from(["Rust systems programming", "Rust borrow checker"])
370        );
371    }
372
373    /// S2: 非法字段被白名单拦截 —— 显式报错,不静默丢弃过滤器回落无过滤检索。
374    #[tokio::test]
375    async fn test_self_query_rejects_disallowed_attribute() {
376        let store = store_with_docs().await;
377        let llm = MockChatModel::new(
378            r#"{"query": "rust", "filter": {"key": "nonexistent", "op": "eq", "value": 1}}"#,
379        );
380        let retriever = build_retriever(llm, store.clone(), &["source"]);
381
382        let err = retriever.retrieve("rust", 10).await.unwrap_err();
383        assert!(matches!(err, RetrieverError::InvalidFilter(_)));
384        assert!(err.to_string().contains("nonexistent"));
385    }
386
387    /// S2: 嵌套 And/Or 里出现白名单外字段同样显式报错。
388    #[tokio::test]
389    async fn test_self_query_rejects_disallowed_attribute_in_nested_and() {
390        let store = store_with_docs().await;
391        let llm = MockChatModel::new(
392            r#"{"query": "rust", "filter": {"And": [{"key": "source", "op": "eq", "value": "docs"}, {"key": "private", "op": "eq", "value": false}]}}"#,
393        );
394        let retriever = build_retriever(llm, store.clone(), &["source"]);
395
396        let err = retriever.retrieve("rust", 10).await.unwrap_err();
397        assert!(matches!(err, RetrieverError::InvalidFilter(_)));
398        assert!(err.to_string().contains("private"));
399    }
400
401    /// S2: 空白名单 = 过滤整体禁用 —— filter 忽略并返回全部文档,不报错。
402    #[tokio::test]
403    async fn test_self_query_empty_whitelist_ignores_filter() {
404        let store = store_with_docs().await;
405        let llm = MockChatModel::new(
406            r#"{"query": "rust", "filter": {"key": "source", "op": "eq", "value": "docs"}}"#,
407        );
408        let retriever = build_retriever(llm, store.clone(), &[]);
409
410        let results = retriever.retrieve("rust", 10).await.unwrap();
411        assert_eq!(results.len(), 3, "filtering disabled -> all docs returned");
412    }
413
414    /// S4: 文本回落 —— 模型输出纯文本时整段当查询词,无过滤。
415    #[tokio::test]
416    async fn test_self_query_text_fallback_query_only() {
417        let store = store_with_docs().await;
418        let llm = MockChatModel::new("rust programming");
419        let retriever = build_retriever(llm, store.clone(), &["source"]);
420
421        let results = retriever.retrieve("rust", 10).await.unwrap();
422        assert_eq!(
423            results.len(),
424            3,
425            "plain-text fallback must search without filter"
426        );
427    }
428
429    /// S4: 派生形状 + 嵌套组合也能反序列化(LLM 输出 And/Or)。
430    #[tokio::test]
431    async fn test_self_query_nested_filter_parses() {
432        let store = store_with_docs().await;
433        let llm = MockChatModel::new(
434            r#"{"query": "rust", "filter": {"And": [{"key": "source", "op": "eq", "value": "docs"}]}}"#,
435        );
436        let retriever = build_retriever(llm, store.clone(), &["source"]);
437
438        let results = retriever.retrieve("rust", 10).await.unwrap();
439        assert_eq!(results.len(), 2);
440    }
441
442    /// S4: `RetrieverRunnable` 组合进 LCEL 链编译并执行。
443    #[tokio::test]
444    async fn test_self_query_pipes_into_retriever_runnable() {
445        use crate::RetrieverRunnable;
446        use lc_core::runnables::RunnableExt;
447
448        let store = store_with_docs().await;
449        let llm = MockChatModel::new(
450            r#"{"query": "rust", "filter": {"key": "source", "op": "eq", "value": "docs"}}"#,
451        );
452        let retriever: Arc<dyn RetrieverTrait> =
453            Arc::new(build_retriever(llm, store.clone(), &["source"]));
454
455        let step = RetrieverRunnable::new(retriever, 10);
456        let docs = step
457            .invoke("告诉我 Rust 的文档".to_string(), None)
458            .await
459            .unwrap();
460        assert_eq!(docs.len(), 2);
461
462        // 链上再挂一步:验证类型链通(Vec<Document> → usize)。
463        let count = step
464            .pipe(lc_core::runnables::RunnableLambda::new_sync(
465                |docs: Vec<Document>| docs.len(),
466            ))
467            .invoke("rust 文档".to_string(), None)
468            .await
469            .unwrap();
470        assert_eq!(count, 2);
471    }
472}