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        self.client
128            .delete_points(DeletePointsBuilder::new(&self.config.collection_name).points(filter))
129            .await
130            .map_err(|e| VectorStoreError::StorageError(format!("按metadata删除失败: {}", e)))?;
131
132        Ok(0)
133    }
134}
135
136#[async_trait]
137impl VectorStore for QdrantVectorStore {
138    async fn add_documents(
139        &self,
140        documents: Vec<Document>,
141        embeddings: Vec<Vec<f32>>,
142    ) -> Result<Vec<String>, VectorStoreError> {
143        if documents.len() != embeddings.len() {
144            return Err(VectorStoreError::StorageError(
145                "文档数量和嵌入向量数量不匹配".to_string(),
146            ));
147        }
148
149        if documents.is_empty() {
150            return Ok(Vec::new());
151        }
152
153        for embedding in &embeddings {
154            if embedding.len() != self.config.vector_size {
155                return Err(VectorStoreError::StorageError(format!(
156                    "向量维度不匹配: 期望 {}, 实际 {}",
157                    self.config.vector_size,
158                    embedding.len()
159                )));
160            }
161        }
162
163        let mut ids = Vec::new();
164        let mut points = Vec::new();
165
166        for (doc, embedding) in documents.into_iter().zip(embeddings) {
167            let user_id = doc.id.clone().unwrap_or_else(|| Uuid::new_v4().to_string());
168
169            // Qdrant PointId 只接受 UUID 或数字,所以生成内部 UUID
170            let internal_uuid = Uuid::new_v4();
171            let point_id = PointId::from(internal_uuid.to_string());
172
173            let mut payload = Payload::new();
174            payload.insert("content", doc.content.clone());
175            payload.insert("doc_id", user_id.clone()); // 用户 ID 存在 payload 中
176
177            for (key, value) in &doc.metadata {
178                payload.insert(key.clone(), value.clone());
179            }
180
181            let point = PointStruct::new(point_id, embedding, payload);
182            points.push(point);
183            ids.push(user_id);
184        }
185
186        self.client
187            .upsert_points(UpsertPointsBuilder::new(
188                &self.config.collection_name,
189                points,
190            ))
191            .await
192            .map_err(|e| VectorStoreError::StorageError(format!("插入文档失败: {}", e)))?;
193
194        Ok(ids)
195    }
196
197    async fn similarity_search(
198        &self,
199        query_embedding: &[f32],
200        k: usize,
201    ) -> Result<Vec<SearchResult>, VectorStoreError> {
202        if query_embedding.len() != self.config.vector_size {
203            return Err(VectorStoreError::StorageError(format!(
204                "查询向量维度不匹配: 期望 {}, 实际 {}",
205                self.config.vector_size,
206                query_embedding.len()
207            )));
208        }
209
210        let search_result = self
211            .client
212            .query(
213                QueryPointsBuilder::new(&self.config.collection_name)
214                    .query(query_embedding.to_vec())
215                    .limit(k as u64)
216                    .with_payload(true),
217            )
218            .await
219            .map_err(|e| VectorStoreError::StorageError(format!("搜索失败: {}", e)))?;
220
221        let results: Vec<SearchResult> = search_result
222            .result
223            .into_iter()
224            .map(|scored_point| {
225                let payload = scored_point.payload;
226
227                let content = payload
228                    .get("content")
229                    .and_then(|v| v.as_str())
230                    .map(|s| s.as_str())
231                    .unwrap_or("")
232                    .to_string();
233
234                let id = payload
235                    .get("doc_id")
236                    .and_then(|v| v.as_str())
237                    .map(|s| s.to_string());
238
239                let mut metadata = HashMap::new();
240                for (key, value) in &payload {
241                    if key != "content" && key != "doc_id" {
242                        if let Some(s) = value.as_str() {
243                            metadata.insert(key.clone(), s.clone());
244                        }
245                    }
246                }
247
248                SearchResult {
249                    document: Document {
250                        content,
251                        metadata,
252                        id,
253                    },
254                    score: scored_point.score,
255                }
256            })
257            .collect();
258
259        Ok(results)
260    }
261
262    async fn get_document(&self, id: &str) -> Result<Option<Document>, VectorStoreError> {
263        let filter = Filter::must([Condition::matches("doc_id", id.to_string())]);
264
265        let results = self
266            .client
267            .query(
268                QueryPointsBuilder::new(&self.config.collection_name)
269                    .query(vec![0.0; self.config.vector_size])
270                    .filter(filter)
271                    .limit(1)
272                    .with_payload(true),
273            )
274            .await
275            .map_err(|e| VectorStoreError::StorageError(format!("获取文档失败: {}", e)))?;
276
277        if let Some(point) = results.result.first() {
278            let payload_map = point.payload.clone();
279
280            let content = payload_map
281                .get("content")
282                .and_then(|v| v.as_str())
283                .map(|s| s.as_str())
284                .unwrap_or("")
285                .to_string();
286
287            let doc_id = payload_map
288                .get("doc_id")
289                .and_then(|v| v.as_str())
290                .map(|s| s.to_string());
291
292            let mut metadata = HashMap::new();
293            for (key, value) in &payload_map {
294                if key != "content" && key != "doc_id" {
295                    if let Some(s) = value.as_str() {
296                        metadata.insert(key.clone(), s.clone());
297                    }
298                }
299            }
300
301            Ok(Some(Document {
302                content,
303                metadata,
304                id: doc_id,
305            }))
306        } else {
307            Ok(None)
308        }
309    }
310
311    async fn get_embedding(&self, id: &str) -> Result<Option<Vec<f32>>, VectorStoreError> {
312        let filter = Filter::must([Condition::matches("doc_id", id.to_string())]);
313
314        let results = self
315            .client
316            .query(
317                QueryPointsBuilder::new(&self.config.collection_name)
318                    .query(vec![0.0; self.config.vector_size])
319                    .filter(filter)
320                    .limit(1)
321                    .with_payload(true),
322            )
323            .await
324            .map_err(|e| VectorStoreError::StorageError(format!("获取向量失败: {}", e)))?;
325
326        if let Some(point) = results.result.first() {
327            if let Some(vectors) = &point.vectors {
328                if let Some(qdrant_client::qdrant::vector_output::Vector::Dense(dense)) =
329                    vectors.get_vector()
330                {
331                    return Ok(Some(dense.data.clone()));
332                }
333            }
334        }
335        Ok(None)
336    }
337
338    async fn delete_document(&self, id: &str) -> Result<(), VectorStoreError> {
339        let filter = Filter::must([Condition::matches("doc_id", id.to_string())]);
340
341        self.client
342            .delete_points(DeletePointsBuilder::new(&self.config.collection_name).points(filter))
343            .await
344            .map_err(|e| VectorStoreError::StorageError(format!("删除文档失败: {}", e)))?;
345
346        Ok(())
347    }
348
349    async fn count(&self) -> usize {
350        let info = self
351            .client
352            .collection_info(&self.config.collection_name)
353            .await;
354
355        info.map(|i| i.result.and_then(|r| r.points_count).unwrap_or(0) as usize)
356            .unwrap_or(0)
357    }
358
359    async fn clear(&self) -> Result<(), VectorStoreError> {
360        let collection_name = self.config.collection_name.clone();
361
362        self.client
363            .delete_collection(&collection_name)
364            .await
365            .map_err(|e| VectorStoreError::StorageError(format!("删除集合失败: {}", e)))?;
366
367        self.client
368            .create_collection(
369                CreateCollectionBuilder::new(&collection_name).vectors_config(
370                    VectorParamsBuilder::new(
371                        self.config.vector_size as u64,
372                        Distance::from(self.config.distance),
373                    ),
374                ),
375            )
376            .await
377            .map_err(|e| VectorStoreError::StorageError(format!("重建集合失败: {}", e)))?;
378
379        Ok(())
380    }
381}
382
383#[cfg(test)]
384mod tests {
385    use super::*;
386
387    #[test]
388    fn test_config_default() {
389        let config = QdrantConfig::default();
390        assert_eq!(config.url, "http://localhost:6334");
391        assert_eq!(config.collection_name, "langchainrust");
392        assert_eq!(config.vector_size, 1536);
393    }
394
395    #[test]
396    fn test_config_builder() {
397        let config = QdrantConfig::new("http://custom:6334", "test_collection")
398            .with_vector_size(3072)
399            .with_distance(QdrantDistance::Euclid);
400
401        assert_eq!(config.url, "http://custom:6334");
402        assert_eq!(config.collection_name, "test_collection");
403        assert_eq!(config.vector_size, 3072);
404        assert!(matches!(config.distance, QdrantDistance::Euclid));
405    }
406
407    #[tokio::test]
408    #[ignore = "需要 Qdrant 服务运行"]
409    async fn test_qdrant_integration() {
410        let config =
411            QdrantConfig::new("http://localhost:6334", "test_collection").with_vector_size(3);
412
413        let store = QdrantVectorStore::new(config).await.unwrap();
414
415        let docs = vec![Document::new("Document 1"), Document::new("Document 2")];
416        let embeddings = vec![vec![1.0, 0.0, 0.0], vec![0.0, 1.0, 0.0]];
417
418        let ids = store.add_documents(docs, embeddings).await.unwrap();
419        assert_eq!(ids.len(), 2);
420
421        let results = store.similarity_search(&[0.9, 0.1, 0.0], 2).await.unwrap();
422        assert_eq!(results.len(), 2);
423
424        store.clear().await.unwrap();
425    }
426}