1use 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#[derive(Debug, Clone, PartialEq, Deserialize)]
28pub struct SelfQueryArgs {
29 pub query: String,
31 #[serde(default)]
33 pub filter: Option<MetadataFilter>,
34}
35
36pub 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 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 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 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 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 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
143fn 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 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 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 #[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 #[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 #[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 #[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 #[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 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}