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