Skip to main content

toolhub_recommender/
rerank.rs

1//! Score-aware reranking, Phase 4 PLAN §8.3.
2//!
3//! After the hybrid (cosine + BM25) combine produces candidate `Hit`s, the
4//! reranker boosts tools with proven track records. Tools without enough
5//! samples or without a `tool_scores` row are passed through unchanged, so
6//! Phase 1–3 behaviour is preserved when telemetry is empty.
7
8use std::collections::HashMap;
9
10use rusqlite::Connection;
11
12use crate::search::Hit;
13
14/// Default boost factor — chosen so a tool with 100 % success rate gets a
15/// 1.3× score multiplier. PLAN §10 #5: track recommender accuracy on the
16/// benchmark set and tune.
17pub const SUCCESS_ALPHA: f32 = 0.3;
18
19/// Minimum sample size before we trust a tool's success_rate enough to apply
20/// the boost. Below this, the score is statistically noisy.
21pub const MIN_SAMPLE_SIZE: i64 = 5;
22
23pub trait Reranker {
24    fn apply(&self, hits: &mut [Hit], conn: &Connection) -> anyhow::Result<()>;
25}
26
27#[derive(Debug, Clone, Copy)]
28pub struct SuccessReranker {
29    pub alpha: f32,
30    pub min_samples: i64,
31}
32
33impl Default for SuccessReranker {
34    fn default() -> Self {
35        Self {
36            alpha: SUCCESS_ALPHA,
37            min_samples: MIN_SAMPLE_SIZE,
38        }
39    }
40}
41
42impl Reranker for SuccessReranker {
43    fn apply(&self, hits: &mut [Hit], conn: &Connection) -> anyhow::Result<()> {
44        if hits.is_empty() {
45            return Ok(());
46        }
47        let scores = load_scores(conn, hits)?;
48        for hit in hits.iter_mut() {
49            if let Some((rate, samples)) = scores.get(&hit.tool_id)
50                && *samples >= self.min_samples
51            {
52                hit.score *= 1.0 + self.alpha * (*rate as f32);
53            }
54        }
55        // Re-sort — ordering may have changed.
56        hits.sort_by(|a, b| {
57            b.score
58                .partial_cmp(&a.score)
59                .unwrap_or(std::cmp::Ordering::Equal)
60        });
61        Ok(())
62    }
63}
64
65fn load_scores(conn: &Connection, hits: &[Hit]) -> anyhow::Result<HashMap<String, (f64, i64)>> {
66    let ids: Vec<&str> = hits.iter().map(|h| h.tool_id.as_str()).collect();
67    let placeholders = std::iter::repeat_n("?", ids.len())
68        .collect::<Vec<_>>()
69        .join(",");
70    let sql = format!(
71        "SELECT tool_id, success_rate, sample_size
72         FROM tool_scores
73         WHERE tool_id IN ({placeholders})"
74    );
75    let mut stmt = conn.prepare(&sql)?;
76    let params: Vec<&dyn rusqlite::ToSql> = ids.iter().map(|s| s as &dyn rusqlite::ToSql).collect();
77    let rows = stmt
78        .query_map(params.as_slice(), |row| {
79            Ok((
80                row.get::<_, String>(0)?,
81                row.get::<_, Option<f64>>(1)?,
82                row.get::<_, Option<i64>>(2)?,
83            ))
84        })?
85        .collect::<Result<Vec<_>, _>>()?;
86    let mut out = HashMap::with_capacity(rows.len());
87    for (id, rate, n) in rows {
88        if let (Some(r), Some(n)) = (rate, n) {
89            out.insert(id, (r, n));
90        }
91    }
92    Ok(out)
93}
94
95#[cfg(test)]
96mod tests {
97    use super::*;
98    use rusqlite::params;
99
100    fn open_with_schema() -> Connection {
101        let conn = Connection::open_in_memory().unwrap();
102        conn.execute_batch(
103            "CREATE TABLE tool_scores (
104                tool_id TEXT PRIMARY KEY,
105                success_rate REAL,
106                sample_size INTEGER,
107                avg_cost_usd REAL,
108                median_duration_ms INTEGER,
109                score_updated_at TEXT
110            );",
111        )
112        .unwrap();
113        conn
114    }
115
116    fn seed(conn: &Connection, id: &str, rate: f64, n: i64) {
117        conn.execute(
118            "INSERT INTO tool_scores VALUES (?, ?, ?, NULL, NULL, '2026-05-03T00:00:00Z')",
119            params![id, rate, n],
120        )
121        .unwrap();
122    }
123
124    #[test]
125    fn alpha_zero_is_identity() {
126        let conn = open_with_schema();
127        seed(&conn, "skill:a", 1.0, 100);
128        let mut hits = vec![Hit {
129            tool_id: "skill:a".into(),
130            score: 0.5,
131        }];
132        let rer = SuccessReranker {
133            alpha: 0.0,
134            min_samples: 5,
135        };
136        rer.apply(&mut hits, &conn).unwrap();
137        assert!((hits[0].score - 0.5).abs() < 1e-6);
138    }
139
140    #[test]
141    fn boost_only_when_min_samples_met() {
142        let conn = open_with_schema();
143        seed(&conn, "skill:trusted", 1.0, 10);
144        seed(&conn, "skill:noisy", 1.0, 2);
145        let mut hits = vec![
146            Hit {
147                tool_id: "skill:trusted".into(),
148                score: 0.5,
149            },
150            Hit {
151                tool_id: "skill:noisy".into(),
152                score: 0.5,
153            },
154        ];
155        let rer = SuccessReranker::default();
156        rer.apply(&mut hits, &conn).unwrap();
157        let trusted = hits.iter().find(|h| h.tool_id == "skill:trusted").unwrap();
158        let noisy = hits.iter().find(|h| h.tool_id == "skill:noisy").unwrap();
159        assert!(trusted.score > noisy.score);
160        assert!((trusted.score - 0.5 * (1.0 + 0.3)).abs() < 1e-6);
161        assert!((noisy.score - 0.5).abs() < 1e-6);
162    }
163
164    #[test]
165    fn missing_score_passes_through() {
166        let conn = open_with_schema();
167        let mut hits = vec![Hit {
168            tool_id: "skill:no-data".into(),
169            score: 0.7,
170        }];
171        let rer = SuccessReranker::default();
172        rer.apply(&mut hits, &conn).unwrap();
173        assert!((hits[0].score - 0.7).abs() < 1e-6);
174    }
175
176    #[test]
177    fn boost_can_change_ordering() {
178        let conn = open_with_schema();
179        seed(&conn, "skill:underdog", 1.0, 50);
180        let mut hits = vec![
181            Hit {
182                tool_id: "skill:leader".into(),
183                score: 0.6,
184            },
185            Hit {
186                tool_id: "skill:underdog".into(),
187                score: 0.5,
188            },
189        ];
190        let rer = SuccessReranker::default();
191        rer.apply(&mut hits, &conn).unwrap();
192        // underdog now 0.5 * 1.3 = 0.65 → wins.
193        assert_eq!(hits[0].tool_id, "skill:underdog");
194    }
195
196    #[test]
197    fn empty_hits_is_noop() {
198        let conn = open_with_schema();
199        let mut hits: Vec<Hit> = Vec::new();
200        SuccessReranker::default().apply(&mut hits, &conn).unwrap();
201        assert!(hits.is_empty());
202    }
203}