diff --git a/src/rerank.rs b/src/rerank.rs index 4b0fd9d..7773132 100644 --- a/src/rerank.rs +++ b/src/rerank.rs @@ -1 +1,26 @@ -// implemented in a later task +use crate::ionos::IonosClient; +use crate::model::AppError; + +/// Cross-encoder rerank. Returns a score per document, aligned to the +/// input `docs` order (results are re-sorted back by their index). +pub fn rerank( + client: &IonosClient, model: &str, query: &str, docs: &[String], +) -> Result, AppError> { + if docs.is_empty() { return Ok(vec![]); } + let body = serde_json::json!({ + "model": model, "query": query, "documents": docs + }); + let resp: serde_json::Value = client.post_json("/rerank", &body)?; + let results = resp["results"].as_array() + .ok_or_else(|| AppError::Ionos("rerank: no results[]".into()))?; + let mut scores = vec![0.0f32; docs.len()]; + for r in results { + let idx = r["index"].as_u64().ok_or_else( + || AppError::Ionos("rerank: no index".into()))? as usize; + let sc = r["relevance_score"].as_f64() + .or_else(|| r["score"].as_f64()) + .ok_or_else(|| AppError::Ionos("rerank: no score".into()))? as f32; + if idx < scores.len() { scores[idx] = sc; } + } + Ok(scores) +} diff --git a/tests/rerank_tests.rs b/tests/rerank_tests.rs new file mode 100644 index 0000000..f26ac2d --- /dev/null +++ b/tests/rerank_tests.rs @@ -0,0 +1,24 @@ +use alpha_id::rerank::rerank; +use alpha_id::ionos::IonosClient; +use std::io::Write; + +#[test] +fn rerank_returns_scores_aligned_to_documents() { + let server = tiny_http::Server::http("127.0.0.1:0").unwrap(); + let url = format!("http://{}", server.server_addr()); + let h = std::thread::spawn(move || { + if let Ok(req) = server.recv() { + let body = r#"{"results":[{"index":0,"relevance_score":0.2}, + {"index":1,"relevance_score":0.9}]}"#; + req.respond(tiny_http::Response::from_string(body)).ok(); + } + }); + let mut tf = tempfile::NamedTempFile::new().unwrap(); + write!(tf, "tok").unwrap(); + let c = IonosClient::new(&url, tf.path().to_str().unwrap()).unwrap(); + let scores = rerank(&c, "Qwen/Qwen3-VL-Reranker-8B", "diabetes", + &["herzinfarkt".to_string(), "diabetes typ 2".to_string()]).unwrap(); + assert_eq!(scores.len(), 2); + assert!(scores[1] > scores[0]); + h.join().ok(); +}