1use 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#[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#[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
75pub 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 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()); 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}