1use 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#[derive(Debug, Clone, PartialEq, Deserialize)]
37pub struct SelfQueryArgs {
38 pub query: String,
40 #[serde(default)]
42 pub filter: Option<MetadataFilter>,
43}
44
45pub 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 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 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 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 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 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 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
184fn 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 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 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 #[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 #[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 #[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 #[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 #[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 #[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 #[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 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}