mod common;
use common::*;
use embellama::{CacheConfigBuilder, EmbeddingEngine, EngineConfigBuilder, NormalizationMode};
use std::time::Instant;
#[test]
#[ignore = "Requires model file"]
fn test_engine_with_cache_enabled() {
let Some(model_path) = get_test_model_path() else {
eprintln!("Skipping test: no test model configured");
return;
};
let config = EngineConfigBuilder::new()
.with_model_path(model_path)
.with_model_name("test-model")
.with_cache_config(
CacheConfigBuilder::new()
.with_enabled(true)
.with_embedding_cache_size(100)
.with_ttl_seconds(3600)
.build()
.expect("Failed to build cache config"),
)
.build()
.expect("Failed to build config");
let engine = EmbeddingEngine::new(config).expect("Failed to create engine");
assert!(engine.is_cache_enabled());
let text = "This is a test for caching functionality";
let start = Instant::now();
let embedding1 = engine
.embed(None, text)
.expect("Failed to generate embedding");
let time_uncached = start.elapsed();
let start = Instant::now();
let embedding2 = engine
.embed(None, text)
.expect("Failed to generate embedding");
let time_cached = start.elapsed();
assert_eq!(embedding1, embedding2);
assert!(
time_cached < time_uncached / 5,
"Cached call should be at least 5x faster. Uncached: {:?}, Cached: {:?}",
time_uncached,
time_cached
);
let stats = engine.get_cache_stats().expect("Cache should be enabled");
assert_eq!(stats.hits, 1);
assert_eq!(stats.misses, 1);
assert_eq!(stats.entry_count, 1);
assert!(stats.hit_rate > 0.0 && stats.hit_rate <= 1.0);
}
#[test]
#[ignore = "Requires model file"]
fn test_engine_with_cache_disabled() {
let Some(model_path) = get_test_model_path() else {
eprintln!("Skipping test: no test model configured");
return;
};
let config = EngineConfigBuilder::new()
.with_model_path(model_path)
.with_model_name("test-model")
.with_cache_disabled()
.build()
.expect("Failed to build config");
let engine = EmbeddingEngine::new(config).expect("Failed to create engine");
assert!(!engine.is_cache_enabled());
let text = "Test without caching";
let embedding1 = engine
.embed(None, text)
.expect("Failed to generate embedding");
let embedding2 = engine
.embed(None, text)
.expect("Failed to generate embedding");
assert_eq!(embedding1, embedding2);
assert!(engine.get_cache_stats().is_none());
}
#[test]
#[ignore = "Requires model file"]
fn test_cache_batch_processing() {
let Some(model_path) = get_test_model_path() else {
eprintln!("Skipping test: no test model configured");
return;
};
let config = EngineConfigBuilder::new()
.with_model_path(model_path)
.with_model_name("test-model")
.with_cache_enabled()
.build()
.expect("Failed to build config");
let engine = EmbeddingEngine::new(config).expect("Failed to create engine");
let texts = vec![
"First unique text",
"Second unique text",
"First unique text", "Third unique text",
"Second unique text", ];
let embeddings1 = engine
.embed_batch(None, &texts)
.expect("Failed to generate batch embeddings");
assert_eq!(embeddings1.len(), 5);
assert_eq!(embeddings1[0], embeddings1[2]); assert_eq!(embeddings1[1], embeddings1[4]);
let start = Instant::now();
let embeddings2 = engine
.embed_batch(None, &texts)
.expect("Failed to generate batch embeddings");
let cached_time = start.elapsed();
for (e1, e2) in embeddings1.iter().zip(embeddings2.iter()) {
assert_eq!(e1, e2);
}
let stats = engine.get_cache_stats().expect("Cache should be enabled");
assert_eq!(stats.entry_count, 3);
assert!(stats.hits >= 5);
println!("Batch cache lookup time: {:?}", cached_time);
}
#[test]
#[ignore = "Requires model file"]
fn test_cache_clearing() {
let Some(model_path) = get_test_model_path() else {
eprintln!("Skipping test: no test model configured");
return;
};
let config = EngineConfigBuilder::new()
.with_model_path(model_path)
.with_model_name("test-model")
.with_cache_enabled()
.build()
.expect("Failed to build config");
let engine = EmbeddingEngine::new(config).expect("Failed to create engine");
let texts = vec!["text1", "text2", "text3"];
for text in &texts {
engine
.embed(None, text)
.expect("Failed to generate embedding");
}
let stats_before = engine.get_cache_stats().expect("Cache should be enabled");
assert_eq!(stats_before.entry_count, 3);
engine.clear_cache();
let stats_after = engine.get_cache_stats().expect("Cache should be enabled");
assert_eq!(stats_after.entry_count, 0);
assert_eq!(stats_after.hits, 0);
assert_eq!(stats_after.misses, 0);
engine
.embed(None, texts[0])
.expect("Failed to generate embedding");
let stats_final = engine.get_cache_stats().expect("Cache should be enabled");
assert_eq!(stats_final.misses, 1);
assert_eq!(stats_final.hits, 0);
}
#[test]
#[ignore = "Requires model file"]
fn test_cache_warm_up() {
let Some(model_path) = get_test_model_path() else {
eprintln!("Skipping test: no test model configured");
return;
};
let config = EngineConfigBuilder::new()
.with_model_path(model_path)
.with_model_name("test-model")
.with_cache_enabled()
.build()
.expect("Failed to build config");
let engine = EmbeddingEngine::new(config).expect("Failed to create engine");
let warm_up_texts = vec![
"Frequently used text 1",
"Frequently used text 2",
"Frequently used text 3",
];
engine
.warm_cache(None, &warm_up_texts)
.expect("Failed to warm cache");
let stats = engine.get_cache_stats().expect("Cache should be enabled");
assert_eq!(stats.entry_count, 3);
for text in &warm_up_texts {
let start = Instant::now();
engine
.embed(None, text)
.expect("Failed to generate embedding");
let elapsed = start.elapsed();
assert!(
elapsed.as_millis() < 10,
"Cache hit took too long: {:?}",
elapsed
);
}
let final_stats = engine.get_cache_stats().expect("Cache should be enabled");
assert_eq!(final_stats.hits, 3);
}
#[test]
#[ignore = "Requires model file"]
fn test_cache_with_different_models() {
let Some(model_path) = get_test_model_path() else {
eprintln!("Skipping test: no test model configured");
return;
};
let config = EngineConfigBuilder::new()
.with_model_path(model_path.clone())
.with_model_name("model1")
.with_cache_enabled()
.build()
.expect("Failed to build config");
let mut engine = EmbeddingEngine::new(config).expect("Failed to create engine");
let config2 = EngineConfigBuilder::new()
.with_model_path(model_path)
.with_model_name("model2")
.with_normalization_mode(NormalizationMode::None) .build()
.expect("Failed to build config");
engine
.load_model(config2)
.expect("Failed to load second model");
let text = "Same text for different models";
let embedding1 = engine
.embed(Some("model1"), text)
.expect("Failed to generate embedding");
let embedding2 = engine
.embed(Some("model2"), text)
.expect("Failed to generate embedding");
assert_ne!(
embedding1, embedding2,
"Different normalization modes should produce different embeddings"
);
let stats = engine.get_cache_stats().expect("Cache should be enabled");
assert_eq!(stats.entry_count, 2);
let _ = engine
.embed(Some("model1"), text)
.expect("Failed to generate embedding");
let _ = engine
.embed(Some("model2"), text)
.expect("Failed to generate embedding");
let final_stats = engine.get_cache_stats().expect("Cache should be enabled");
assert_eq!(final_stats.hits, 2);
}
#[test]
fn test_cache_key_differences() {
use embellama::PoolingStrategy;
use embellama::cache::embedding_cache::EmbeddingCache;
let text = "Test text";
let model = "test-model";
let key1 =
EmbeddingCache::compute_key(text, model, PoolingStrategy::Mean, NormalizationMode::L2);
let key2 =
EmbeddingCache::compute_key(text, model, PoolingStrategy::Mean, NormalizationMode::None);
let key3 =
EmbeddingCache::compute_key(text, model, PoolingStrategy::Cls, NormalizationMode::L2);
let key4 =
EmbeddingCache::compute_key(text, model, PoolingStrategy::Max, NormalizationMode::L2);
let key5 = EmbeddingCache::compute_key(
text,
"different-model",
PoolingStrategy::Mean,
NormalizationMode::L2,
);
assert_ne!(key1, key2); assert_ne!(key1, key3); assert_ne!(key1, key4); assert_ne!(key1, key5);
let key1_duplicate =
EmbeddingCache::compute_key(text, model, PoolingStrategy::Mean, NormalizationMode::L2);
assert_eq!(key1, key1_duplicate);
}
#[test]
#[ignore = "Requires model file"]
fn test_mixed_batch_cache_hits() {
let Some(model_path) = get_test_model_path() else {
eprintln!("Skipping test: no test model configured");
return;
};
let config = EngineConfigBuilder::new()
.with_model_path(model_path)
.with_model_name("test-model")
.with_cache_enabled()
.build()
.expect("Failed to build config");
let engine = EmbeddingEngine::new(config).expect("Failed to create engine");
let cached_texts = vec!["cached1", "cached2"];
for text in &cached_texts {
engine
.embed(None, text)
.expect("Failed to generate embedding");
}
let mixed_batch = vec![
"cached1", "new1", "cached2", "new2", "cached1", ];
let embeddings = engine
.embed_batch(None, &mixed_batch)
.expect("Failed to generate batch");
assert_eq!(embeddings.len(), 5);
let stats = engine.get_cache_stats().expect("Cache should be enabled");
assert_eq!(stats.entry_count, 4);
println!("Cache stats: {:?}", stats);
}