Skip to main content

lc_vector_stores/
qdrant.rs

1// lc-vector-stores/src/qdrant.rs
2//! Qdrant 向量存储实现
3
4use crate::{Document, SearchResult, VectorStore, VectorStoreError};
5use async_trait::async_trait;
6use qdrant_client::{
7    qdrant::{
8        Condition, CreateCollectionBuilder, DeletePointsBuilder, Distance, Filter, PointId,
9        PointStruct, QueryPointsBuilder, UpsertPointsBuilder, VectorParamsBuilder,
10    },
11    Payload, Qdrant,
12};
13use std::collections::HashMap;
14use std::sync::Arc;
15use uuid::Uuid;
16
17/// Qdrant 配置
18#[derive(Debug, Clone)]
19pub struct QdrantConfig {
20    pub url: String,
21    pub collection_name: String,
22    pub vector_size: usize,
23    pub distance: QdrantDistance,
24}
25
26/// Qdrant 距离度量类型
27#[derive(Debug, Clone, Copy)]
28pub enum QdrantDistance {
29    Cosine,
30    Euclid,
31    Dot,
32}
33
34impl From<QdrantDistance> for Distance {
35    fn from(dist: QdrantDistance) -> Self {
36        match dist {
37            QdrantDistance::Cosine => Distance::Cosine,
38            QdrantDistance::Euclid => Distance::Euclid,
39            QdrantDistance::Dot => Distance::Dot,
40        }
41    }
42}
43
44impl Default for QdrantConfig {
45    fn default() -> Self {
46        Self {
47            url: "http://localhost:6334".to_string(),
48            collection_name: "langchainrust".to_string(),
49            vector_size: 1536,
50            distance: QdrantDistance::Cosine,
51        }
52    }
53}
54
55impl QdrantConfig {
56    pub fn new(url: impl Into<String>, collection_name: impl Into<String>) -> Self {
57        Self {
58            url: url.into(),
59            collection_name: collection_name.into(),
60            ..Default::default()
61        }
62    }
63
64    pub fn with_vector_size(mut self, size: usize) -> Self {
65        self.vector_size = size;
66        self
67    }
68
69    pub fn with_distance(mut self, distance: QdrantDistance) -> Self {
70        self.distance = distance;
71        self
72    }
73}
74
75/// Qdrant 向量存储
76pub struct QdrantVectorStore {
77    client: Arc<Qdrant>,
78    config: QdrantConfig,
79}
80
81impl QdrantVectorStore {
82    pub async fn new(config: QdrantConfig) -> Result<Self, VectorStoreError> {
83        let client = Qdrant::from_url(&config.url)
84            .build()
85            .map_err(|e| VectorStoreError::ConnectionError(format!("连接 Qdrant 失败: {}", e)))?;
86
87        let client = Arc::new(client);
88
89        let exists = client
90            .collection_exists(&config.collection_name)
91            .await
92            .map_err(|e| VectorStoreError::StorageError(format!("检查集合失败: {}", e)))?;
93
94        if !exists {
95            client
96                .create_collection(
97                    CreateCollectionBuilder::new(&config.collection_name).vectors_config(
98                        VectorParamsBuilder::new(
99                            config.vector_size as u64,
100                            Distance::from(config.distance),
101                        ),
102                    ),
103                )
104                .await
105                .map_err(|e| VectorStoreError::StorageError(format!("创建集合失败: {}", e)))?;
106        }
107
108        Ok(Self { client, config })
109    }
110
111    pub async fn from_env() -> Result<Self, VectorStoreError> {
112        let url =
113            std::env::var("QDRANT_URL").unwrap_or_else(|_| "http://localhost:6334".to_string());
114        let collection_name =
115            std::env::var("QDRANT_COLLECTION").unwrap_or_else(|_| "langchainrust".to_string());
116
117        Self::new(QdrantConfig::new(url, collection_name)).await
118    }
119
120    pub async fn delete_by_metadata(
121        &self,
122        key: &str,
123        value: &str,
124    ) -> Result<usize, VectorStoreError> {
125        let filter = Filter::must([Condition::matches(key, value.to_string())]);
126
127        // Q4: 先按 metadata 过滤统计匹配点,再删除,返回真实删除数。
128        // 旧实现删完直接 Ok(0) —— 无论删除是否生效,上层都误以为"没删任何数据"。
129        let total = self.count().await as u64;
130        let matched = self
131            .client
132            .query(
133                QueryPointsBuilder::new(&self.config.collection_name)
134                    .query(vec![0.0; self.config.vector_size])
135                    .filter(filter.clone())
136                    .limit(total.max(1))
137                    .with_payload(false),
138            )
139            .await
140            .map_err(|e| VectorStoreError::StorageError(format!("按metadata统计失败: {}", e)))?;
141
142        let deleted = matched.result.len();
143
144        if deleted > 0 {
145            self.client
146                .delete_points(
147                    DeletePointsBuilder::new(&self.config.collection_name).points(filter),
148                )
149                .await
150                .map_err(|e| {
151                    VectorStoreError::StorageError(format!("按metadata删除失败: {}", e))
152                })?;
153        }
154
155        Ok(deleted)
156    }
157}
158
159#[async_trait]
160impl VectorStore for QdrantVectorStore {
161    async fn add_documents(
162        &self,
163        documents: Vec<Document>,
164        embeddings: Vec<Vec<f32>>,
165    ) -> Result<Vec<String>, VectorStoreError> {
166        if documents.len() != embeddings.len() {
167            return Err(VectorStoreError::StorageError(
168                "文档数量和嵌入向量数量不匹配".to_string(),
169            ));
170        }
171
172        if documents.is_empty() {
173            return Ok(Vec::new());
174        }
175
176        for embedding in &embeddings {
177            if embedding.len() != self.config.vector_size {
178                return Err(VectorStoreError::StorageError(format!(
179                    "向量维度不匹配: 期望 {}, 实际 {}",
180                    self.config.vector_size,
181                    embedding.len()
182                )));
183            }
184        }
185
186        let mut ids = Vec::new();
187        let mut points = Vec::new();
188
189        for (doc, embedding) in documents.into_iter().zip(embeddings) {
190            let user_id = doc.id.clone().unwrap_or_else(|| Uuid::new_v4().to_string());
191
192            // Qdrant PointId 只接受 UUID 或数字,所以生成内部 UUID
193            let internal_uuid = Uuid::new_v4();
194            let point_id = PointId::from(internal_uuid.to_string());
195
196            let mut payload = Payload::new();
197            payload.insert("content", doc.content.clone());
198            payload.insert("doc_id", user_id.clone()); // 用户 ID 存在 payload 中
199
200            for (key, value) in &doc.metadata {
201                payload.insert(key.clone(), value.clone());
202            }
203
204            let point = PointStruct::new(point_id, embedding, payload);
205            points.push(point);
206            ids.push(user_id);
207        }
208
209        self.client
210            .upsert_points(UpsertPointsBuilder::new(
211                &self.config.collection_name,
212                points,
213            ))
214            .await
215            .map_err(|e| VectorStoreError::StorageError(format!("插入文档失败: {}", e)))?;
216
217        Ok(ids)
218    }
219
220    async fn similarity_search(
221        &self,
222        query_embedding: &[f32],
223        k: usize,
224    ) -> Result<Vec<SearchResult>, VectorStoreError> {
225        if query_embedding.len() != self.config.vector_size {
226            return Err(VectorStoreError::StorageError(format!(
227                "查询向量维度不匹配: 期望 {}, 实际 {}",
228                self.config.vector_size,
229                query_embedding.len()
230            )));
231        }
232
233        let search_result = self
234            .client
235            .query(
236                QueryPointsBuilder::new(&self.config.collection_name)
237                    .query(query_embedding.to_vec())
238                    .limit(k as u64)
239                    .with_payload(true),
240            )
241            .await
242            .map_err(|e| VectorStoreError::StorageError(format!("搜索失败: {}", e)))?;
243
244        let results: Vec<SearchResult> = search_result
245            .result
246            .into_iter()
247            .map(|scored_point| {
248                let payload = scored_point.payload;
249
250                let content = payload
251                    .get("content")
252                    .and_then(|v| v.as_str())
253                    .map(|s| s.as_str())
254                    .unwrap_or("")
255                    .to_string();
256
257                let id = payload
258                    .get("doc_id")
259                    .and_then(|v| v.as_str())
260                    .map(|s| s.to_string());
261
262                let mut metadata = HashMap::new();
263                for (key, value) in &payload {
264                    if key != "content" && key != "doc_id" {
265                        if let Some(s) = value.as_str() {
266                            metadata.insert(key.clone(), s.clone());
267                        }
268                    }
269                }
270
271                SearchResult {
272                    document: Document {
273                        content,
274                        metadata,
275                        id,
276                    },
277                    score: scored_point.score,
278                }
279            })
280            .collect();
281
282        Ok(results)
283    }
284
285    async fn get_document(&self, id: &str) -> Result<Option<Document>, VectorStoreError> {
286        let filter = Filter::must([Condition::matches("doc_id", id.to_string())]);
287
288        let results = self
289            .client
290            .query(
291                QueryPointsBuilder::new(&self.config.collection_name)
292                    .query(vec![0.0; self.config.vector_size])
293                    .filter(filter)
294                    .limit(1)
295                    .with_payload(true),
296            )
297            .await
298            .map_err(|e| VectorStoreError::StorageError(format!("获取文档失败: {}", e)))?;
299
300        if let Some(point) = results.result.first() {
301            let payload_map = point.payload.clone();
302
303            let content = payload_map
304                .get("content")
305                .and_then(|v| v.as_str())
306                .map(|s| s.as_str())
307                .unwrap_or("")
308                .to_string();
309
310            let doc_id = payload_map
311                .get("doc_id")
312                .and_then(|v| v.as_str())
313                .map(|s| s.to_string());
314
315            let mut metadata = HashMap::new();
316            for (key, value) in &payload_map {
317                if key != "content" && key != "doc_id" {
318                    if let Some(s) = value.as_str() {
319                        metadata.insert(key.clone(), s.clone());
320                    }
321                }
322            }
323
324            Ok(Some(Document {
325                content,
326                metadata,
327                id: doc_id,
328            }))
329        } else {
330            Ok(None)
331        }
332    }
333
334    async fn get_embedding(&self, id: &str) -> Result<Option<Vec<f32>>, VectorStoreError> {
335        let filter = Filter::must([Condition::matches("doc_id", id.to_string())]);
336
337        let results = self
338            .client
339            .query(
340                QueryPointsBuilder::new(&self.config.collection_name)
341                    .query(vec![0.0; self.config.vector_size])
342                    .filter(filter)
343                    .limit(1)
344                    .with_payload(true),
345            )
346            .await
347            .map_err(|e| VectorStoreError::StorageError(format!("获取向量失败: {}", e)))?;
348
349        if let Some(point) = results.result.first() {
350            if let Some(vectors) = &point.vectors {
351                if let Some(qdrant_client::qdrant::vector_output::Vector::Dense(dense)) =
352                    vectors.get_vector()
353                {
354                    return Ok(Some(dense.data.clone()));
355                }
356            }
357        }
358        Ok(None)
359    }
360
361    async fn delete_document(&self, id: &str) -> Result<(), VectorStoreError> {
362        let filter = Filter::must([Condition::matches("doc_id", id.to_string())]);
363
364        self.client
365            .delete_points(DeletePointsBuilder::new(&self.config.collection_name).points(filter))
366            .await
367            .map_err(|e| VectorStoreError::StorageError(format!("删除文档失败: {}", e)))?;
368
369        Ok(())
370    }
371
372    async fn count(&self) -> usize {
373        let info = self
374            .client
375            .collection_info(&self.config.collection_name)
376            .await;
377
378        info.map(|i| i.result.and_then(|r| r.points_count).unwrap_or(0) as usize)
379            .unwrap_or(0)
380    }
381
382    async fn clear(&self) -> Result<(), VectorStoreError> {
383        let collection_name = self.config.collection_name.clone();
384
385        self.client
386            .delete_collection(&collection_name)
387            .await
388            .map_err(|e| VectorStoreError::StorageError(format!("删除集合失败: {}", e)))?;
389
390        self.client
391            .create_collection(
392                CreateCollectionBuilder::new(&collection_name).vectors_config(
393                    VectorParamsBuilder::new(
394                        self.config.vector_size as u64,
395                        Distance::from(self.config.distance),
396                    ),
397                ),
398            )
399            .await
400            .map_err(|e| VectorStoreError::StorageError(format!("重建集合失败: {}", e)))?;
401
402        Ok(())
403    }
404}
405
406#[cfg(test)]
407mod tests {
408    use super::*;
409
410    #[test]
411    fn test_config_default() {
412        let config = QdrantConfig::default();
413        assert_eq!(config.url, "http://localhost:6334");
414        assert_eq!(config.collection_name, "langchainrust");
415        assert_eq!(config.vector_size, 1536);
416    }
417
418    #[test]
419    fn test_config_builder() {
420        let config = QdrantConfig::new("http://custom:6334", "test_collection")
421            .with_vector_size(3072)
422            .with_distance(QdrantDistance::Euclid);
423
424        assert_eq!(config.url, "http://custom:6334");
425        assert_eq!(config.collection_name, "test_collection");
426        assert_eq!(config.vector_size, 3072);
427        assert!(matches!(config.distance, QdrantDistance::Euclid));
428    }
429}