132 lines
4.2 KiB
Rust
132 lines
4.2 KiB
Rust
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("/rerank"))
|
|||
|
|
.respond_with(ResponseTemplate::new(200).set_body_json(vec![
|
|||
|
|
serde_json::json!({"index": 0, "score": 0.98}),
|
|||
|
|
serde_json::json!({"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("/rerank"))
|
|||
|
|
.respond_with(ResponseTemplate::new(200).set_body_json(vec![
|
|||
|
|
serde_json::json!({"index": 1, "score": 0.99}),
|
|||
|
|
serde_json::json!({"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("/rerank"))
|
|||
|
|
.respond_with(ResponseTemplate::new(200).set_body_json(vec![
|
|||
|
|
serde_json::json!({"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),
|
|||
|
|
}
|
|||
|
|
}
|