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,
22 pub collection_name: String,
24 pub vector_size: usize,
26 pub distance: QdrantDistance,
28}
29
30#[derive(Debug, Clone, Copy)]
32pub enum QdrantDistance {
33 Cosine,
35 Euclid,
37 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 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 pub fn with_vector_size(mut self, size: usize) -> Self {
74 self.vector_size = size;
75 self
76 }
77
78 pub fn with_distance(mut self, distance: QdrantDistance) -> Self {
80 self.distance = distance;
81 self
82 }
83}
84
85pub struct QdrantVectorStore {
87 client: Arc<Qdrant>,
88 config: QdrantConfig,
89}
90
91impl QdrantVectorStore {
92 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 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 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 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 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()); 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}