mod db;
mod similarity;
pub use db::{
BatchInsertEmbeddingsStruct, CacheDB, Collection, CollectionHandlerStruct,
CreateCollectionStruct, Distance, Embedding, Error, GetSimilarityStruct, InsertEmbeddingStruct,
SimilarityResult,
};
pub use similarity::{ScoreIndex, get_cache_attr, get_distance_fn, normalize};
pub fn create_database() -> CacheDB {
CacheDB::new()
}
pub fn add(left: u64, right: u64) -> u64 {
left + right
}
#[cfg(test)]
mod tests {
use super::*;
use std::collections::HashMap;
#[test]
fn test_add_function() {
assert_eq!(add(2, 2), 4);
assert_eq!(add(0, 0), 0);
assert_eq!(add(100, 200), 300);
}
#[test]
fn test_library_exports() {
let db = CacheDB::new();
assert!(db.collections.is_empty());
let _euclidean = Distance::Euclidean;
let _cosine = Distance::Cosine;
let _dot_product = Distance::DotProduct;
let _create_struct = CreateCollectionStruct {
collection_name: "test".to_string(),
dimension: 128,
distance: Distance::Euclidean,
};
}
#[test]
fn test_create_database_function() {
let db = create_database();
assert!(db.collections.is_empty());
let db2 = CacheDB::new();
assert_eq!(db.collections.len(), db2.collections.len());
}
#[test]
fn test_embedding_creation() {
let mut id = HashMap::new();
id.insert("id".to_string(), "test_embedding".to_string());
let mut metadata = HashMap::new();
metadata.insert("type".to_string(), "document".to_string());
metadata.insert("source".to_string(), "test".to_string());
let embedding = Embedding {
id,
vector: vec![0.1, 0.2, 0.3, 0.4],
metadata: Some(metadata),
};
assert_eq!(embedding.vector.len(), 4);
assert!(embedding.metadata.is_some());
assert_eq!(embedding.id.get("id"), Some(&"test_embedding".to_string()));
}
#[test]
fn test_distance_enum_equality() {
assert_eq!(Distance::Euclidean, Distance::Euclidean);
assert_eq!(Distance::Cosine, Distance::Cosine);
assert_eq!(Distance::DotProduct, Distance::DotProduct);
assert_ne!(Distance::Euclidean, Distance::Cosine);
}
#[test]
fn test_similarity_functions_accessible() {
let vector = vec![1.0, 2.0, 3.0];
let euclidean_cache = get_cache_attr(Distance::Euclidean, &vector);
let cosine_cache = get_cache_attr(Distance::Cosine, &vector);
let dot_cache = get_cache_attr(Distance::DotProduct, &vector);
assert_eq!(euclidean_cache, 0.0);
assert_eq!(dot_cache, 0.0);
assert!(cosine_cache > 0.0);
let normalized = normalize(&vector);
assert_eq!(normalized.len(), vector.len());
let magnitude: f32 = normalized.iter().map(|x| x.powi(2)).sum::<f32>().sqrt();
assert!((magnitude - 1.0).abs() < 1e-6);
}
#[test]
fn test_score_index_ordering() {
let score1 = ScoreIndex {
score: 0.5,
index: 0,
};
let score2 = ScoreIndex {
score: 0.8,
index: 1,
};
let score3 = ScoreIndex {
score: 0.3,
index: 2,
};
assert!(score3 > score1); assert!(score1 > score2); }
#[test]
fn test_error_enum() {
let unique_error = Error::UniqueViolation;
let not_found_error = Error::NotFound;
let dimension_error = Error::DimensionMismatch;
assert_eq!(unique_error, Error::UniqueViolation);
assert_ne!(unique_error, not_found_error);
let _error_string = format!("{:?}", dimension_error);
}
#[test]
fn test_collection_creation_and_basic_operations() {
let mut db = CacheDB::new();
let result = db.create_collection("test_collection".to_string(), 3, Distance::Euclidean);
assert!(result.is_ok());
let collection = db.get_collection("test_collection");
assert!(collection.is_some());
assert_eq!(collection.unwrap().dimension, 3);
}
#[test]
fn test_end_to_end_workflow() {
let mut db = CacheDB::new();
db.create_collection("documents".to_string(), 4, Distance::Cosine)
.unwrap();
let mut id1 = HashMap::new();
id1.insert("doc_id".to_string(), "doc1".to_string());
let mut metadata1 = HashMap::new();
metadata1.insert("title".to_string(), "Document 1".to_string());
let embedding1 = Embedding {
id: id1,
vector: vec![1.0, 0.0, 0.0, 0.0],
metadata: Some(metadata1),
};
let mut id2 = HashMap::new();
id2.insert("doc_id".to_string(), "doc2".to_string());
let embedding2 = Embedding {
id: id2,
vector: vec![0.0, 1.0, 0.0, 0.0],
metadata: None,
};
assert!(db.insert_into_collection("documents", embedding1).is_ok());
assert!(db.insert_into_collection("documents", embedding2).is_ok());
let embeddings = db.get_embeddings("documents");
assert!(embeddings.is_some());
assert_eq!(embeddings.unwrap().len(), 2);
let collection = db.get_collection("documents").unwrap();
let query_vector = vec![1.0, 0.1, 0.0, 0.0];
let results = collection.get_similarity(&query_vector, 2);
assert_eq!(results.len(), 2);
assert!(results[0].score >= results[1].score);
}
#[test]
fn test_normalize_function_edge_cases() {
let zero_vec = vec![0.0, 0.0, 0.0];
let normalized_zero = normalize(&zero_vec);
assert_eq!(normalized_zero, zero_vec);
let small_vec = vec![1e-10, 1e-10, 1e-10];
let normalized_small = normalize(&small_vec);
assert_eq!(normalized_small, small_vec);
let unit_vec = vec![1.0, 0.0, 0.0];
let normalized_unit = normalize(&unit_vec);
let magnitude: f32 = normalized_unit
.iter()
.map(|x| x.powi(2))
.sum::<f32>()
.sqrt();
assert!((magnitude - 1.0).abs() < 1e-6);
}
#[test]
fn test_distance_functions() {
let vec1 = vec![1.0, 2.0, 3.0];
let vec2 = vec![4.0, 5.0, 6.0];
let euclidean_fn = get_distance_fn(Distance::Euclidean);
let euclidean_dist = euclidean_fn(&vec1, &vec2, 0.0);
assert!(euclidean_dist > 0.0);
let dot_fn = get_distance_fn(Distance::DotProduct);
let dot_result = dot_fn(&vec1, &vec2, 0.0);
assert_eq!(dot_result, 32.0);
let cosine_fn = get_distance_fn(Distance::Cosine);
let cosine_result = cosine_fn(&vec1, &vec2, 0.0);
assert_eq!(cosine_result, dot_result); }
}