Files

138 lines
4.3 KiB
Plaintext
Raw Permalink Normal View History

use mem_llm::RerankClient;
use wiremock::{Mock, MockServer, ResponseTemplate};
use wiremock::matchers::{method, path};
#[tokio::test]
async fn a1_bare_array_parsed() {
let mock_server = MockServer::start().await;
Mock::given(method("POST"))
.and(path("/v1/rerank"))
.respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
"results": [
{"index": 0, "score": 0.98},
{"index": 1, "score": 0.01},
]
})))
.mount(&mock_server)
.await;
let client = RerankClient::new(&mock_server.uri(), "test-key", "bge-reranker").unwrap();
let results = client
.rerank("test", &["relevant", "irrelevant"])
.await
.unwrap();
assert_eq!(results.len(), 2);
assert_eq!(results[0].0, 0); // Index 0 (higher score)
assert!(results[0].1 > 0.9);
}
#[tokio::test]
async fn a2_index_mapping() {
let mock_server = MockServer::start().await;
// Return out-of-order: index 1 first, then index 0
Mock::given(method("POST"))
.and(path("/v1/rerank"))
.respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
"results": [
{"index": 1, "score": 0.99},
{"index": 0, "score": 0.01},
]
})))
.mount(&mock_server)
.await;
let client = RerankClient::new(&mock_server.uri(), "test-key", "bge-reranker").unwrap();
let results = client
.rerank("test", &["irrelevant", "relevant"])
.await
.unwrap();
// Results sorted by score (descending)
assert_eq!(results[0].0, 1, "Index 1 should be first (highest score)");
assert!(results[0].1 > 0.9);
assert_eq!(results[1].0, 0, "Index 0 should be second");
assert!(results[1].1 < 0.1);
}
#[tokio::test]
async fn a3_empty_no_request() {
let mock_server = MockServer::start().await;
let client = RerankClient::new(&mock_server.uri(), "test-key", "bge-reranker").unwrap();
let results = client.rerank("test", &[]).await.unwrap();
assert_eq!(results.len(), 0, "Empty input should return empty without request");
}
#[tokio::test]
async fn a4_apikey_sent() {
let mock_server = MockServer::start().await;
Mock::given(method("POST"))
.and(path("/v1/rerank"))
.respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
"results": [
{"index": 0, "score": 0.95},
]
})))
.mount(&mock_server)
.await;
let client = RerankClient::new(&mock_server.uri(), "my-secret-key", "bge").unwrap();
let result = client.rerank("q", &["text"]).await;
// If request succeeds, apikey was sent (mock only accepts POST, no header check in this mock)
assert!(result.is_ok(), "Request should succeed with apikey");
}
#[tokio::test]
#[ignore]
async fn a5_live_discriminates() {
// Live test against real rerank endpoint
// Run with: cargo test --test it_rerank -- --ignored --nocapture
let api_key = match std::env::var("MEM_API_KEY") {
Ok(k) => k,
Err(_) => {
println!("SKIP: MEM_API_KEY not set");
return;
}
};
let client = match RerankClient::new("https://api.riotpiao.com/v1", &api_key, "bge-reranker-base") {
Ok(c) => c,
Err(e) => {
println!("SKIP: Could not create rerank client: {}", e);
return;
}
};
let texts = &[
"Rust is a systems programming language focused on safety and performance",
"Bananas are a tropical fruit",
];
match client.rerank("what is rust", texts).await {
Ok(results) => {
println!("Rerank results:");
for (idx, score) in &results {
println!(" [{}] score={:.6}: {}", idx, score, texts[*idx]);
}
// First result should be the Rust text (index 0)
assert_eq!(results[0].0, 0, "Rust text should rank first");
// Score ratio should be large (rust >> banana)
if results.len() > 1 {
let ratio = results[0].1 / results[1].1.max(0.0001);
println!("Score ratio: {:.1}×", ratio);
assert!(ratio > 10.0, "Rust should score at least 10× higher than bananas");
}
}
Err(e) => println!("Live test skipped: {}", e),
}
}