1use ares_types::types::{AppError, Document, Result, SearchResult};
44use async_trait::async_trait;
45use serde::{Deserialize, Serialize};
46
47#[derive(Debug, Clone, Serialize, Deserialize)]
56#[serde(tag = "provider", rename_all = "lowercase")]
57pub enum VectorStoreProvider {
58 #[cfg(feature = "ares-vector")]
63 AresVector {
64 path: Option<String>,
66 },
67
68 #[cfg(feature = "lancedb")]
73 LanceDB {
74 path: String,
76 },
77
78 #[cfg(feature = "qdrant")]
82 Qdrant {
83 url: String,
85 api_key: Option<String>,
87 },
88
89 #[cfg(feature = "pgvector")]
93 PgVector {
94 connection_string: String,
96 },
97
98 #[cfg(feature = "chromadb")]
102 ChromaDB {
103 url: String,
105 },
106
107 #[cfg(feature = "pinecone")]
111 Pinecone {
112 api_key: String,
114 environment: String,
116 index_name: String,
118 },
119
120 InMemory,
124}
125
126impl VectorStoreProvider {
127 pub async fn create_store(&self) -> Result<Box<dyn VectorStore>> {
134 match self {
135 #[cfg(feature = "ares-vector")]
136 VectorStoreProvider::AresVector { path } => {
137 let store = super::ares_vector::AresVectorStore::new(path.clone()).await?;
138 Ok(Box::new(store))
139 }
140
141 #[cfg(feature = "lancedb")]
142 VectorStoreProvider::LanceDB { path } => {
143 let store = super::lancedb::LanceDBStore::new(path).await?;
144 Ok(Box::new(store))
145 }
146
147 #[cfg(feature = "qdrant")]
148 VectorStoreProvider::Qdrant { url, api_key } => {
149 let store =
150 super::qdrant::QdrantVectorStore::new(url.clone(), api_key.clone()).await?;
151 Ok(Box::new(store))
152 }
153
154 #[cfg(feature = "pgvector")]
155 VectorStoreProvider::PgVector { connection_string } => {
156 let store = super::pgvector::PgVectorStore::new(connection_string).await?;
157 Ok(Box::new(store))
158 }
159
160 #[cfg(feature = "chromadb")]
161 VectorStoreProvider::ChromaDB { url } => {
162 let store = super::chromadb::ChromaDBStore::new(url).await?;
163 Ok(Box::new(store))
164 }
165
166 #[cfg(feature = "pinecone")]
167 VectorStoreProvider::Pinecone {
168 api_key,
169 environment,
170 index_name,
171 } => {
172 let store =
173 super::pinecone::PineconeStore::new(api_key, environment, index_name).await?;
174 Ok(Box::new(store))
175 }
176
177 VectorStoreProvider::InMemory => {
178 let store = InMemoryVectorStore::new();
179 Ok(Box::new(store))
180 }
181
182 #[allow(unreachable_patterns)]
183 _ => Err(AppError::Configuration(
184 "Vector store provider not enabled. Check feature flags.".into(),
185 )),
186 }
187 }
188
189 pub fn from_env() -> Self {
200 #[cfg(feature = "ares-vector")]
201 if let Ok(path) = std::env::var("ARES_VECTOR_PATH") {
202 return VectorStoreProvider::AresVector { path: Some(path) };
203 }
204
205 #[cfg(feature = "lancedb")]
206 if let Ok(path) = std::env::var("LANCEDB_PATH") {
207 return VectorStoreProvider::LanceDB { path };
208 }
209
210 #[cfg(feature = "qdrant")]
211 if let Ok(url) = std::env::var("QDRANT_URL") {
212 let api_key = std::env::var("QDRANT_API_KEY").ok();
213 return VectorStoreProvider::Qdrant { url, api_key };
214 }
215
216 #[cfg(feature = "pgvector")]
217 if let Ok(connection_string) = std::env::var("PGVECTOR_URL") {
218 return VectorStoreProvider::PgVector { connection_string };
219 }
220
221 #[cfg(feature = "chromadb")]
222 if let Ok(url) = std::env::var("CHROMADB_URL") {
223 return VectorStoreProvider::ChromaDB { url };
224 }
225
226 #[cfg(feature = "pinecone")]
227 if let Ok(api_key) = std::env::var("PINECONE_API_KEY") {
228 let environment =
229 std::env::var("PINECONE_ENVIRONMENT").unwrap_or_else(|_| "us-east-1".into());
230 let index_name =
231 std::env::var("PINECONE_INDEX").unwrap_or_else(|_| "ares-documents".into());
232 return VectorStoreProvider::Pinecone {
233 api_key,
234 environment,
235 index_name,
236 };
237 }
238
239 #[cfg(feature = "ares-vector")]
241 return VectorStoreProvider::AresVector { path: None };
242
243 #[cfg(not(feature = "ares-vector"))]
244 VectorStoreProvider::InMemory
245 }
246}
247
248#[derive(Debug, Clone, Serialize, Deserialize)]
254pub struct CollectionStats {
255 pub name: String,
257 pub document_count: usize,
259 pub dimensions: usize,
261 pub index_size_bytes: Option<u64>,
263 pub distance_metric: String,
265}
266
267#[derive(Debug, Clone, Serialize, Deserialize)]
269pub struct CollectionInfo {
270 pub name: String,
272 pub document_count: usize,
274 pub dimensions: usize,
276}
277
278#[async_trait]
296pub trait VectorStore: Send + Sync {
297 fn provider_name(&self) -> &'static str;
299
300 async fn create_collection(&self, name: &str, dimensions: usize) -> Result<()>;
311
312 async fn delete_collection(&self, name: &str) -> Result<()>;
322
323 async fn list_collections(&self) -> Result<Vec<CollectionInfo>>;
325
326 async fn collection_exists(&self, name: &str) -> Result<bool>;
328
329 async fn collection_stats(&self, name: &str) -> Result<CollectionStats>;
331
332 async fn upsert(&self, collection: &str, documents: &[Document]) -> Result<usize>;
347
348 async fn search(
361 &self,
362 collection: &str,
363 embedding: &[f32],
364 limit: usize,
365 threshold: f32,
366 ) -> Result<Vec<SearchResult>>;
367
368 async fn search_with_filters(
382 &self,
383 collection: &str,
384 embedding: &[f32],
385 limit: usize,
386 threshold: f32,
387 _filters: &[(String, String)],
388 ) -> Result<Vec<SearchResult>> {
389 self.search(collection, embedding, limit, threshold).await
392 }
393
394 async fn delete(&self, collection: &str, ids: &[String]) -> Result<usize>;
405
406 async fn get(&self, collection: &str, id: &str) -> Result<Option<Document>>;
417
418 async fn count(&self, collection: &str) -> Result<usize> {
420 let stats = self.collection_stats(collection).await?;
421 Ok(stats.document_count)
422 }
423}
424
425use parking_lot::RwLock;
430use std::collections::HashMap;
431use std::sync::Arc;
432
433pub struct InMemoryVectorStore {
438 collections: Arc<RwLock<HashMap<String, InMemoryCollection>>>,
439}
440
441struct InMemoryCollection {
442 dimensions: usize,
443 documents: HashMap<String, Document>,
444}
445
446impl InMemoryVectorStore {
447 pub fn new() -> Self {
449 Self {
450 collections: Arc::new(RwLock::new(HashMap::new())),
451 }
452 }
453
454 fn cosine_similarity(a: &[f32], b: &[f32]) -> f32 {
456 if a.len() != b.len() {
457 return 0.0;
458 }
459
460 let dot_product: f32 = a.iter().zip(b.iter()).map(|(x, y)| x * y).sum();
461 let norm_a: f32 = a.iter().map(|x| x * x).sum::<f32>().sqrt();
462 let norm_b: f32 = b.iter().map(|x| x * x).sum::<f32>().sqrt();
463
464 if norm_a == 0.0 || norm_b == 0.0 {
465 return 0.0;
466 }
467
468 dot_product / (norm_a * norm_b)
469 }
470}
471
472impl Default for InMemoryVectorStore {
473 fn default() -> Self {
474 Self::new()
475 }
476}
477
478#[async_trait]
479impl VectorStore for InMemoryVectorStore {
480 fn provider_name(&self) -> &'static str {
481 "in-memory"
482 }
483
484 async fn create_collection(&self, name: &str, dimensions: usize) -> Result<()> {
485 let mut collections = self.collections.write();
486 if collections.contains_key(name) {
487 return Err(AppError::InvalidInput(format!(
488 "Collection '{}' already exists",
489 name
490 )));
491 }
492 collections.insert(
493 name.to_string(),
494 InMemoryCollection {
495 dimensions,
496 documents: HashMap::new(),
497 },
498 );
499 Ok(())
500 }
501
502 async fn delete_collection(&self, name: &str) -> Result<()> {
503 let mut collections = self.collections.write();
504 collections
505 .remove(name)
506 .ok_or_else(|| AppError::NotFound(format!("Collection '{}' not found", name)))?;
507 Ok(())
508 }
509
510 async fn list_collections(&self) -> Result<Vec<CollectionInfo>> {
511 let collections = self.collections.read();
512 Ok(collections
513 .iter()
514 .map(|(name, col)| CollectionInfo {
515 name: name.clone(),
516 document_count: col.documents.len(),
517 dimensions: col.dimensions,
518 })
519 .collect())
520 }
521
522 async fn collection_exists(&self, name: &str) -> Result<bool> {
523 let collections = self.collections.read();
524 Ok(collections.contains_key(name))
525 }
526
527 async fn collection_stats(&self, name: &str) -> Result<CollectionStats> {
528 let collections = self.collections.read();
529 let col = collections
530 .get(name)
531 .ok_or_else(|| AppError::NotFound(format!("Collection '{}' not found", name)))?;
532
533 Ok(CollectionStats {
534 name: name.to_string(),
535 document_count: col.documents.len(),
536 dimensions: col.dimensions,
537 index_size_bytes: None,
538 distance_metric: "cosine".to_string(),
539 })
540 }
541
542 async fn upsert(&self, collection: &str, documents: &[Document]) -> Result<usize> {
543 let mut collections = self.collections.write();
544 let col = collections
545 .get_mut(collection)
546 .ok_or_else(|| AppError::NotFound(format!("Collection '{}' not found", collection)))?;
547
548 let mut count = 0;
549 for doc in documents {
550 if doc.embedding.is_none() {
551 return Err(AppError::InvalidInput(format!(
552 "Document '{}' is missing embedding",
553 doc.id
554 )));
555 }
556 col.documents.insert(doc.id.clone(), doc.clone());
557 count += 1;
558 }
559
560 Ok(count)
561 }
562
563 async fn search(
564 &self,
565 collection: &str,
566 embedding: &[f32],
567 limit: usize,
568 threshold: f32,
569 ) -> Result<Vec<SearchResult>> {
570 let collections = self.collections.read();
571 let col = collections
572 .get(collection)
573 .ok_or_else(|| AppError::NotFound(format!("Collection '{}' not found", collection)))?;
574
575 let mut results: Vec<SearchResult> = col
576 .documents
577 .values()
578 .filter_map(|doc| {
579 let doc_embedding = doc.embedding.as_ref()?;
580 let score = Self::cosine_similarity(embedding, doc_embedding);
581 if score >= threshold {
582 Some(SearchResult {
583 document: Document {
584 id: doc.id.clone(),
585 content: doc.content.clone(),
586 metadata: doc.metadata.clone(),
587 embedding: None, },
589 score,
590 })
591 } else {
592 None
593 }
594 })
595 .collect();
596
597 results.sort_by(|a, b| {
599 b.score
600 .partial_cmp(&a.score)
601 .unwrap_or(std::cmp::Ordering::Equal)
602 });
603
604 results.truncate(limit);
606
607 Ok(results)
608 }
609
610 async fn delete(&self, collection: &str, ids: &[String]) -> Result<usize> {
611 let mut collections = self.collections.write();
612 let col = collections
613 .get_mut(collection)
614 .ok_or_else(|| AppError::NotFound(format!("Collection '{}' not found", collection)))?;
615
616 let mut count = 0;
617 for id in ids {
618 if col.documents.remove(id).is_some() {
619 count += 1;
620 }
621 }
622
623 Ok(count)
624 }
625
626 async fn get(&self, collection: &str, id: &str) -> Result<Option<Document>> {
627 let collections = self.collections.read();
628 let col = collections
629 .get(collection)
630 .ok_or_else(|| AppError::NotFound(format!("Collection '{}' not found", collection)))?;
631
632 Ok(col.documents.get(id).cloned())
633 }
634}
635
636#[cfg(test)]
641mod tests {
642 use super::*;
643 use ares_types::types::DocumentMetadata;
644 use chrono::Utc;
645
646 fn clear_vector_env() {
647 std::env::remove_var("ARES_VECTOR_PATH");
648 std::env::remove_var("LANCEDB_PATH");
649 std::env::remove_var("QDRANT_URL");
650 std::env::remove_var("QDRANT_API_KEY");
651 std::env::remove_var("PGVECTOR_URL");
652 std::env::remove_var("CHROMADB_URL");
653 std::env::remove_var("PINECONE_API_KEY");
654 std::env::remove_var("PINECONE_ENVIRONMENT");
655 std::env::remove_var("PINECONE_INDEX");
656 }
657
658 fn create_test_document(id: &str, content: &str, embedding: Vec<f32>) -> Document {
659 Document {
660 id: id.to_string(),
661 content: content.to_string(),
662 metadata: DocumentMetadata {
663 title: format!("Test Doc {}", id),
664 source: "test".to_string(),
665 created_at: Utc::now(),
666 tags: vec!["test".to_string()],
667 },
668 embedding: Some(embedding),
669 }
670 }
671
672 #[tokio::test]
675 async fn test_inmemory_create_collection() {
676 let store = InMemoryVectorStore::new();
677 store.create_collection("test", 384).await.unwrap();
678 assert!(store.collection_exists("test").await.unwrap());
679 }
680
681 #[tokio::test]
682 async fn test_inmemory_duplicate_collection_error() {
683 let store = InMemoryVectorStore::new();
684 store.create_collection("test", 384).await.unwrap();
685 let result = store.create_collection("test", 384).await;
686 assert!(result.is_err());
687 }
688
689 #[tokio::test]
690 async fn test_inmemory_upsert_and_search() {
691 let store = InMemoryVectorStore::new();
692 store.create_collection("test", 3).await.unwrap();
693
694 let doc1 = create_test_document("doc1", "Hello world", vec![1.0, 0.0, 0.0]);
695 let doc2 = create_test_document("doc2", "Goodbye world", vec![0.0, 1.0, 0.0]);
696 let doc3 = create_test_document("doc3", "Hello again", vec![0.9, 0.1, 0.0]);
697
698 store.upsert("test", &[doc1, doc2, doc3]).await.unwrap();
699
700 let results = store
701 .search("test", &[1.0, 0.0, 0.0], 10, 0.5)
702 .await
703 .unwrap();
704
705 assert_eq!(results.len(), 2);
706 assert_eq!(results[0].document.id, "doc1");
707 assert_eq!(results[1].document.id, "doc3");
708 }
709
710 #[tokio::test]
711 async fn test_inmemory_delete() {
712 let store = InMemoryVectorStore::new();
713 store.create_collection("test", 3).await.unwrap();
714
715 let doc = create_test_document("doc1", "Test", vec![1.0, 0.0, 0.0]);
716 store.upsert("test", &[doc]).await.unwrap();
717 assert_eq!(store.count("test").await.unwrap(), 1);
718
719 let deleted = store.delete("test", &["doc1".to_string()]).await.unwrap();
720 assert_eq!(deleted, 1);
721 assert_eq!(store.count("test").await.unwrap(), 0);
722 }
723
724 #[tokio::test]
725 async fn test_inmemory_get() {
726 let store = InMemoryVectorStore::new();
727 store.create_collection("test", 3).await.unwrap();
728
729 let doc = create_test_document("doc1", "Test content", vec![1.0, 0.0, 0.0]);
730 store.upsert("test", &[doc]).await.unwrap();
731
732 let retrieved = store.get("test", "doc1").await.unwrap();
733 assert!(retrieved.is_some());
734 assert_eq!(retrieved.unwrap().content, "Test content");
735
736 let not_found = store.get("test", "nonexistent").await.unwrap();
737 assert!(not_found.is_none());
738 }
739
740 #[tokio::test]
741 async fn test_inmemory_list_collections() {
742 let store = InMemoryVectorStore::new();
743 store.create_collection("col1", 384).await.unwrap();
744 store.create_collection("col2", 768).await.unwrap();
745 let collections = store.list_collections().await.unwrap();
746 assert_eq!(collections.len(), 2);
747 }
748
749 #[tokio::test]
750 async fn test_cosine_similarity() {
751 assert!(
752 (InMemoryVectorStore::cosine_similarity(&[1.0, 0.0], &[1.0, 0.0]) - 1.0).abs() < 0.001
753 );
754 assert!(InMemoryVectorStore::cosine_similarity(&[1.0, 0.0], &[0.0, 1.0]).abs() < 0.001);
755 assert!(
756 (InMemoryVectorStore::cosine_similarity(&[1.0, 0.0], &[-1.0, 0.0]) + 1.0).abs() < 0.001
757 );
758 }
759
760 #[test]
765 fn test_cosine_similarity_mismatched_lengths_returns_zero() {
766 let score = InMemoryVectorStore::cosine_similarity(&[1.0, 0.0], &[1.0, 0.0, 0.0]);
767 assert_eq!(score, 0.0);
768 }
769
770 #[test]
771 fn test_cosine_similarity_empty_vectors_returns_zero() {
772 let score = InMemoryVectorStore::cosine_similarity(&[], &[]);
773 assert_eq!(score, 0.0);
774 }
775
776 #[test]
777 fn test_cosine_similarity_first_zero_vector_returns_zero() {
778 let score = InMemoryVectorStore::cosine_similarity(&[0.0, 0.0], &[1.0, 1.0]);
779 assert_eq!(score, 0.0);
780 }
781
782 #[test]
783 fn test_cosine_similarity_second_zero_vector_returns_zero() {
784 let score = InMemoryVectorStore::cosine_similarity(&[1.0, 1.0], &[0.0, 0.0]);
785 assert_eq!(score, 0.0);
786 }
787
788 #[test]
789 fn test_cosine_similarity_high_dimensional() {
790 let a: Vec<f32> = (0..384).map(|i| i as f32).collect();
791 let b: Vec<f32> = (0..384).map(|i| i as f32).collect();
792 let score = InMemoryVectorStore::cosine_similarity(&a, &b);
793 assert!((score - 1.0).abs() < 0.001);
794 }
795
796 #[test]
797 fn test_cosine_similarity_negative_values() {
798 let score = InMemoryVectorStore::cosine_similarity(&[-1.0, -2.0], &[-3.0, -6.0]);
799 assert!((score - 1.0).abs() < 0.001);
800 }
801
802 #[test]
803 fn test_cosine_similarity_partial_orthogonality() {
804 let score = InMemoryVectorStore::cosine_similarity(&[1.0, 1.0, 0.0], &[0.0, 0.0, 1.0]);
805 assert!(score.abs() < 0.001);
806 }
807
808 #[test]
809 fn test_cosine_similarity_symmetry() {
810 let a = [1.0, 2.0, 3.0];
811 let b = [4.0, 5.0, 6.0];
812 let score_ab = InMemoryVectorStore::cosine_similarity(&a, &b);
813 let score_ba = InMemoryVectorStore::cosine_similarity(&b, &a);
814 assert!((score_ab - score_ba).abs() < 0.001);
815 }
816
817 #[tokio::test]
820 async fn test_inmemory_default_is_empty() {
821 let store = InMemoryVectorStore::default();
822 let collections = store.list_collections().await.unwrap();
823 assert!(collections.is_empty());
824 }
825
826 #[test]
829 fn test_vector_store_provider_inmemory_serde_roundtrip() {
830 let provider = VectorStoreProvider::InMemory;
831 let json = serde_json::to_string(&provider).unwrap();
832 let deserialized: VectorStoreProvider = serde_json::from_str(&json).unwrap();
833 assert!(matches!(deserialized, VectorStoreProvider::InMemory));
834 }
835
836 #[test]
837 fn test_vector_store_provider_inmemory_json_value() {
838 let provider = VectorStoreProvider::InMemory;
839 let json = serde_json::to_string(&provider).unwrap();
840 assert_eq!(json, r#"{"provider":"inmemory"}"#);
841 }
842
843 #[test]
844 fn test_vector_store_provider_deserialize_from_json() {
845 let json = r#"{"provider":"inmemory"}"#;
846 let provider: VectorStoreProvider = serde_json::from_str(json).unwrap();
847 assert!(matches!(provider, VectorStoreProvider::InMemory));
848 }
849
850 #[test]
851 fn test_collection_stats_serde_roundtrip() {
852 let stats = CollectionStats {
853 name: "test-col".to_string(),
854 document_count: 42,
855 dimensions: 768,
856 index_size_bytes: Some(1024 * 1024),
857 distance_metric: "cosine".to_string(),
858 };
859 let json = serde_json::to_string(&stats).unwrap();
860 let deserialized: CollectionStats = serde_json::from_str(&json).unwrap();
861 assert_eq!(deserialized.name, "test-col");
862 assert_eq!(deserialized.document_count, 42);
863 assert_eq!(deserialized.dimensions, 768);
864 assert_eq!(deserialized.index_size_bytes, Some(1024 * 1024));
865 assert_eq!(deserialized.distance_metric, "cosine");
866 }
867
868 #[test]
869 fn test_collection_stats_none_index_size() {
870 let stats = CollectionStats {
871 name: "col".to_string(),
872 document_count: 0,
873 dimensions: 384,
874 index_size_bytes: None,
875 distance_metric: "euclidean".to_string(),
876 };
877 let json = serde_json::to_string(&stats).unwrap();
878 let deserialized: CollectionStats = serde_json::from_str(&json).unwrap();
879 assert_eq!(deserialized.index_size_bytes, None);
880 }
881
882 #[test]
883 fn test_collection_info_serde_roundtrip() {
884 let info = CollectionInfo {
885 name: "docs".to_string(),
886 document_count: 100,
887 dimensions: 512,
888 };
889 let json = serde_json::to_string(&info).unwrap();
890 let deserialized: CollectionInfo = serde_json::from_str(&json).unwrap();
891 assert_eq!(deserialized.name, "docs");
892 assert_eq!(deserialized.document_count, 100);
893 assert_eq!(deserialized.dimensions, 512);
894 }
895
896 #[tokio::test]
899 async fn test_inmemory_provider_create_store() {
900 let provider = VectorStoreProvider::InMemory;
901 let store = provider.create_store().await.unwrap();
902 assert_eq!(store.provider_name(), "in-memory");
903 }
904
905 #[tokio::test]
906 async fn test_collection_exists_false_when_missing() {
907 let store = InMemoryVectorStore::new();
908 assert!(!store.collection_exists("missing").await.unwrap());
909 }
910
911 #[tokio::test]
912 async fn test_search_with_filters_default_matches_search() {
913 let store = InMemoryVectorStore::new();
914 store.create_collection("test", 3).await.unwrap();
915 let doc = create_test_document("d1", "x", vec![1.0, 0.0, 0.0]);
916 store.upsert("test", &[doc]).await.unwrap();
917
918 let plain = store
919 .search("test", &[1.0, 0.0, 0.0], 5, 0.0)
920 .await
921 .unwrap();
922 let filtered = store
923 .search_with_filters(
924 "test",
925 &[1.0, 0.0, 0.0],
926 5,
927 0.0,
928 &[("tag".to_string(), "test".to_string())],
929 )
930 .await
931 .unwrap();
932 assert_eq!(plain.len(), filtered.len());
933 assert_eq!(plain[0].document.id, filtered[0].document.id);
934 }
935
936 #[cfg(feature = "pgvector")]
937 #[test]
938 fn test_pgvector_provider_serde_roundtrip() {
939 let provider = VectorStoreProvider::PgVector {
940 connection_string: "postgres://localhost/ares".into(),
941 };
942 let json = serde_json::to_string(&provider).unwrap();
943 assert!(json.contains("pgvector"));
944 let restored: VectorStoreProvider = serde_json::from_str(&json).unwrap();
945 match (provider, restored) {
946 (
947 VectorStoreProvider::PgVector {
948 connection_string: a,
949 },
950 VectorStoreProvider::PgVector {
951 connection_string: b,
952 },
953 ) => assert_eq!(a, b),
954 _ => panic!("pgvector roundtrip mismatch"),
955 }
956 }
957
958 #[cfg(feature = "lancedb")]
959 #[test]
960 fn test_lancedb_provider_serde_roundtrip() {
961 let provider = VectorStoreProvider::LanceDB {
962 path: "/data/lancedb".into(),
963 };
964 let json = serde_json::to_string(&provider).unwrap();
965 let restored: VectorStoreProvider = serde_json::from_str(&json).unwrap();
966 match (provider, restored) {
967 (
968 VectorStoreProvider::LanceDB { path: a },
969 VectorStoreProvider::LanceDB { path: b },
970 ) => assert_eq!(a, b),
971 _ => panic!("lancedb roundtrip mismatch"),
972 }
973 }
974
975 #[cfg(feature = "qdrant")]
976 #[test]
977 fn test_qdrant_provider_serde_roundtrip() {
978 let provider = VectorStoreProvider::Qdrant {
979 url: "http://localhost:6334".into(),
980 api_key: Some("secret".into()),
981 };
982 let json = serde_json::to_string(&provider).unwrap();
983 let restored: VectorStoreProvider = serde_json::from_str(&json).unwrap();
984 match (provider, restored) {
985 (
986 VectorStoreProvider::Qdrant {
987 url: a,
988 api_key: ka,
989 },
990 VectorStoreProvider::Qdrant {
991 url: b,
992 api_key: kb,
993 },
994 ) => {
995 assert_eq!(a, b);
996 assert_eq!(ka, kb);
997 }
998 _ => panic!("qdrant roundtrip mismatch"),
999 }
1000 }
1001
1002 #[cfg(feature = "pinecone")]
1003 #[test]
1004 fn test_pinecone_provider_serde_roundtrip() {
1005 let provider = VectorStoreProvider::Pinecone {
1006 api_key: "key".into(),
1007 environment: "us-east-1".into(),
1008 index_name: "idx".into(),
1009 };
1010 let json = serde_json::to_string(&provider).unwrap();
1011 let restored: VectorStoreProvider = serde_json::from_str(&json).unwrap();
1012 match (provider, restored) {
1013 (
1014 VectorStoreProvider::Pinecone {
1015 api_key: a,
1016 environment: e,
1017 index_name: i,
1018 },
1019 VectorStoreProvider::Pinecone {
1020 api_key: b,
1021 environment: f,
1022 index_name: j,
1023 },
1024 ) => {
1025 assert_eq!(a, b);
1026 assert_eq!(e, f);
1027 assert_eq!(i, j);
1028 }
1029 _ => panic!("pinecone roundtrip mismatch"),
1030 }
1031 }
1032
1033 #[cfg(feature = "pgvector")]
1034 #[test]
1035 fn test_from_env_pgvector_url() {
1036 clear_vector_env();
1037 std::env::set_var("PGVECTOR_URL", "postgres://localhost/vec");
1038 let provider = VectorStoreProvider::from_env();
1039 match provider {
1040 VectorStoreProvider::PgVector { connection_string } => {
1041 assert_eq!(connection_string, "postgres://localhost/vec");
1042 }
1043 other => panic!("expected PgVector, got {:?}", other),
1044 }
1045 std::env::remove_var("PGVECTOR_URL");
1046 }
1047
1048 #[cfg(feature = "lancedb")]
1049 #[test]
1050 fn test_from_env_lancedb_path() {
1051 clear_vector_env();
1052 std::env::set_var("LANCEDB_PATH", "/tmp/lance");
1053 let provider = VectorStoreProvider::from_env();
1054 match provider {
1055 VectorStoreProvider::LanceDB { path } => assert_eq!(path, "/tmp/lance"),
1056 other => panic!("expected LanceDB, got {:?}", other),
1057 }
1058 std::env::remove_var("LANCEDB_PATH");
1059 }
1060
1061 #[cfg(feature = "qdrant")]
1062 #[test]
1063 fn test_from_env_qdrant_url() {
1064 clear_vector_env();
1065 std::env::set_var("QDRANT_URL", "http://qdrant:6334");
1066 std::env::set_var("QDRANT_API_KEY", "token");
1067 let provider = VectorStoreProvider::from_env();
1068 match provider {
1069 VectorStoreProvider::Qdrant { url, api_key } => {
1070 assert_eq!(url, "http://qdrant:6334");
1071 assert_eq!(api_key.as_deref(), Some("token"));
1072 }
1073 other => panic!("expected Qdrant, got {:?}", other),
1074 }
1075 std::env::remove_var("QDRANT_URL");
1076 std::env::remove_var("QDRANT_API_KEY");
1077 }
1078
1079 #[cfg(feature = "pinecone")]
1080 #[test]
1081 fn test_from_env_pinecone_api_key() {
1082 clear_vector_env();
1083 std::env::set_var("PINECONE_API_KEY", "pk-test");
1084 std::env::set_var("PINECONE_ENVIRONMENT", "eu-west-1");
1085 std::env::set_var("PINECONE_INDEX", "my-index");
1086 let provider = VectorStoreProvider::from_env();
1087 match provider {
1088 VectorStoreProvider::Pinecone {
1089 api_key,
1090 environment,
1091 index_name,
1092 } => {
1093 assert_eq!(api_key, "pk-test");
1094 assert_eq!(environment, "eu-west-1");
1095 assert_eq!(index_name, "my-index");
1096 }
1097 other => panic!("expected Pinecone, got {:?}", other),
1098 }
1099 std::env::remove_var("PINECONE_API_KEY");
1100 std::env::remove_var("PINECONE_ENVIRONMENT");
1101 std::env::remove_var("PINECONE_INDEX");
1102 }
1103
1104 #[test]
1107 fn test_from_env_defaults_to_inmemory() {
1108 clear_vector_env();
1109
1110 let provider = VectorStoreProvider::from_env();
1111 #[cfg(not(feature = "ares-vector"))]
1112 assert!(matches!(provider, VectorStoreProvider::InMemory));
1113 #[cfg(feature = "ares-vector")]
1114 match provider {
1115 VectorStoreProvider::AresVector { path } => {
1116 assert_eq!(path, None, "AresVector default should have no path");
1117 }
1118 other => panic!("Unexpected default provider: {:?}", other),
1119 }
1120 }
1121
1122 #[tokio::test]
1125 async fn test_upsert_missing_embedding_errors() {
1126 let store = InMemoryVectorStore::new();
1127 store.create_collection("test", 3).await.unwrap();
1128
1129 let doc_no_embedding = Document {
1130 id: "bad-doc".to_string(),
1131 content: "No embedding".to_string(),
1132 metadata: DocumentMetadata::default(),
1133 embedding: None,
1134 };
1135 let result = store.upsert("test", &[doc_no_embedding]).await;
1136 assert!(result.is_err());
1137 match result.unwrap_err() {
1138 AppError::InvalidInput(msg) => assert!(msg.contains("missing embedding")),
1139 other => panic!("Expected InvalidInput, got {:?}", other),
1140 }
1141 }
1142
1143 #[tokio::test]
1144 async fn test_upsert_on_nonexistent_collection_errors() {
1145 let store = InMemoryVectorStore::new();
1146 let doc = create_test_document("d1", "content", vec![1.0]);
1147 let result = store.upsert("nope", &[doc]).await;
1148 assert!(result.is_err());
1149 match result.unwrap_err() {
1150 AppError::NotFound(_) => {}
1151 other => panic!("Expected NotFound, got {:?}", other),
1152 }
1153 }
1154
1155 #[tokio::test]
1156 async fn test_search_threshold_filters_results() {
1157 let store = InMemoryVectorStore::new();
1158 store.create_collection("test", 3).await.unwrap();
1159
1160 let doc1 = create_test_document("d1", "a", vec![1.0, 0.0, 0.0]);
1161 let doc2 = create_test_document("d2", "b", vec![0.5, 0.866, 0.0]);
1162 let doc3 = create_test_document("d3", "c", vec![0.0, 1.0, 0.0]);
1163
1164 store.upsert("test", &[doc1, doc2, doc3]).await.unwrap();
1165
1166 let results = store
1167 .search("test", &[1.0, 0.0, 0.0], 10, 0.8)
1168 .await
1169 .unwrap();
1170 assert!(results.iter().all(|r| r.score >= 0.8));
1171 assert!(!results.iter().any(|r| r.document.id == "d3"));
1172 }
1173
1174 #[tokio::test]
1175 async fn test_search_limit_truncation() {
1176 let store = InMemoryVectorStore::new();
1177 store.create_collection("test", 2).await.unwrap();
1178
1179 let docs: Vec<Document> = (0..10)
1180 .map(|i| create_test_document(&format!("d{}", i), "doc", vec![1.0, 0.0]))
1181 .collect();
1182 store.upsert("test", &docs).await.unwrap();
1183
1184 let results = store.search("test", &[1.0, 0.0], 3, 0.0).await.unwrap();
1185 assert_eq!(results.len(), 3);
1186 }
1187
1188 #[tokio::test]
1189 async fn test_delete_nonexistent_collection_errors() {
1190 let store = InMemoryVectorStore::new();
1191 let result = store.delete_collection("nope").await;
1192 assert!(result.is_err());
1193 match result.unwrap_err() {
1194 AppError::NotFound(_) => {}
1195 other => panic!("Expected NotFound, got {:?}", other),
1196 }
1197 }
1198
1199 #[tokio::test]
1200 async fn test_delete_mixed_existing_and_nonexistent_ids() {
1201 let store = InMemoryVectorStore::new();
1202 store.create_collection("test", 2).await.unwrap();
1203
1204 let doc = create_test_document("d1", "content", vec![1.0, 0.0]);
1205 store.upsert("test", &[doc]).await.unwrap();
1206
1207 let deleted = store
1208 .delete("test", &["d1".to_string(), "d2".to_string()])
1209 .await
1210 .unwrap();
1211 assert_eq!(deleted, 1);
1212 assert_eq!(store.count("test").await.unwrap(), 0);
1213 }
1214
1215 #[tokio::test]
1216 async fn test_delete_returns_zero_for_all_nonexistent_ids() {
1217 let store = InMemoryVectorStore::new();
1218 store.create_collection("test", 2).await.unwrap();
1219
1220 let deleted = store
1221 .delete("test", &["a".to_string(), "b".to_string()])
1222 .await
1223 .unwrap();
1224 assert_eq!(deleted, 0);
1225 }
1226
1227 #[tokio::test]
1228 async fn test_collection_stats_after_operations() {
1229 let store = InMemoryVectorStore::new();
1230 store.create_collection("test", 4).await.unwrap();
1231
1232 let doc = create_test_document("d1", "hello", vec![1.0, 0.0, 0.0, 0.0]);
1233 store.upsert("test", &[doc]).await.unwrap();
1234
1235 let stats = store.collection_stats("test").await.unwrap();
1236 assert_eq!(stats.document_count, 1);
1237 assert_eq!(stats.dimensions, 4);
1238 assert_eq!(stats.distance_metric, "cosine");
1239 assert_eq!(stats.index_size_bytes, None);
1240
1241 store.delete("test", &["d1".to_string()]).await.unwrap();
1242 let stats = store.collection_stats("test").await.unwrap();
1243 assert_eq!(stats.document_count, 0);
1244 }
1245
1246 #[tokio::test]
1247 async fn test_collection_stats_nonexistent_errors() {
1248 let store = InMemoryVectorStore::new();
1249 let result = store.collection_stats("nope").await;
1250 assert!(result.is_err());
1251 }
1252
1253 #[tokio::test]
1254 async fn test_list_collections_after_delete() {
1255 let store = InMemoryVectorStore::new();
1256 store.create_collection("a", 128).await.unwrap();
1257 store.create_collection("b", 256).await.unwrap();
1258 store.delete_collection("a").await.unwrap();
1259
1260 let cols = store.list_collections().await.unwrap();
1261 assert_eq!(cols.len(), 1);
1262 assert_eq!(cols[0].name, "b");
1263 }
1264
1265 #[tokio::test]
1266 async fn test_get_nonexistent_collection_errors() {
1267 let store = InMemoryVectorStore::new();
1268 let result = store.get("nope", "id").await;
1269 assert!(result.is_err());
1270 }
1271
1272 #[tokio::test]
1273 async fn test_provider_name() {
1274 let store = InMemoryVectorStore::new();
1275 assert_eq!(store.provider_name(), "in-memory");
1276 }
1277
1278 #[tokio::test]
1279 async fn test_search_empty_collection_returns_empty() {
1280 let store = InMemoryVectorStore::new();
1281 store.create_collection("test", 3).await.unwrap();
1282
1283 let results = store
1284 .search("test", &[1.0, 0.0, 0.0], 10, 0.0)
1285 .await
1286 .unwrap();
1287 assert!(results.is_empty());
1288 }
1289
1290 #[tokio::test]
1291 async fn test_search_results_excludes_embeddings() {
1292 let store = InMemoryVectorStore::new();
1293 store.create_collection("test", 3).await.unwrap();
1294
1295 let doc = create_test_document("d1", "content", vec![1.0, 0.0, 0.0]);
1296 store.upsert("test", &[doc]).await.unwrap();
1297
1298 let results = store
1299 .search("test", &[1.0, 0.0, 0.0], 10, 0.0)
1300 .await
1301 .unwrap();
1302 assert_eq!(results.len(), 1);
1303 assert!(results[0].document.embedding.is_none());
1304 }
1305
1306 #[tokio::test]
1307 async fn test_upsert_updates_existing_document() {
1308 let store = InMemoryVectorStore::new();
1309 store.create_collection("test", 3).await.unwrap();
1310
1311 let doc1 = create_test_document("d1", "original", vec![1.0, 0.0, 0.0]);
1312 store.upsert("test", &[doc1]).await.unwrap();
1313
1314 let doc1_updated = create_test_document("d1", "updated", vec![0.0, 1.0, 0.0]);
1315 store.upsert("test", &[doc1_updated]).await.unwrap();
1316
1317 let retrieved = store.get("test", "d1").await.unwrap().unwrap();
1318 assert_eq!(retrieved.content, "updated");
1319 assert_eq!(store.count("test").await.unwrap(), 1);
1320 }
1321
1322 #[tokio::test]
1323 async fn test_search_score_ordering_descending() {
1324 let store = InMemoryVectorStore::new();
1325 store.create_collection("test", 3).await.unwrap();
1326
1327 let doc1 = create_test_document("d1", "far", vec![0.0, 0.0, 1.0]);
1328 let doc2 = create_test_document("d2", "near", vec![0.9, 0.1, 0.0]);
1329 let doc3 = create_test_document("d3", "exact", vec![1.0, 0.0, 0.0]);
1330
1331 store.upsert("test", &[doc1, doc2, doc3]).await.unwrap();
1332
1333 let results = store
1334 .search("test", &[1.0, 0.0, 0.0], 10, 0.0)
1335 .await
1336 .unwrap();
1337
1338 for i in 1..results.len() {
1339 assert!(results[i - 1].score >= results[i].score);
1340 }
1341 assert_eq!(results[0].document.id, "d3");
1342 }
1343
1344 #[tokio::test]
1345 async fn test_search_nonexistent_collection_errors() {
1346 let store = InMemoryVectorStore::new();
1347 let result = store.search("nope", &[1.0, 0.0], 5, 0.0).await;
1348 assert!(result.is_err());
1349 match result.unwrap_err() {
1350 AppError::NotFound(msg) => assert!(msg.contains("nope")),
1351 other => panic!("Expected NotFound, got {:?}", other),
1352 }
1353 }
1354
1355 #[tokio::test]
1356 async fn test_delete_on_nonexistent_collection_errors() {
1357 let store = InMemoryVectorStore::new();
1358 let result = store.delete("nope", &["id".to_string()]).await;
1359 assert!(result.is_err());
1360 match result.unwrap_err() {
1361 AppError::NotFound(msg) => assert!(msg.contains("nope")),
1362 other => panic!("Expected NotFound, got {:?}", other),
1363 }
1364 }
1365
1366 #[tokio::test]
1367 async fn test_count_matches_collection_stats() {
1368 let store = InMemoryVectorStore::new();
1369 store.create_collection("test", 2).await.unwrap();
1370
1371 let doc = create_test_document("d1", "x", vec![1.0, 0.0]);
1372 store.upsert("test", &[doc]).await.unwrap();
1373
1374 assert_eq!(store.count("test").await.unwrap(), 1);
1375 let stats = store.collection_stats("test").await.unwrap();
1376 assert_eq!(stats.document_count, 1);
1377 assert_eq!(stats.name, "test");
1378 }
1379
1380 #[test]
1381 fn test_public_types_clone_and_debug() {
1382 let stats = CollectionStats {
1383 name: "c".to_string(),
1384 document_count: 1,
1385 dimensions: 3,
1386 index_size_bytes: None,
1387 distance_metric: "cosine".to_string(),
1388 };
1389 let stats_dbg = format!("{:?}", stats.clone());
1390 assert!(stats_dbg.contains("c"));
1391
1392 let info = CollectionInfo {
1393 name: "docs".to_string(),
1394 document_count: 2,
1395 dimensions: 128,
1396 };
1397 let info_dbg = format!("{:?}", info.clone());
1398 assert!(info_dbg.contains("docs"));
1399
1400 let provider = VectorStoreProvider::InMemory;
1401 let provider_dbg = format!("{:?}", provider.clone());
1402 assert!(provider_dbg.contains("InMemory"));
1403 }
1404
1405 #[tokio::test]
1406 async fn test_search_sort_handles_nan_scores() {
1407 let store = InMemoryVectorStore::new();
1408 store.create_collection("test", 2).await.unwrap();
1409
1410 let nan_doc = create_test_document("nan", "bad", vec![f32::NAN, 0.0]);
1411 let ok_doc = create_test_document("ok", "good", vec![1.0, 0.0]);
1412 store.upsert("test", &[nan_doc, ok_doc]).await.unwrap();
1413
1414 let results = store.search("test", &[1.0, 0.0], 10, 0.0).await.unwrap();
1415 assert!(!results.is_empty());
1416 }
1417}