1use std::collections::HashMap;
9
10use rusqlite::Connection;
11
12use crate::search::Hit;
13
14pub const SUCCESS_ALPHA: f32 = 0.3;
18
19pub 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 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 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}