1use 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#[derive(Debug, Clone, PartialEq, Deserialize)]
33pub struct SelfQueryArgs {
34 pub query: String,
36 #[serde(default)]
38 pub filter: Option<MetadataFilter>,
39}
40
41pub 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 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 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 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 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 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 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
173fn 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 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 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 #[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 #[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 #[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 #[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 #[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 #[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 #[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 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}