#![forbid(unsafe_code)]
use renkin::reranker::LightGbmModel;
use serde::{Deserialize, Serialize};
use std::io::{BufRead, Write};
#[derive(Deserialize)]
struct InRow {
features: Vec<Option<f64>>,
}
#[derive(Serialize)]
struct OutRow {
rust_score: f64,
}
fn arg_value(flag: &str) -> Option<String> {
let args: Vec<String> = std::env::args().collect();
args.iter()
.position(|a| a == flag)
.and_then(|i| args.get(i + 1))
.cloned()
}
fn main() {
let model_path = arg_value("--model").unwrap_or_else(|| {
eprintln!("Usage: renkin-reranker-predict --model <path/to/model.txt> < features.jsonl");
std::process::exit(2);
});
let model = LightGbmModel::from_path(&model_path)
.unwrap_or_else(|e| panic!("load model {model_path}: {e}"));
let stdin = std::io::stdin();
let stdout = std::io::stdout();
let mut out = stdout.lock();
for (i, line) in stdin.lock().lines().enumerate() {
let line = line.unwrap_or_else(|e| panic!("read stdin line {i}: {e}"));
if line.trim().is_empty() {
continue;
}
let row: InRow =
serde_json::from_str(&line).unwrap_or_else(|e| panic!("parse line {i}: {e}"));
let features: Vec<f64> = row.features.iter().map(|v| v.unwrap_or(f64::NAN)).collect();
let score = model
.predict(&features)
.unwrap_or_else(|e| panic!("predict line {i}: {e}"));
let out_row = OutRow { rust_score: score };
writeln!(out, "{}", serde_json::to_string(&out_row).unwrap())
.unwrap_or_else(|e| panic!("write stdout line {i}: {e}"));
}
}