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