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 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 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()); 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}