- Add ingest_worker and query_worker modules to main.rs - Update rerank tests to use /v1/rerank endpoint with new response format - All 253+ tests now passing
138 lines
4.3 KiB
Rust
138 lines
4.3 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("/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),
|
||
}
|
||
}
|