mod common;
use embellama::{EmbeddingEngine, EngineConfig, NormalizationMode, PoolingStrategy, RerankResult};
use serial_test::serial;
use std::path::PathBuf;
fn get_rerank_model_path() -> Option<PathBuf> {
std::env::var("EMBELLAMA_TEST_RERANK_MODEL")
.ok()
.map(PathBuf::from)
.filter(|p| p.exists())
}
fn should_run_rerank_tests() -> bool {
get_rerank_model_path().is_some()
}
#[test]
fn test_rerank_result_struct() {
let result = RerankResult {
index: 2,
relevance_score: 0.95,
};
assert_eq!(result.index, 2);
assert!((result.relevance_score - 0.95).abs() < f32::EPSILON);
}
#[test]
fn test_rerank_result_serialization() {
let result = RerankResult {
index: 0,
relevance_score: 0.75,
};
let json = serde_json::to_string(&result).unwrap();
assert!(json.contains("\"index\":0"));
assert!(json.contains("\"relevance_score\":0.75"));
let deserialized: RerankResult = serde_json::from_str(&json).unwrap();
assert_eq!(deserialized.index, result.index);
assert!((deserialized.relevance_score - result.relevance_score).abs() < f32::EPSILON);
}
#[test]
fn test_rerank_result_equality() {
let a = RerankResult {
index: 1,
relevance_score: 0.5,
};
let b = RerankResult {
index: 1,
relevance_score: 0.5,
};
assert_eq!(a, b);
}
#[test]
fn test_pooling_strategy_rank_is_not_default() {
assert_ne!(PoolingStrategy::default(), PoolingStrategy::Rank);
assert_eq!(PoolingStrategy::default(), PoolingStrategy::Mean);
}
#[test]
fn test_pooling_strategy_rank_serde() {
let strategy = PoolingStrategy::Rank;
let json = serde_json::to_string(&strategy).unwrap();
let deserialized: PoolingStrategy = serde_json::from_str(&json).unwrap();
assert_eq!(deserialized, PoolingStrategy::Rank);
}
#[test]
fn test_sigmoid_normalization_properties() {
let sigmoid = |x: f32| 1.0 / (1.0 + (-x).exp());
assert!((sigmoid(0.0) - 0.5).abs() < 1e-6);
assert!(sigmoid(10.0) > 0.999);
assert!(sigmoid(-10.0) < 0.001);
assert!(sigmoid(100.0) > 0.999_999);
assert!(sigmoid(-100.0) < 0.000_001);
assert!(sigmoid(-1.0) < sigmoid(0.0));
assert!(sigmoid(0.0) < sigmoid(1.0));
assert!(sigmoid(1.0) < sigmoid(2.0));
}
#[test]
fn test_engine_config_with_rank_pooling() {
let (_dir, model_path) = common::create_dummy_model();
let config = EngineConfig::builder()
.with_model_path(&model_path)
.with_model_name("reranker")
.with_pooling_strategy(PoolingStrategy::Rank)
.with_normalization_mode(NormalizationMode::None)
.build()
.unwrap();
assert_eq!(
config.model_config.pooling_strategy,
Some(PoolingStrategy::Rank)
);
assert_eq!(
config.model_config.normalization_mode,
Some(NormalizationMode::None)
);
}
#[test]
#[serial]
fn test_rerank_basic() {
if !should_run_rerank_tests() {
eprintln!("Skipping: EMBELLAMA_TEST_RERANK_MODEL not set");
return;
}
common::init_test_logger();
let model_path = get_rerank_model_path().unwrap();
let config = EngineConfig::builder()
.with_model_path(&model_path)
.with_model_name("reranker")
.with_pooling_strategy(PoolingStrategy::Rank)
.with_normalization_mode(NormalizationMode::None)
.with_n_seq_max(4)
.build()
.unwrap();
let engine = EmbeddingEngine::new(config).unwrap();
let query = "What is the capital of France?";
let documents = [
"Paris is the capital and largest city of France.",
"Berlin is the capital of Germany.",
"The weather is nice today.",
];
let results = engine
.rerank(Some("reranker"), query, &documents, None, true)
.unwrap();
assert_eq!(results.len(), 3);
for i in 1..results.len() {
assert!(
results[i - 1].relevance_score >= results[i].relevance_score,
"Results should be sorted descending: {} >= {}",
results[i - 1].relevance_score,
results[i].relevance_score
);
}
assert_eq!(results[0].index, 0, "Paris document should be ranked first");
for r in &results {
assert!(
r.relevance_score >= 0.0 && r.relevance_score <= 1.0,
"Normalized score should be in [0, 1]: {}",
r.relevance_score
);
}
}
#[test]
#[serial]
fn test_rerank_top_n() {
if !should_run_rerank_tests() {
eprintln!("Skipping: EMBELLAMA_TEST_RERANK_MODEL not set");
return;
}
common::init_test_logger();
let model_path = get_rerank_model_path().unwrap();
let config = EngineConfig::builder()
.with_model_path(&model_path)
.with_model_name("reranker")
.with_pooling_strategy(PoolingStrategy::Rank)
.with_normalization_mode(NormalizationMode::None)
.with_n_seq_max(4)
.build()
.unwrap();
let engine = EmbeddingEngine::new(config).unwrap();
let query = "machine learning";
let documents = [
"Deep learning is a subset of machine learning.",
"The stock market rose today.",
"Neural networks power modern AI.",
"I went to the grocery store.",
];
let results = engine
.rerank(Some("reranker"), query, &documents, Some(2), true)
.unwrap();
assert_eq!(results.len(), 2, "Should return only top 2 results");
assert!(results[0].relevance_score >= results[1].relevance_score);
}
#[test]
#[serial]
fn test_rerank_without_normalization() {
if !should_run_rerank_tests() {
eprintln!("Skipping: EMBELLAMA_TEST_RERANK_MODEL not set");
return;
}
common::init_test_logger();
let model_path = get_rerank_model_path().unwrap();
let config = EngineConfig::builder()
.with_model_path(&model_path)
.with_model_name("reranker")
.with_pooling_strategy(PoolingStrategy::Rank)
.with_normalization_mode(NormalizationMode::None)
.with_n_seq_max(4)
.build()
.unwrap();
let engine = EmbeddingEngine::new(config).unwrap();
let query = "What is the capital of France?";
let documents = ["Paris is the capital of France.", "Berlin is in Germany."];
let results = engine
.rerank(Some("reranker"), query, &documents, None, false)
.unwrap();
assert_eq!(results.len(), 2);
assert!(results[0].relevance_score >= results[1].relevance_score);
}
#[test]
#[serial]
fn test_rerank_empty_documents_returns_empty() {
if !should_run_rerank_tests() {
eprintln!("Skipping: EMBELLAMA_TEST_RERANK_MODEL not set");
return;
}
common::init_test_logger();
let model_path = get_rerank_model_path().unwrap();
let config = EngineConfig::builder()
.with_model_path(&model_path)
.with_model_name("reranker")
.with_pooling_strategy(PoolingStrategy::Rank)
.with_normalization_mode(NormalizationMode::None)
.build()
.unwrap();
let engine = EmbeddingEngine::new(config).unwrap();
let documents: &[&str] = &[];
let results = engine
.rerank(Some("reranker"), "query", documents, None, true)
.unwrap();
assert!(results.is_empty());
}
#[test]
#[serial]
fn test_rerank_single_document() {
if !should_run_rerank_tests() {
eprintln!("Skipping: EMBELLAMA_TEST_RERANK_MODEL not set");
return;
}
common::init_test_logger();
let model_path = get_rerank_model_path().unwrap();
let config = EngineConfig::builder()
.with_model_path(&model_path)
.with_model_name("reranker")
.with_pooling_strategy(PoolingStrategy::Rank)
.with_normalization_mode(NormalizationMode::None)
.build()
.unwrap();
let engine = EmbeddingEngine::new(config).unwrap();
let results = engine
.rerank(
Some("reranker"),
"test query",
&["single document"],
None,
true,
)
.unwrap();
assert_eq!(results.len(), 1);
assert_eq!(results[0].index, 0);
assert!(results[0].relevance_score >= 0.0 && results[0].relevance_score <= 1.0);
}
#[test]
#[serial]
fn test_rerank_batch_exceeds_n_seq_max() {
if !should_run_rerank_tests() {
eprintln!("Skipping: EMBELLAMA_TEST_RERANK_MODEL not set");
return;
}
common::init_test_logger();
let model_path = get_rerank_model_path().unwrap();
let config = EngineConfig::builder()
.with_model_path(&model_path)
.with_model_name("reranker")
.with_pooling_strategy(PoolingStrategy::Rank)
.with_normalization_mode(NormalizationMode::None)
.with_n_seq_max(2)
.build()
.unwrap();
let engine = EmbeddingEngine::new(config).unwrap();
let query = "What is machine learning?";
let documents = [
"Machine learning is a branch of AI.",
"The sun is a star.",
"Deep learning uses neural networks.",
"I like pizza.",
"Supervised learning uses labeled data.",
];
let results = engine
.rerank(Some("reranker"), query, &documents, None, true)
.unwrap();
assert_eq!(results.len(), 5);
let mut seen_indices: Vec<usize> = results.iter().map(|r| r.index).collect();
seen_indices.sort();
assert_eq!(seen_indices, vec![0, 1, 2, 3, 4]);
}
#[test]
#[serial]
fn test_embed_on_rank_model_fails() {
if !should_run_rerank_tests() {
eprintln!("Skipping: EMBELLAMA_TEST_RERANK_MODEL not set");
return;
}
common::init_test_logger();
let model_path = get_rerank_model_path().unwrap();
let config = EngineConfig::builder()
.with_model_path(&model_path)
.with_model_name("reranker")
.with_pooling_strategy(PoolingStrategy::Rank)
.with_normalization_mode(NormalizationMode::None)
.build()
.unwrap();
let engine = EmbeddingEngine::new(config).unwrap();
let result = engine.embed(Some("reranker"), "test text");
assert!(result.is_err());
assert!(
result
.unwrap_err()
.to_string()
.contains("PoolingStrategy::Rank")
);
}
#[test]
#[serial]
fn test_rerank_on_embedding_model_fails() {
if !common::should_run_model_tests() {
eprintln!("Skipping: EMBELLAMA_TEST_MODEL not set");
return;
}
common::init_test_logger();
let model_path = common::get_test_model_path().unwrap();
let config = EngineConfig::builder()
.with_model_path(&model_path)
.with_model_name("embedder")
.with_pooling_strategy(PoolingStrategy::Mean)
.build()
.unwrap();
let engine = EmbeddingEngine::new(config).unwrap();
let result = engine.rerank(Some("embedder"), "query", &["document"], None, true);
assert!(result.is_err());
assert!(
result
.unwrap_err()
.to_string()
.contains("PoolingStrategy::Rank")
);
}
#[test]
#[serial]
fn test_rerank_auto_detect_from_gguf() {
if !should_run_rerank_tests() {
eprintln!("Skipping: EMBELLAMA_TEST_RERANK_MODEL not set");
return;
}
common::init_test_logger();
let model_path = get_rerank_model_path().unwrap();
let config = EngineConfig::builder()
.with_model_path(&model_path)
.with_model_name("reranker-auto")
.with_n_seq_max(4)
.build()
.unwrap();
let engine = EmbeddingEngine::new(config).unwrap();
let query = "What is the capital of France?";
let documents = [
"Paris is the capital and largest city of France.",
"Berlin is the capital of Germany.",
"The weather is nice today.",
];
let results = engine
.rerank(Some("reranker-auto"), query, &documents, None, true)
.unwrap();
assert_eq!(results.len(), 3);
for i in 1..results.len() {
assert!(
results[i - 1].relevance_score >= results[i].relevance_score,
"Results should be sorted descending"
);
}
assert_eq!(results[0].index, 0, "Paris document should be ranked first");
}